This commit is contained in:
firestar5683
2026-05-28 17:38:14 -05:00
parent 742c104964
commit a9930207ba
2 changed files with 53 additions and 20 deletions
+47 -16
View File
@@ -51,6 +51,25 @@ def make_random_images(keys, shape, device=None):
return {key: Tensor.randint(shape, low=0, high=256, dtype="uint8", device=device).realize() for key in keys}
class _BlobTensorInputs(dict):
_backing_arrays: dict[str, np.ndarray]
def make_random_blob_images(keys, shape, device=None):
blob_shape = shape if isinstance(shape, tuple) else (shape,)
backing_arrays = {
key: np.random.randint(0, 256, size=blob_shape, dtype=np.uint8)
for key in keys
}
inputs = _BlobTensorInputs({
key: Tensor.from_blob(array.ctypes.data, array.shape, dtype="uint8", device=device).realize()
for key, array in backing_arrays.items()
})
# Keep the numpy storage alive for the duration of the JIT capture/replay call.
inputs._backing_arrays = backing_arrays
return inputs
def warp_perspective_tinygrad(src_flat, matrix_inverse, dst_shape, src_shape, stride_pad, border_fill_val=None):
width_dst, height_dst = dst_shape
height_src, width_src = src_shape
@@ -133,13 +152,30 @@ def make_frame_prepare(nv12: NV12Frame, model_w, model_h):
return frame_prepare_tinygrad
def make_input_queues(vision_input_shapes, policy_input_shapes, frame_skip, device):
def make_tensor_inputs(vision_input_shapes, policy_input_shapes, frame_skip, device):
img = vision_input_shapes["img"]
n_frames = img[1] // 6
img_buf_shape = (frame_skip * (n_frames - 1) + 1, 6, img[2], img[3])
features_buffer = policy_input_shapes["features_buffer"]
desire_pulse = policy_input_shapes["desire_pulse"]
return {
"img_q": Tensor(np.zeros(img_buf_shape, dtype=np.uint8), device=device).contiguous().realize(),
"big_img_q": Tensor(np.zeros(img_buf_shape, dtype=np.uint8), device=device).contiguous().realize(),
"feat_q": Tensor(
np.zeros((frame_skip * (features_buffer[1] - 1) + 1, features_buffer[0], features_buffer[2]), dtype=np.float32),
device=device,
).contiguous().realize(),
"desire_q": Tensor(
np.zeros((frame_skip * desire_pulse[1], desire_pulse[0], desire_pulse[2]), dtype=np.float32),
device=device,
).contiguous().realize(),
}
def make_npy_inputs(policy_input_shapes):
desire_pulse = policy_input_shapes["desire_pulse"]
traffic_convention = policy_input_shapes["traffic_convention"]
npy = {
@@ -150,21 +186,14 @@ def make_input_queues(vision_input_shapes, policy_input_shapes, frame_skip, devi
}
if "action_t" in policy_input_shapes:
npy["action_t"] = np.zeros(policy_input_shapes["action_t"], dtype=np.float32)
npy_tensors = {key: Tensor(value, device="NPY").realize() for key, value in npy.items()}
return npy, npy_tensors
input_queues = {
"img_q": Tensor(np.zeros(img_buf_shape, dtype=np.uint8), device=device).contiguous().realize(),
"big_img_q": Tensor(np.zeros(img_buf_shape, dtype=np.uint8), device=device).contiguous().realize(),
"feat_q": Tensor(
np.zeros((frame_skip * (features_buffer[1] - 1) + 1, features_buffer[0], features_buffer[2]), dtype=np.float32),
device=device,
).contiguous().realize(),
"desire_q": Tensor(
np.zeros((frame_skip * desire_pulse[1], desire_pulse[0], desire_pulse[2]), dtype=np.float32),
device=device,
).contiguous().realize(),
**{key: Tensor(value, device="NPY").realize() for key, value in npy.items()},
}
return input_queues, npy
def make_input_queues(vision_input_shapes, policy_input_shapes, frame_skip, device):
tensor_inputs = make_tensor_inputs(vision_input_shapes, policy_input_shapes, frame_skip, device)
npy, npy_tensors = make_npy_inputs(policy_input_shapes)
return {**tensor_inputs, **npy_tensors}, npy
def shift_and_sample(buf, new_val, sample_fn):
@@ -343,13 +372,15 @@ if __name__ == "__main__":
out["metadata"]["vision"] = vision_metadata
out["metadata"]["off_policy"] = off_policy_metadata
out["metadata"]["on_policy"] = on_policy_metadata
out["tensor_inputs"] = make_tensor_inputs(vision_metadata["input_shapes"], on_policy_metadata["input_shapes"], args.frame_skip, Device.DEFAULT)
make_random_model_inputs = partial(make_random_images, keys=["img", "big_img"], shape=vision_metadata["input_shapes"]["img"])
out["run_policy"] = compile_jit(run_policy_jit, make_random_model_inputs, POLICY_INPUTS, args.frame_skip, vision_metadata, on_policy_metadata)
for cam_w, cam_h in args.camera_resolutions:
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)
# Capture warp against blob-backed frames so the JIT ABI matches runtime VisionBuf inputs.
make_random_warp_inputs = partial(make_random_blob_images, keys=["frame", "big_frame"], shape=nv12.size, device=WARP_DEV)
warp_enqueue = TinyJit(make_warp(nv12, model_w, model_h, args.frame_skip), prune=True)
out[(cam_w, cam_h)] = compile_jit(warp_enqueue, make_random_warp_inputs, WARP_INPUTS, args.frame_skip, vision_metadata, on_policy_metadata)
+6 -4
View File
@@ -28,7 +28,7 @@ from openpilot.common.transformations.camera import DEVICE_CAMERAS
from openpilot.common.transformations.model import get_warp_matrix
from openpilot.selfdrive.controls.lib.desire_helper import DesireHelper
from openpilot.selfdrive.controls.lib.drive_helpers import smooth_value
from openpilot.selfdrive.modeld.compile_modeld import POLICY_INPUTS, WARP_INPUTS, make_input_queues
from openpilot.selfdrive.modeld.compile_modeld import POLICY_INPUTS, WARP_INPUTS, make_npy_inputs, make_tensor_inputs
from openpilot.selfdrive.modeld.constants import ModelConstants
from openpilot.selfdrive.modeld.fill_model_msg import PublishState, fill_model_msg, fill_pose_msg
from openpilot.selfdrive.modeld.helpers import get_tg_input_devices
@@ -184,9 +184,11 @@ class ModelState:
self.frame_skip = ModelConstants.MODEL_RUN_FREQ // ModelConstants.MODEL_CONTEXT_FREQ
input_devices = get_tg_input_devices(PROCESS_NAME, usbgpu)
self.WARP_DEV, self.QUEUE_DEV = input_devices["WARP_DEV"], input_devices["QUEUE_DEV"]
self.input_queues, self.npy = make_input_queues(
self.vision_input_shapes, self.policy_input_shapes, self.frame_skip, device=self.QUEUE_DEV
)
tensor_inputs = jits.get("tensor_inputs")
if tensor_inputs is None:
tensor_inputs = make_tensor_inputs(self.vision_input_shapes, self.policy_input_shapes, self.frame_skip, device=self.QUEUE_DEV)
self.npy, npy_tensors = make_npy_inputs(self.policy_input_shapes)
self.input_queues = {**tensor_inputs, **npy_tensors}
self.full_frames: dict[str, Tensor] = {}
self._blob_cache: dict[tuple[str, int], Tensor] = {}
self.parser = Parser()