From 047ae41c0d9fe3ae5656d1541a376d6fbf17c9a3 Mon Sep 17 00:00:00 2001 From: James Vecellio-Grant <159560811+Discountchubbs@users.noreply.github.com> Date: Fri, 4 Sep 2026 21:15:00 -0700 Subject: [PATCH] modeld_v2: one dev warp and enqueue (#1990) --- .github/workflows/sunnypilot-build-model.yaml | 2 +- .../sunnypilot/modeld_v2/compile_modeld.py | 16 +++--- openpilot/sunnypilot/modeld_v2/modeld.py | 56 +++++++------------ openpilot/sunnypilot/models/fetcher.py | 2 +- 4 files changed, 32 insertions(+), 44 deletions(-) diff --git a/.github/workflows/sunnypilot-build-model.yaml b/.github/workflows/sunnypilot-build-model.yaml index 10be63ee22..837c53be69 100644 --- a/.github/workflows/sunnypilot-build-model.yaml +++ b/.github/workflows/sunnypilot-build-model.yaml @@ -188,7 +188,7 @@ jobs: if [ "${{ inputs.target_hardware }}" == "chestnut" ]; then echo "CHESTNUT build" export CHESTNUT=1 - TG_FLAGS="DEBUG=1 DEV=USB+AMD:LLVM WARP_DEV=QCOM FLOAT16=1 JIT_BATCH_SIZE=0 GMMU=0 TC_OPT=2" + TG_FLAGS="DEBUG=1 DEV=USB+AMD:LLVM FLOAT16=1 JIT_BATCH_SIZE=0 GMMU=0 TC_OPT=2 TC_OCCUPANCY_OPT=1" OUTPUT_PKL="${{ env.MODELS_DIR }}/big_driving_tinygrad.pkl" else echo "QCOM build" diff --git a/openpilot/sunnypilot/modeld_v2/compile_modeld.py b/openpilot/sunnypilot/modeld_v2/compile_modeld.py index 7a58352c36..a274097563 100755 --- a/openpilot/sunnypilot/modeld_v2/compile_modeld.py +++ b/openpilot/sunnypilot/modeld_v2/compile_modeld.py @@ -41,7 +41,6 @@ from tinygrad.tensor import Tensor MODEL_TYPES = ('vision_policy', 'supercombo', 'vision_multi_policy') WARP_INPUTS = ['tfm', 'big_tfm'] POLICY_INPUTS = ['img_q', 'big_img_q', 'feat_q', 'desire_q', 'packed_npy_inputs'] -WARP_DEV = os.getenv('WARP_DEV') def _detect_desire_key(shapes: dict) -> str | None: @@ -154,12 +153,13 @@ def make_warp_queues(device=Device.DEFAULT): def make_warp(nv12: NV12Frame, model_w: int, model_h: int): frame_prepare = make_frame_prepare(nv12, model_w, model_h) - WARP_DEV = os.getenv('WARP_DEV', Device.DEFAULT) def warp(tfm, big_tfm, frame, big_frame): - tfm = tfm.to(WARP_DEV) - big_tfm = big_tfm.to(WARP_DEV) - Tensor.realize(tfm, big_tfm) + tfm = tfm.to(Device.DEFAULT) + big_tfm = big_tfm.to(Device.DEFAULT) + frame = frame.to(Device.DEFAULT) + big_frame = big_frame.to(Device.DEFAULT) + Tensor.realize(tfm, big_tfm, frame, big_frame) warped_frame = frame_prepare(frame, tfm).unsqueeze(0) warped_big_frame = frame_prepare(big_frame, big_tfm).unsqueeze(0) @@ -368,16 +368,18 @@ if __name__ == "__main__": run_policy_func = make_run_policy(vision_runner, policy_runners, features_slice, derived_frame_skip, all_shapes) run_policy_jit = TinyJit(run_policy_func, prune=True) make_policy_queues = partial(generate_queues_and_npy, all_shapes, derived_frame_skip, is_supercombo=is_supercombo) - make_random_model_inputs = partial(make_random_images, keys=['warped'], shape=(2, 6, model_h // 2, model_w // 2), device=WARP_DEV) + make_random_model_inputs = partial(make_random_images, keys=['warped'], shape=(2, 6, model_h // 2, model_w // 2), device=Device.DEFAULT) output_data['run_policy'] = compile_jit(run_policy_jit, make_random_model_inputs, POLICY_INPUTS, make_policy_queues) for cam_w, cam_h in args.camera_resolutions: print(f"Compiling warp JIT for {cam_w}x{cam_h}...") nv12 = NV12Frame(cam_w, cam_h, *get_nv12_info(cam_w, cam_h)) - make_random_warp_inputs = partial(make_random_images, keys=['frame', 'big_frame'], shape=nv12.size, device=WARP_DEV) + make_random_warp_inputs = partial(make_random_images, keys=['frame', 'big_frame'], shape=nv12.size, device=Device.DEFAULT) warp = TinyJit(make_warp(nv12, model_w, model_h), prune=True) output_data[(cam_w, cam_h)] = compile_jit(warp, make_random_warp_inputs, WARP_INPUTS, make_warp_queues) + output_data['metadata']['warp_dev'] = Device.DEFAULT + with open(args.output, "wb") as file: dump_oob(output_data, file) diff --git a/openpilot/sunnypilot/modeld_v2/modeld.py b/openpilot/sunnypilot/modeld_v2/modeld.py index cd421d70d5..2ab96934ef 100755 --- a/openpilot/sunnypilot/modeld_v2/modeld.py +++ b/openpilot/sunnypilot/modeld_v2/modeld.py @@ -6,6 +6,7 @@ This file is part of sunnypilot and is licensed under the MIT License. See the LICENSE.md file in the root directory for more details. """ +from collections.abc import Callable import os os.environ['GMMU'] = '0' import numpy as np @@ -110,18 +111,12 @@ class ModelState(ModelStateBase): cloudlog.warning(f"loading combined pkl: {pkl_path}") jits = load_oob(open_file_chunked(pkl_path)) - self.WARP_DEV = 'QCOM' if COMMA_HARDWARE else 'CPU' - self.DEV = 'AMD' if self.chestnut else self.WARP_DEV - self.QUEUE_DEV = self.DEV metadata = jits['metadata'] - - self.is_legacy_model = 'run_policy' not in jits # remove after next recompile - if self.is_legacy_model: - self.warp = jits[(cam_w, cam_h)]['warp_enqueue'] - self.run_policy = jits[(cam_w, cam_h)]['run_policy'] - else: - self.run_policy = jits['run_policy'] - self.warp = jits[(cam_w, cam_h)] + self.WARP_DEV = metadata.get('warp_dev', 'QCOM' if COMMA_HARDWARE else 'CPU') + self.DEV = 'AMD' if self.chestnut else ('QCOM' if COMMA_HARDWARE else 'CPU') + self.QUEUE_DEV = self.DEV + self.run_policy = jits['run_policy'] + self.warp = jits[(cam_w, cam_h)] if 'model' in metadata: model_metadata = metadata['model'] @@ -180,11 +175,7 @@ class ModelState(ModelStateBase): yuv_size = self.frame_buf_params[self._road_key][3] frame_tensor = Tensor(np.zeros(yuv_size, dtype=np.uint8), device=self.WARP_DEV).contiguous().realize() big_frame_tensor = Tensor(np.zeros(yuv_size, dtype=np.uint8), device=self.WARP_DEV).contiguous().realize() - - if self.is_legacy_model: # Remove this conditional hack after recompile - self.warp(**self.input_queues, frame=frame_tensor, big_frame=big_frame_tensor) - else: - self.warp(**{k: self.input_queues[k] for k in WARP_INPUTS}, frame=frame_tensor, big_frame=big_frame_tensor) + self.warp(**{k: self.input_queues[k] for k in WARP_INPUTS}, frame=frame_tensor, big_frame=big_frame_tensor) def warmup(self) -> None: dummy_frames = {k: np.zeros(self.frame_buf_params[k][3], dtype=np.uint8) for k in self._vision_input_names} @@ -217,7 +208,8 @@ class ModelState(ModelStateBase): return self._desire_key def run(self, bufs: dict[str, VisionBuf], transforms: dict[str, np.ndarray], - inputs: dict[str, np.ndarray], prepare_only: bool) -> dict[str, np.ndarray] | None: + inputs: dict[str, np.ndarray], prepare_only: bool, + after_enqueue: Callable[[], None] | None = None) -> dict[str, np.ndarray] | None: for key in bufs.keys(): ptr = np.frombuffer(bufs[key].data, dtype=np.uint8).ctypes.data yuv_size = self.frame_buf_params[key][3] @@ -239,20 +231,18 @@ class ModelState(ModelStateBase): self.numpy_inputs['tfm'][:, :] = transforms[road_key].reshape(3, 3) self.numpy_inputs['big_tfm'][:, :] = transforms[wide_key].reshape(3, 3) - if self.is_legacy_model: # remove after next recompile - if prepare_only: - self.warp(**self.input_queues, frame=self.full_frames[road_key], big_frame=self.full_frames[wide_key]) - return None - raw_outputs = self.run_policy(**self.input_queues, frame=self.full_frames[road_key], big_frame=self.full_frames[wide_key]) - else: - if prepare_only: - self.warp(**{k: self.input_queues[k] for k in WARP_INPUTS}, frame=self.full_frames[road_key], big_frame=self.full_frames[wide_key]) - return None - warped = self.warp(**{k: self.input_queues[k] for k in WARP_INPUTS}, frame=self.full_frames[road_key], big_frame=self.full_frames[wide_key]) - raw_outputs = self.run_policy(**{k: self.input_queues[k] for k in POLICY_INPUTS if k in self.input_queues}, warped=warped) + if prepare_only: + self.warp(**{k: self.input_queues[k] for k in WARP_INPUTS}, frame=self.full_frames[road_key], big_frame=self.full_frames[wide_key]) + return None + warped = self.warp(**{k: self.input_queues[k] for k in WARP_INPUTS}, frame=self.full_frames[road_key], big_frame=self.full_frames[wide_key]) + raw_outputs = self.run_policy(**{k: self.input_queues[k] for k in POLICY_INPUTS if k in self.input_queues}, warped=warped) + if after_enqueue is not None: + after_enqueue() if self._combined_model_type == 'supercombo': model_output = raw_outputs.numpy().flatten() + if self.chestnut and not np.all(np.isfinite(model_output)): + raise RuntimeError("model output not finite") sliced = {k: model_output[np.newaxis, v] for k, v in self.vision_output_slices.items()} outputs = self.parser.parse_outputs(sliced) if 'prev_feat' in self.numpy_inputs: @@ -285,9 +275,6 @@ class ModelState(ModelStateBase): buf[0, :-1] = buf[0, 1:] buf[0, -1, :] = outputs['desired_curvature'][0, :] if not self.mlsim else 0 - if self.chestnut and not np.all(np.isfinite(outputs.get('plan', np.array([0.])))): - raise RuntimeError("model output not finite") - return outputs def get_action_from_model(self, model_output: dict[str, np.ndarray], prev_action: log.ModelDataV2.Action, @@ -512,7 +499,9 @@ def main(demo=False): mt1 = time.perf_counter() try: - model_output = model.run(bufs, transforms, inputs, prepare_only) + send_chestnut = (chestnut_state is not None and + run_count % round(model.constants.MODEL_FREQ / SERVICE_LIST['chestnutState'].frequency) == 0) + model_output = model.run(bufs, transforms, inputs, prepare_only, chestnut_state.send if send_chestnut else None) except Exception: if not params.get_bool("ChestnutActive"): raise @@ -559,9 +548,6 @@ def main(demo=False): pm.send('modelDataV2SP', mdv2sp_send) last_vipc_frame_id = meta_main.frame_id - if chestnut_state is not None and run_count % round(model.constants.MODEL_FREQ / SERVICE_LIST['chestnutState'].frequency) == 0: - chestnut_state.send() - if __name__ == "__main__": try: import argparse diff --git a/openpilot/sunnypilot/models/fetcher.py b/openpilot/sunnypilot/models/fetcher.py index f997dbbc51..810b9b7e49 100644 --- a/openpilot/sunnypilot/models/fetcher.py +++ b/openpilot/sunnypilot/models/fetcher.py @@ -139,7 +139,7 @@ class ModelCache: class ModelFetcher: """Handles fetching and caching of model data from remote source""" MODEL_URL = "https://raw.githubusercontent.com/sunnypilot/sunnypilot-models/refs/heads/gh-pages/docs/driving_models_v22.json" - MODEL_URL_CHESTNUT = "https://raw.githubusercontent.com/sunnypilot/sunnypilot-models/refs/heads/gh-pages/docs/driving_models_chestnut_v23.json" + MODEL_URL_CHESTNUT = "https://raw.githubusercontent.com/sunnypilot/sunnypilot-models/refs/heads/gh-pages/docs/driving_models_chestnut_v24.json" MODEL_SOURCES = { "qcom": (MODEL_URL, ""),