Compare commits

...

1 Commits

Author SHA1 Message Date
firestar5683 5a9deaca5f Cinquev3 2026-09-18 13:43:23 -05:00
2 changed files with 93 additions and 14 deletions
+79 -12
View File
@@ -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
+14 -2
View File
@@ -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,
)