mirror of
https://github.com/firestar5683/StarPilot.git
synced 2026-10-03 21:03:52 +08:00
Compare commits
1 Commits
NikTest
...
Codename-Domi
| Author | SHA1 | Date | |
|---|---|---|---|
| 5a9deaca5f |
@@ -467,7 +467,54 @@ def make_run_supercombo(model_runner, metadata, frame_skip, image_history_pipeli
|
|||||||
return run_policy
|
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
|
seed = 42
|
||||||
|
|
||||||
def random_inputs_run(fn, current_seed, test_values=None, test_buffers=None, expect_match=True):
|
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)
|
np.random.seed(current_seed)
|
||||||
Tensor.manual_seed(current_seed)
|
Tensor.manual_seed(current_seed)
|
||||||
testing = test_values is not None or test_buffers is not None
|
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 index in range(run_count):
|
||||||
for value in npy.values():
|
for value in npy.values():
|
||||||
@@ -489,9 +537,9 @@ def compile_jit(jit, make_random_inputs, input_keys, make_queues):
|
|||||||
end = time.perf_counter()
|
end = time.perf_counter()
|
||||||
print(f" [{index + 1}/{run_count}] enqueue {(mid - start) * 1e3:6.2f} ms -- total {(end - start) * 1e3:6.2f} ms")
|
print(f" [{index + 1}/{run_count}] enqueue {(mid - start) * 1e3:6.2f} ms -- total {(end - start) * 1e3:6.2f} ms")
|
||||||
|
|
||||||
if index == 0:
|
if index < validation_runs:
|
||||||
values = [np.copy(value.numpy()) for value in outputs]
|
values.extend(np.copy(value.numpy()) for value in outputs)
|
||||||
buffers = [np.copy(value.numpy()) for value in input_queues.values()]
|
buffers.extend(np.copy(value.numpy()) for value in input_queues.values())
|
||||||
if not all(np.isfinite(value).all() for value in values):
|
if not all(np.isfinite(value).all() for value in values):
|
||||||
raise ValueError("Compiled JIT produced non-finite outputs")
|
raise ValueError("Compiled JIT produced non-finite outputs")
|
||||||
|
|
||||||
@@ -602,13 +650,31 @@ def main():
|
|||||||
output["metadata"]["model"] = make_metadata_dict(model_path)
|
output["metadata"]["model"] = make_metadata_dict(model_path)
|
||||||
validate_metadata(output["metadata"]["model"])
|
validate_metadata(output["metadata"]["model"])
|
||||||
policy_shapes = output["metadata"]["model"]["input_shapes"]
|
policy_shapes = output["metadata"]["model"]["input_shapes"]
|
||||||
frame_skip = args.frame_skip or derive_frame_skip(policy_shapes)
|
if 'new_img' in policy_shapes:
|
||||||
make_policy_queues = partial(make_supercombo_input_queues, policy_shapes, frame_skip)
|
if args.image_history_pipeline != IMAGE_HISTORY_IN_POLICY:
|
||||||
run_policy = make_run_supercombo(
|
parser.error('ONNX-managed history requires --image-history-pipeline policy')
|
||||||
model_runner, output["metadata"], frame_skip, args.image_history_pipeline,
|
metadata = output['metadata']['model']
|
||||||
)
|
metadata['state_pairs'] = {name: f'next_{name}' for name in policy_shapes
|
||||||
image_shapes = policy_shapes
|
if f'next_{name}' in metadata['output_shapes']}
|
||||||
policy_input_keys = FAST_POLICY_INPUTS if args.image_history_pipeline == IMAGE_HISTORY_IN_POLICY else SUPERCOMBO_POLICY_INPUTS
|
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:
|
else:
|
||||||
if not args.vision_onnx:
|
if not args.vision_onnx:
|
||||||
parser.error("--vision-onnx is required for split models")
|
parser.error("--vision-onnx is required for split models")
|
||||||
@@ -675,6 +741,7 @@ def main():
|
|||||||
)
|
)
|
||||||
output["run_policy"] = compile_jit(
|
output["run_policy"] = compile_jit(
|
||||||
run_policy_jit, make_random_model_inputs, policy_input_keys, make_policy_queues,
|
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
|
model_w, model_h = args.model_size
|
||||||
|
|||||||
@@ -48,6 +48,9 @@ from openpilot.selfdrive.modeld.compile_modeld import (
|
|||||||
derive_frame_skip,
|
derive_frame_skip,
|
||||||
make_split_input_queues,
|
make_split_input_queues,
|
||||||
make_supercombo_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.helpers import get_tg_input_devices, load_oob, tinygrad_dev_config, usbgpu_present
|
||||||
from openpilot.selfdrive.modeld.usbgpu_link import wait_usbgpu_link
|
from openpilot.selfdrive.modeld.usbgpu_link import wait_usbgpu_link
|
||||||
@@ -587,8 +590,15 @@ class ModelState:
|
|||||||
self.run_policy = artifact["run_policy"]
|
self.run_policy = artifact["run_policy"]
|
||||||
self.warp_enqueue = artifact[(cam_w, cam_h)]
|
self.warp_enqueue = artifact[(cam_w, cam_h)]
|
||||||
self.can_prepare_only = self.image_history_pipeline == IMAGE_HISTORY_IN_WARP
|
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"]
|
input_shapes = self.metadata["model"]["input_shapes"]
|
||||||
self.output_slices = self.metadata["model"]["output_slices"]
|
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)
|
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
|
return parsed
|
||||||
|
|
||||||
def _reset_state(self) -> None:
|
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.input_queues, self.npy = make_supercombo_input_queues(
|
||||||
self.policy_input_shapes, self.frame_skip, self.QUEUE_DEV,
|
self.policy_input_shapes, self.frame_skip, self.QUEUE_DEV,
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user