mirror of
https://github.com/sunnypilot/sunnypilot.git
synced 2026-08-22 09:03:45 +08:00
Driving Model Selector: Cleanup
This commit is contained in:
+16
-22
@@ -63,36 +63,30 @@ class ModelState:
|
||||
self.frame = ModelFrame(context)
|
||||
self.wide_frame = ModelFrame(context)
|
||||
self.prev_desire = np.zeros(ModelConstants.DESIRE_LEN, dtype=np.float32)
|
||||
_inputs_start = {
|
||||
'desire': np.zeros(ModelConstants.DESIRE_LEN * (ModelConstants.HISTORY_BUFFER_LEN+1), dtype=np.float32),
|
||||
'traffic_convention': np.zeros(ModelConstants.TRAFFIC_CONVENTION_LEN, dtype=np.float32),
|
||||
}
|
||||
_inputs_middle = {
|
||||
_inputs = {
|
||||
'lateral_control_params': np.zeros(ModelConstants.LATERAL_CONTROL_PARAMS_LEN, dtype=np.float32), # gen2/3
|
||||
'prev_desired_curv': np.zeros(ModelConstants.PREV_DESIRED_CURV_LEN * (ModelConstants.HISTORY_BUFFER_LEN+1), dtype=np.float32), # gen3
|
||||
}
|
||||
_inputs_end = {
|
||||
if self.custom_model:
|
||||
if self.model_gen == 1:
|
||||
_inputs = {
|
||||
'lat_planner_state': np.zeros(ModelConstants.LAT_PLANNER_STATE_LEN, dtype=np.float32), # gen1
|
||||
}
|
||||
if self.model_gen == 2: # gen2
|
||||
_inputs = {
|
||||
'lateral_control_params': np.zeros(ModelConstants.LATERAL_CONTROL_PARAMS_LEN, dtype=np.float32), # gen2/3
|
||||
'prev_desired_curvs': np.zeros(ModelConstants.PREV_DESIRED_CURVS_LEN, dtype=np.float32), # gen2
|
||||
}
|
||||
|
||||
self.inputs = {
|
||||
'desire': np.zeros(ModelConstants.DESIRE_LEN * (ModelConstants.HISTORY_BUFFER_LEN+1), dtype=np.float32),
|
||||
'traffic_convention': np.zeros(ModelConstants.TRAFFIC_CONVENTION_LEN, dtype=np.float32),
|
||||
**_inputs,
|
||||
'nav_features': np.zeros(ModelConstants.NAV_FEATURE_LEN, dtype=np.float32),
|
||||
'nav_instructions': np.zeros(ModelConstants.NAV_INSTRUCTION_LEN, dtype=np.float32),
|
||||
'features_buffer': np.zeros(ModelConstants.HISTORY_BUFFER_LEN * ModelConstants.FEATURE_LEN, dtype=np.float32),
|
||||
}
|
||||
|
||||
if self.custom_model and self.model_gen == 1:
|
||||
_inputs_middle = {
|
||||
'lat_planner_state': np.zeros(ModelConstants.LAT_PLANNER_STATE_LEN, dtype=np.float32), # gen1
|
||||
}
|
||||
if self.custom_model and self.model_gen == 2: # gen2
|
||||
_inputs_middle = {
|
||||
'lateral_control_params': np.zeros(ModelConstants.LATERAL_CONTROL_PARAMS_LEN, dtype=np.float32), # gen2/3
|
||||
'prev_desired_curvs': np.zeros(ModelConstants.PREV_DESIRED_CURVS_LEN, dtype=np.float32), # gen2
|
||||
}
|
||||
|
||||
self.inputs = {
|
||||
**_inputs_start,
|
||||
**_inputs_middle,
|
||||
**_inputs_end,
|
||||
}
|
||||
|
||||
self.param_s = Params()
|
||||
|
||||
if self.param_s.get_bool("CustomDrivingModel"):
|
||||
|
||||
Reference in New Issue
Block a user