mirror of
https://github.com/sunnypilot/sunnypilot.git
synced 2026-09-30 19:13:43 +08:00
Add 20Hz model state, smart input, and model switcher classes
Introduce `ModelState20Hz`, `ModelSmartInput`, and `ModelSwitcher` for enhanced modularity and flexibility in modeld. Refactor `ModelState` to inherit from these new classes, enabling support for 20Hz processing and smart input initialization. Update associated files to handle the new buffer length parameter and metadata management.
This commit is contained in:
+18
-23
@@ -2,6 +2,10 @@
|
||||
import os
|
||||
from openpilot.system.hardware import TICI
|
||||
|
||||
from openpilot.sunnypilot.modeld_v2.model_smart_input import ModelSmartInput
|
||||
from openpilot.sunnypilot.modeld_v2.model_switcher import ModelSwitcher
|
||||
from openpilot.sunnypilot.modeld_v2.model_state_20hz import ModelState20Hz
|
||||
|
||||
#
|
||||
if TICI:
|
||||
from tinygrad.tensor import Tensor
|
||||
@@ -50,34 +54,33 @@ class FrameMeta:
|
||||
if vipc is not None:
|
||||
self.frame_id, self.timestamp_sof, self.timestamp_eof = vipc.frame_id, vipc.timestamp_sof, vipc.timestamp_eof
|
||||
|
||||
class ModelState:
|
||||
class ModelState(ModelState20Hz, ModelSwitcher, ModelSmartInput):
|
||||
frames: dict[str, DrivingModelFrame]
|
||||
inputs: dict[str, np.ndarray]
|
||||
output: np.ndarray
|
||||
prev_desire: np.ndarray # for tracking the rising edge of the pulse
|
||||
|
||||
def __init__(self, context: CLContext):
|
||||
self.is_20hz = False
|
||||
ModelState20Hz.__init__(self, context)
|
||||
ModelSmartInput.__init__(self, METADATA_PATH)
|
||||
|
||||
buffer_length = 5 if self.is_20hz else 2
|
||||
self.frames = {'input_imgs': DrivingModelFrame(context, buffer_length), 'big_input_imgs': DrivingModelFrame(context, buffer_length)}
|
||||
self.prev_desire = np.zeros(ModelConstants.DESIRE_LEN, dtype=np.float32)
|
||||
self.full_features_20Hz = np.zeros((ModelConstants.FULL_HISTORY_BUFFER_LEN, ModelConstants.FEATURE_LEN), dtype=np.float32)
|
||||
self.desire_20Hz = np.zeros((ModelConstants.FULL_HISTORY_BUFFER_LEN + 1, ModelConstants.DESIRE_LEN), dtype=np.float32)
|
||||
self.frames = self.frames or {'input_imgs': DrivingModelFrame(context), 'big_input_imgs': DrivingModelFrame(context)}
|
||||
self.prev_desire = self.prev_desire or np.zeros(ModelConstants.DESIRE_LEN, dtype=np.float32)
|
||||
|
||||
# img buffers are managed in openCL transform code
|
||||
self.numpy_inputs = {
|
||||
'desire': np.zeros((1, (ModelConstants.FULL_HISTORY_BUFFER_LEN+1), ModelConstants.DESIRE_LEN), dtype=np.float32),
|
||||
'traffic_convention': np.zeros((1, ModelConstants.TRAFFIC_CONVENTION_LEN), dtype=np.float32),
|
||||
'lateral_control_params': np.zeros((1, ModelConstants.LATERAL_CONTROL_PARAMS_LEN), dtype=np.float32),
|
||||
'prev_desired_curv': np.zeros((1, (ModelConstants.FULL_HISTORY_BUFFER_LEN+1), ModelConstants.PREV_DESIRED_CURV_LEN), dtype=np.float32),
|
||||
'features_buffer': np.zeros((1, ModelConstants.FULL_HISTORY_BUFFER_LEN, ModelConstants.FEATURE_LEN), dtype=np.float32),
|
||||
}
|
||||
|
||||
with open(METADATA_PATH, 'rb') as f:
|
||||
model_metadata = pickle.load(f)
|
||||
self.input_shapes = model_metadata['input_shapes']
|
||||
self.input_shapes = model_metadata['input_shapes']
|
||||
|
||||
self.output_slices = model_metadata['output_slices']
|
||||
|
||||
# img buffers are managed in openCL transform code
|
||||
self.numpy_inputs = {}
|
||||
for key, shape in self.input_shapes.items():
|
||||
if key not in self.frames: # Managed by opencl
|
||||
self.numpy_inputs[key] = np.zeros(shape, dtype=np.float32)
|
||||
|
||||
net_output_size = model_metadata['output_shapes']['outputs'][1]
|
||||
self.output = np.zeros(net_output_size, dtype=np.float32)
|
||||
self.parser = Parser()
|
||||
@@ -89,14 +92,6 @@ class ModelState:
|
||||
else:
|
||||
self.onnx_cpu_runner = make_onnx_cpu_runner(MODEL_PATH)
|
||||
|
||||
net_output_size = model_metadata['output_shapes']['outputs'][1]
|
||||
self.output = np.zeros(net_output_size, dtype=np.float32)
|
||||
|
||||
num_elements = self.numpy_inputs['features_buffer'].shape[1]
|
||||
step_size = int(-100 / num_elements)
|
||||
self.full_features_20Hz_idxs = np.arange(step_size, step_size * (num_elements + 1), step_size)[::-1]
|
||||
self.desire_reshape_dims = (self.numpy_inputs['desire'].shape[0], self.numpy_inputs['desire'].shape[1], -1, self.numpy_inputs['desire'].shape[2])
|
||||
|
||||
def slice_outputs(self, model_outputs: np.ndarray) -> dict[str, np.ndarray]:
|
||||
parsed_model_outputs = {k: model_outputs[np.newaxis, v] for k,v in self.output_slices.items()}
|
||||
if SEND_RAW_PRED:
|
||||
|
||||
@@ -64,7 +64,7 @@ protected:
|
||||
|
||||
class DrivingModelFrame : public ModelFrame {
|
||||
public:
|
||||
DrivingModelFrame(cl_device_id device_id, cl_context context, uint8_t buffer_length = 2);
|
||||
DrivingModelFrame(cl_device_id device_id, cl_context context, uint8_t buffer_length);
|
||||
~DrivingModelFrame();
|
||||
cl_mem* prepare(cl_mem yuv_cl, int frame_width, int frame_height, int frame_stride, int frame_uv_offset, const mat3& projection);
|
||||
|
||||
|
||||
@@ -20,7 +20,7 @@ cdef extern from "selfdrive/modeld/models/commonmodel.h":
|
||||
|
||||
cppclass DrivingModelFrame:
|
||||
int buf_size
|
||||
DrivingModelFrame(cl_device_id, cl_context)
|
||||
DrivingModelFrame(cl_device_id, cl_context, unsigned char)
|
||||
|
||||
cppclass MonitoringModelFrame:
|
||||
int buf_size
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
import numpy as np
|
||||
cimport numpy as cnp
|
||||
from libc.string cimport memcpy
|
||||
from libc.stdint cimport uintptr_t
|
||||
from libc.stdint cimport uintptr_t, uint8_t
|
||||
|
||||
from msgq.visionipc.visionipc cimport cl_mem
|
||||
from msgq.visionipc.visionipc_pyx cimport VisionBuf, CLContext as BaseCLContext
|
||||
@@ -60,7 +60,7 @@ cdef class DrivingModelFrame(ModelFrame):
|
||||
cdef cppDrivingModelFrame * _frame
|
||||
|
||||
def __cinit__(self, CLContext context, int buffer_length=2):
|
||||
self._frame = new cppDrivingModelFrame(context.device_id, context.context)
|
||||
self._frame = new cppDrivingModelFrame(context.device_id, context.context, buffer_length)
|
||||
self.frame = <cppModelFrame*>(self._frame)
|
||||
self.buf_size = self._frame.buf_size
|
||||
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
import pickle
|
||||
from abc import abstractmethod, ABC
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
class ModelSmartInput(ABC):
|
||||
def __init__(self, METADATA_PATH):
|
||||
self._using_smart_input = True
|
||||
self.desire_reshape_dims = None
|
||||
self.output = None
|
||||
self.full_features_20Hz_idxs = None
|
||||
self._output_slices = None
|
||||
self._input_shapes = None
|
||||
self._numpy_inputs = {}
|
||||
|
||||
if self._using_smart_input:
|
||||
self.initialize_smart_input(METADATA_PATH)
|
||||
|
||||
def initialize_smart_input(self, METADATA_PATH):
|
||||
with open(METADATA_PATH, 'rb') as f:
|
||||
model_metadata = pickle.load(f)
|
||||
|
||||
self._input_shapes = model_metadata['input_shapes']
|
||||
self._output_slices = model_metadata['output_slices']
|
||||
|
||||
for key, shape in self.input_shapes.items():
|
||||
if key not in ['input_imgs', 'big_input_imgs']: # Managed by opencl
|
||||
self._numpy_inputs[key] = np.zeros(shape, dtype=np.float32)
|
||||
|
||||
net_output_size = model_metadata['output_shapes']['outputs'][1]
|
||||
self.output = np.zeros(net_output_size, dtype=np.float32)
|
||||
|
||||
num_elements = self.numpy_inputs['features_buffer'].shape[1]
|
||||
step_size = int(-100 / num_elements)
|
||||
self.full_features_20Hz_idxs = np.arange(step_size, step_size * (num_elements + 1), step_size)[::-1]
|
||||
self.desire_reshape_dims = (self.numpy_inputs['desire'].shape[0], self.numpy_inputs['desire'].shape[1], -1, self.numpy_inputs['desire'].shape[2])
|
||||
|
||||
@property
|
||||
def input_shapes(self):
|
||||
return self._input_shapes
|
||||
|
||||
@input_shapes.setter
|
||||
def input_shapes(self, value):
|
||||
if not self._input_shapes:
|
||||
self._input_shapes = value
|
||||
print("Waring: ignoring input_shapes setter because ModelSmartInput is in use.")
|
||||
|
||||
@property
|
||||
def output_slices(self):
|
||||
return self._output_slices
|
||||
|
||||
@output_slices.setter
|
||||
def output_slices(self, value):
|
||||
if not self._output_slices:
|
||||
self._output_slices = value
|
||||
print("Waring: ignoring output_slices setter because ModelSmartInput is in use.")
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def frames(self):
|
||||
raise NotImplementedError
|
||||
|
||||
@property
|
||||
def numpy_inputs(self):
|
||||
return self._numpy_inputs
|
||||
|
||||
@numpy_inputs.setter
|
||||
def numpy_inputs(self, value):
|
||||
if not self._numpy_inputs:
|
||||
self._numpy_inputs = value
|
||||
print("Waring: ignoring numpy_inputs setter because ModelSmartInput is in use.")
|
||||
Executable
+43
@@ -0,0 +1,43 @@
|
||||
from abc import abstractmethod, ABC
|
||||
|
||||
import numpy as np
|
||||
from openpilot.selfdrive.modeld.constants import ModelConstants
|
||||
from openpilot.selfdrive.modeld.models.commonmodel_pyx import DrivingModelFrame
|
||||
|
||||
|
||||
class ModelState20Hz(ABC):
|
||||
def __init__(self, context):
|
||||
self.is_20hz = False
|
||||
self._context = context
|
||||
self.desire_20Hz = None
|
||||
self.full_features_20Hz = None
|
||||
self.frames = None
|
||||
self.prev_desire = None
|
||||
if self.is_20hz:
|
||||
self.initialize_20hz_buffers()
|
||||
|
||||
def initialize_20hz_buffers(self):
|
||||
self.full_features_20Hz = np.zeros((ModelConstants.FULL_HISTORY_BUFFER_LEN, ModelConstants.FEATURE_LEN), dtype=np.float32)
|
||||
self.desire_20Hz = np.zeros((ModelConstants.FULL_HISTORY_BUFFER_LEN + 1, ModelConstants.DESIRE_LEN), dtype=np.float32)
|
||||
self.frames = {'input_imgs': DrivingModelFrame(self._context, self.buffer_length), 'big_input_imgs': DrivingModelFrame(self._context, self.buffer_length)}
|
||||
self.prev_desire = np.zeros(ModelConstants.DESIRE_LEN, dtype=np.float32)
|
||||
|
||||
@property
|
||||
def frames(self):
|
||||
return self._frames
|
||||
|
||||
@frames.setter
|
||||
def frames(self, value):
|
||||
self._frames = value
|
||||
|
||||
@property
|
||||
def prev_desire(self):
|
||||
return self._prev_desire
|
||||
|
||||
@prev_desire.setter
|
||||
def prev_desire(self, value):
|
||||
self._prev_desire = value
|
||||
|
||||
@property
|
||||
def buffer_length(self):
|
||||
return 5 if self.is_20hz else 2
|
||||
@@ -0,0 +1,2 @@
|
||||
class ModelSwitcher:
|
||||
pass
|
||||
Reference in New Issue
Block a user