From 038b83ac4a725466b59219a5591876bebefac36b Mon Sep 17 00:00:00 2001 From: firestar5683 <168790843+firestar5683@users.noreply.github.com> Date: Thu, 4 Jun 2026 23:52:55 -0500 Subject: [PATCH] this outta be good --- selfdrive/modeld/modeld_v16.py | 28 ++++++++++-- tinygrad_repo/tinygrad/engine/jit.py | 65 ++++++++++++++++++++++++++-- 2 files changed, 87 insertions(+), 6 deletions(-) diff --git a/selfdrive/modeld/modeld_v16.py b/selfdrive/modeld/modeld_v16.py index 287dc48f7..bfd31764b 100644 --- a/selfdrive/modeld/modeld_v16.py +++ b/selfdrive/modeld/modeld_v16.py @@ -18,6 +18,7 @@ from msgq.visionipc import VisionBuf, VisionIpcClient, VisionStreamType from opendbc.car.car_helpers import get_demo_car_params from setproctitle import setproctitle from tinygrad.dtype import dtypes +from tinygrad.engine.jit import get_out_buffers_for_ei from tinygrad.tensor import Tensor from openpilot.common.file_chunker import read_file_chunked @@ -214,6 +215,23 @@ class ModelState: def slice_outputs(self, model_outputs: np.ndarray, output_slices: dict[str, slice]) -> dict[str, np.ndarray]: return {key: model_outputs[np.newaxis, value] for key, value in output_slices.items()} + def read_captured_outputs(self) -> tuple[np.ndarray, np.ndarray, np.ndarray] | None: + captured = getattr(self.run_policy, "captured", None) + ret_output_map = getattr(captured, "ret_output_map", None) + if captured is None or ret_output_map is None or len(ret_output_map) != 3: + return None + + jit_outs = [] + for ji in captured.jit_cache: + jit_outs.extend(get_out_buffers_for_ei(ji)) + + outputs = [] + for idx in ret_output_map: + if idx is None or idx >= len(jit_outs): + return None + outputs.append(np.frombuffer(bytes(jit_outs[idx].as_memoryview()), dtype=np.float32).copy()) + return tuple(outputs) + 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[self.desire_key][0] = 0 self.npy["desire"][:] = np.where(inputs[self.desire_key] - self.prev_desire > 0.99, inputs[self.desire_key], 0) @@ -245,9 +263,13 @@ class ModelState: big_img=self.vision_inputs["big_img"], ) - vision_output = vision_output.numpy().flatten() - policy_output = policy_output.numpy().flatten() - off_policy_output = off_policy_output.numpy().flatten() + captured_outputs = self.read_captured_outputs() + if captured_outputs is not None: + vision_output, policy_output, off_policy_output = captured_outputs + else: + vision_output = vision_output.numpy().flatten() + policy_output = policy_output.numpy().flatten() + off_policy_output = off_policy_output.numpy().flatten() vision_outputs_dict = self.parser.parse_vision_outputs(self.slice_outputs(vision_output, self.vision_output_slices)) off_policy_outputs_dict = self.parser.parse_off_policy_outputs(self.slice_outputs(off_policy_output, self.off_policy_output_slices)) diff --git a/tinygrad_repo/tinygrad/engine/jit.py b/tinygrad_repo/tinygrad/engine/jit.py index 79fe034d3..2103c87f6 100644 --- a/tinygrad_repo/tinygrad/engine/jit.py +++ b/tinygrad_repo/tinygrad/engine/jit.py @@ -166,6 +166,48 @@ def update_depends(depends:set[Buffer|None], jit_cache:list[ExecItem]): for ei in jit_cache: if any(b in depends for b in ei.bufs): depends.update(get_out_buffers_for_ei(ei)) +def get_ret_tensors(ret:Any) -> list[Tensor]: + if isinstance(ret, Tensor): return [ret] + if isinstance(ret, (tuple, list)): return flatten([get_ret_tensors(x) for x in ret]) + if isinstance(ret, dict): return flatten([get_ret_tensors(x) for x in ret.values()]) + return [] + +def get_jit_outs(jit_cache:list[ExecItem]) -> list[Buffer]: + return flatten([get_out_buffers_for_ei(ei) for ei in jit_cache]) + +def get_ret_output_map(ret:Any, jit_cache:list[ExecItem]) -> list[int|None]: + output_map = {id(buf): idx for idx, buf in enumerate(get_jit_outs(jit_cache))} + ret_output_map: list[int|None] = [] + for t in get_ret_tensors(ret): + realized = t.uop.base.realized + ret_output_map.append(output_map.get(id(realized)) if realized is not None else None) + return ret_output_map + +def get_ret_spec(ret:Any, jit_cache:list[ExecItem]) -> Any: + output_map = {id(buf): idx for idx, buf in enumerate(get_jit_outs(jit_cache))} + if isinstance(ret, Tensor): + realized = ret.uop.base.realized + if realized is not None and (out_idx:=output_map.get(id(realized))) is not None: + return ("tensor", out_idx, ret.uop, ret.requires_grad) + return ("value", ret) + if isinstance(ret, tuple): return ("tuple", tuple(get_ret_spec(x, jit_cache) for x in ret)) + if isinstance(ret, list): return ("list", [get_ret_spec(x, jit_cache) for x in ret]) + if isinstance(ret, dict): return ("dict", [(k, get_ret_spec(v, jit_cache)) for k, v in ret.items()]) + return ("value", ret) + +def rebuild_ret_from_spec(spec:Any, jit_outs:list[Buffer]) -> Any: + tag, payload = spec[0], spec[1:] + if tag == "tensor": + out_idx, template_uop, requires_grad = payload + target_buf = jit_outs[out_idx] + buf_uop = UOp.new_buffer(target_buf.device, target_buf.size, template_uop.base.dtype) + bound = UOp(buf_uop.op, buf_uop.dtype, buf_uop.src, buf_uop.arg, buf_uop.tag, _buffer=target_buf) + return Tensor(template_uop.substitute({template_uop.base: bound}, name="rebuild captured jit ret"), requires_grad=requires_grad) + if tag == "tuple": return tuple(rebuild_ret_from_spec(x, jit_outs) for x in payload[0]) + if tag == "list": return [rebuild_ret_from_spec(x, jit_outs) for x in payload[0]] + if tag == "dict": return {k: rebuild_ret_from_spec(v, jit_outs) for k, v in payload[0]} + return payload[0] + ReturnType = TypeVar('ReturnType') @dataclass class CapturedJit(Generic[ReturnType]): @@ -175,10 +217,14 @@ class CapturedJit(Generic[ReturnType]): extra_view_inputs: list[tuple[int, int, str, int, DType]] expected_names: list[int|str] expected_input_info: list[tuple[UOp, tuple[Variable, ...], DType, str]] # (view, variables, dtype, device) per input + ret_output_map: list[int|None]|None = None + ret_spec: Any = None def __reduce__(self): # TODO: free_intermediates here? replan_buffers_memory_layout here? - return self.__class__, (self.ret, self.jit_cache, self.input_replace, self.extra_view_inputs, self.expected_names, self.expected_input_info) + return self.__class__, ( + self.ret, self.jit_cache, self.input_replace, self.extra_view_inputs, self.expected_names, self.expected_input_info, self.ret_output_map, self.ret_spec, + ) def __post_init__(self): self._jit_cache: list[ExecItem] = self.jit_cache @@ -189,6 +235,18 @@ class CapturedJit(Generic[ReturnType]): self._input_to_max_reader: dict[int, int] = {} for (j, _), idx in self.input_replace.items(): self._input_to_max_reader[idx] = max(self._input_to_max_reader.get(idx, -1), j) self._clear_inputs() + self._rebind_ret_outputs() + + def _rebind_ret_outputs(self): + if self.ret_output_map is None: return + jit_outs = get_jit_outs(self.jit_cache) + for t, out_idx in zip(get_ret_tensors(self.ret), self.ret_output_map): + if out_idx is None or out_idx >= len(jit_outs) or t.uop.base.op is not Ops.BUFFER: continue + target_buf = jit_outs[out_idx] + if t.uop.base.realized is target_buf: continue + buf_uop = UOp.new_buffer(target_buf.device, target_buf.size, t.uop.base.dtype) + new_base = UOp(buf_uop.op, buf_uop.dtype, buf_uop.src, buf_uop.arg, buf_uop.tag, _buffer=target_buf) + t.uop = t.uop.substitute({t.uop.base: new_base}, name="rebind captured jit outputs") def _clear_inputs(self): for (j,i) in self._input_replace.keys(): self._jit_cache[j].bufs[i] = None @@ -244,7 +302,7 @@ class CapturedJit(Generic[ReturnType]): if DEBUG >= 1 and len(self._jit_cache) >= 10: print(f"jit execs {len(self._jit_cache)} kernels") for ei in self._jit_cache: ei.run(var_vals, jit=True) self._clear_inputs() - return self.ret + return rebuild_ret_from_spec(self.ret_spec, get_jit_outs(self.jit_cache)) if self.ret_spec is not None else self.ret def _prepare_jit_inputs(args, kwargs): input_tensors: list[tuple[int|str, Tensor]] = [(name,t) for name,t in list(enumerate(args))+sorted(kwargs.items()) if t.__class__ is Tensor] @@ -361,7 +419,8 @@ class TinyJit(Generic[ReturnType]): if DEBUG >= 1 and len(set(input_replace.values())) != len(input_buffers): print("WARNING: some input tensors not found") # set this for next run - self.captured = CapturedJit(ret, jit_cache, input_replace, extra_view_inputs, names, expected_input_info) + self.captured = CapturedJit(ret, jit_cache, input_replace, extra_view_inputs, names, expected_input_info, + get_ret_output_map(ret, jit_cache), get_ret_spec(ret, jit_cache)) if self.optimize: self.captured.replan_buffers_memory_layout() elif self.cnt >= 2: # jit exec