diff --git a/selfdrive/modeld/compile_modeld.py b/selfdrive/modeld/compile_modeld.py index af8cb482f9..75bb0945ba 100644 --- a/selfdrive/modeld/compile_modeld.py +++ b/selfdrive/modeld/compile_modeld.py @@ -467,7 +467,54 @@ def make_run_supercombo(model_runner, metadata, frame_skip, image_history_pipeli return run_policy -def compile_jit(jit, make_random_inputs, input_keys, make_queues): +def stateful_image_shapes(metadata): + shape = tuple(metadata['input_shapes']['new_img']) + if len(shape) != 4 or shape[:2] != (2, 6): + raise ValueError(f"Unsupported stateful image shape: {shape}") + return dict.fromkeys(('img', 'big_img'), (1, *shape[1:])) + + +def stateful_host_shapes(metadata): + return {name: shape for name, shape in metadata['input_shapes'].items() + if name != 'new_img' and name not in metadata['state_pairs']} + + +def make_stateful_input_queues(metadata, device): + queues, npy = make_warp_input_queues(stateful_image_shapes(metadata), 1, device) + shapes = stateful_host_shapes(metadata) + sizes = [math.prod(shape) for shape in shapes.values()] + packed = np.zeros(sum(sizes), dtype=np.float32) + npy.update({name: value.reshape(shape) for (name, shape), value in + zip(shapes.items(), np.split(packed, np.cumsum(sizes[:-1])), strict=True)}) + queues['packed_npy_inputs'] = Tensor(packed, device='NPY').realize() + for name in metadata['state_pairs']: + queues[name] = Tensor(np.zeros(metadata['input_shapes'][name], dtype=metadata['input_dtypes'][name]), + device=device).contiguous().realize() + return queues, npy + + +def make_run_stateful_supercombo(model_runner, metadata): + shapes = stateful_host_shapes(metadata) + sizes = [math.prod(shape) for shape in shapes.values()] + + def run_policy(warped, packed_npy_inputs, **state): + packed = packed_npy_inputs.to(Device.DEFAULT).realize() + inputs = {name: value.reshape(shape).cast(model_runner.graph_inputs[name].dtype) + for (name, shape), value in zip(shapes.items(), packed.split(sizes), strict=True)} + inputs['new_img'] = warped.to(Device.DEFAULT).cast(model_runner.graph_inputs['new_img'].dtype) + outputs = {name: value.contiguous() for name, value in model_runner(inputs | state).items()} + for name, next_name in metadata['state_pairs'].items(): + if outputs[next_name].dtype != state[name].dtype: + raise ValueError(f'State dtype mismatch: {name} -> {next_name}') + # All reads of the previous state must finish before updating any history buffer. + Tensor.realize(*outputs.values()) + Tensor.realize(*(state[name].assign(outputs[next_name]) for name, next_name in metadata['state_pairs'].items())) + return outputs['outputs'].cast('float32'), + + return run_policy + + +def compile_jit(jit, make_random_inputs, input_keys, make_queues, validation_runs=1): seed = 42 def random_inputs_run(fn, current_seed, test_values=None, test_buffers=None, expect_match=True): @@ -475,7 +522,8 @@ def compile_jit(jit, make_random_inputs, input_keys, make_queues): np.random.seed(current_seed) Tensor.manual_seed(current_seed) testing = test_values is not None or test_buffers is not None - run_count = 1 if testing else 3 + run_count = validation_runs if testing else max(3, validation_runs) + values, buffers = [], [] for index in range(run_count): for value in npy.values(): @@ -489,9 +537,9 @@ def compile_jit(jit, make_random_inputs, input_keys, make_queues): end = time.perf_counter() print(f" [{index + 1}/{run_count}] enqueue {(mid - start) * 1e3:6.2f} ms -- total {(end - start) * 1e3:6.2f} ms") - if index == 0: - values = [np.copy(value.numpy()) for value in outputs] - buffers = [np.copy(value.numpy()) for value in input_queues.values()] + if index < validation_runs: + values.extend(np.copy(value.numpy()) for value in outputs) + buffers.extend(np.copy(value.numpy()) for value in input_queues.values()) if not all(np.isfinite(value).all() for value in values): raise ValueError("Compiled JIT produced non-finite outputs") @@ -602,13 +650,31 @@ def main(): output["metadata"]["model"] = make_metadata_dict(model_path) validate_metadata(output["metadata"]["model"]) policy_shapes = output["metadata"]["model"]["input_shapes"] - frame_skip = args.frame_skip or derive_frame_skip(policy_shapes) - make_policy_queues = partial(make_supercombo_input_queues, policy_shapes, frame_skip) - run_policy = make_run_supercombo( - model_runner, output["metadata"], frame_skip, args.image_history_pipeline, - ) - image_shapes = policy_shapes - policy_input_keys = FAST_POLICY_INPUTS if args.image_history_pipeline == IMAGE_HISTORY_IN_POLICY else SUPERCOMBO_POLICY_INPUTS + if 'new_img' in policy_shapes: + if args.image_history_pipeline != IMAGE_HISTORY_IN_POLICY: + parser.error('ONNX-managed history requires --image-history-pipeline policy') + metadata = output['metadata']['model'] + metadata['state_pairs'] = {name: f'next_{name}' for name in policy_shapes + if f'next_{name}' in metadata['output_shapes']} + if not metadata['state_pairs']: + raise ValueError('Stateful supercombo is missing next-state outputs') + metadata['input_dtypes'] = {name: np.dtype(spec.dtype.fmt).name for name, spec in model_runner.graph_inputs.items()} + for name, next_name in metadata['state_pairs'].items(): + if policy_shapes[name] != metadata['output_shapes'][next_name]: + raise ValueError(f'State shape mismatch: {name} -> {next_name}') + frame_skip = 1 + make_policy_queues = partial(make_stateful_input_queues, metadata) + run_policy = make_run_stateful_supercombo(model_runner, metadata) + image_shapes = stateful_image_shapes(metadata) + policy_input_keys = ('packed_npy_inputs', *metadata['state_pairs']) + else: + frame_skip = args.frame_skip or derive_frame_skip(policy_shapes) + make_policy_queues = partial(make_supercombo_input_queues, policy_shapes, frame_skip) + run_policy = make_run_supercombo( + model_runner, output["metadata"], frame_skip, args.image_history_pipeline, + ) + image_shapes = policy_shapes + policy_input_keys = FAST_POLICY_INPUTS if args.image_history_pipeline == IMAGE_HISTORY_IN_POLICY else SUPERCOMBO_POLICY_INPUTS else: if not args.vision_onnx: parser.error("--vision-onnx is required for split models") @@ -675,6 +741,7 @@ def main(): ) output["run_policy"] = compile_jit( run_policy_jit, make_random_model_inputs, policy_input_keys, make_policy_queues, + validation_runs=5 if output['metadata'].get('model', {}).get('state_pairs') else 1, ) model_w, model_h = args.model_size diff --git a/selfdrive/modeld/modeld.py b/selfdrive/modeld/modeld.py index 797b29ea25..1073b72b71 100755 --- a/selfdrive/modeld/modeld.py +++ b/selfdrive/modeld/modeld.py @@ -48,6 +48,9 @@ from openpilot.selfdrive.modeld.compile_modeld import ( derive_frame_skip, make_split_input_queues, make_supercombo_input_queues, + make_stateful_input_queues, + stateful_host_shapes, + stateful_image_shapes, ) from openpilot.selfdrive.modeld.helpers import get_tg_input_devices, load_oob, tinygrad_dev_config, usbgpu_present from openpilot.selfdrive.modeld.usbgpu_link import wait_usbgpu_link @@ -587,8 +590,15 @@ class ModelState: self.run_policy = artifact["run_policy"] self.warp_enqueue = artifact[(cam_w, cam_h)] self.can_prepare_only = self.image_history_pipeline == IMAGE_HISTORY_IN_WARP + self.onnx_history = self.model_type == 'supercombo' and bool(self.metadata['model'].get('state_pairs')) - if self.model_type == "supercombo": + if self.onnx_history: + metadata = self.metadata['model'] + input_shapes = stateful_image_shapes(metadata) + self.output_slices = metadata['output_slices'] + self.input_queues, self.npy = make_stateful_input_queues(metadata, self.QUEUE_DEV) + self.policy_input_shapes = stateful_host_shapes(metadata) + elif self.model_type == "supercombo": input_shapes = self.metadata["model"]["input_shapes"] self.output_slices = self.metadata["model"]["output_slices"] self.input_queues, self.npy = make_supercombo_input_queues(input_shapes, self.frame_skip, self.QUEUE_DEV) @@ -698,7 +708,9 @@ class ModelState: return parsed def _reset_state(self) -> None: - if self.model_type == "supercombo": + if self.onnx_history: + self.input_queues, self.npy = make_stateful_input_queues(self.metadata['model'], self.QUEUE_DEV) + elif self.model_type == "supercombo": self.input_queues, self.npy = make_supercombo_input_queues( self.policy_input_shapes, self.frame_skip, self.QUEUE_DEV, )