mirror of
https://github.com/sunnypilot/sunnypilot.git
synced 2026-10-01 06:03:43 +08:00
cherry pick from devtekve as base
This commit is contained in:
@@ -52,11 +52,9 @@ class ModelState:
|
||||
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)
|
||||
|
||||
# img buffers are managed in openCL transform code
|
||||
self.inputs = prepare_inputs()
|
||||
|
||||
model_paths = load_model()
|
||||
model_metadata = load_metadata()
|
||||
self.inputs = prepare_inputs(model_metadata)
|
||||
|
||||
self.output_slices = model_metadata['output_slices']
|
||||
net_output_size = model_metadata['output_shapes']['outputs'][1]
|
||||
@@ -239,7 +237,10 @@ def main(demo=False):
|
||||
if prepare_only:
|
||||
cloudlog.error(f"skipping model eval. Dropped {vipc_dropped_frames} frames")
|
||||
|
||||
inputs = parse_runner_inputs(vec_desire, traffic_convention)
|
||||
inputs: dict[str, np.ndarray] = {
|
||||
'desire': vec_desire,
|
||||
'traffic_convention': traffic_convention,
|
||||
}
|
||||
|
||||
mt1 = time.perf_counter()
|
||||
model_output = model.run(buf_main, buf_extra, model_transform_main, model_transform_extra, inputs, prepare_only)
|
||||
|
||||
@@ -13,30 +13,30 @@ from openpilot.sunnypilot.models.helpers import get_active_bundle
|
||||
from openpilot.system.hardware import PC
|
||||
from openpilot.system.hardware.hw import Paths
|
||||
|
||||
USE_ONNX = int(os.getenv('USE_ONNX', str(int(PC))))
|
||||
USE_ONNX = os.getenv('USE_ONNX', PC)
|
||||
|
||||
CUSTOM_MODEL_PATH = Paths.model_root()
|
||||
METADATA_PATH = Path(__file__).parent / 'models/supercombo_metadata.pkl'
|
||||
METADATA_PATH = Path(__file__).parent / '../models/supercombo_metadata.pkl'
|
||||
|
||||
ModelManager = custom.ModelManagerSP
|
||||
|
||||
|
||||
def load_model():
|
||||
if USE_ONNX:
|
||||
model_paths = {ModelRunner.ONNX: Path(__file__).parent / 'models/supercombo.onnx'}
|
||||
model_paths = {ModelRunner.ONNX: Path(__file__).parent / '../models/supercombo.onnx'}
|
||||
elif bundle := get_active_bundle():
|
||||
drive_model = next(model for model in bundle.models if model.type == ModelManager.Type.drive)
|
||||
model_paths = {ModelRunner.THNEED: f"{CUSTOM_MODEL_PATH}/{drive_model.fileName}"}
|
||||
else:
|
||||
model_paths = {ModelRunner.THNEED: Path(__file__).parent / 'models/supercombo.thneed'}
|
||||
model_paths = {ModelRunner.THNEED: Path(__file__).parent / '../models/supercombo.thneed'}
|
||||
|
||||
return model_paths
|
||||
|
||||
|
||||
def load_metadata():
|
||||
if bundle := get_active_bundle():
|
||||
drive_model = next(model for model in bundle.models if model.type == ModelManager.Type.metadata)
|
||||
metadata_path = f"{CUSTOM_MODEL_PATH}/{drive_model.fileName}"
|
||||
metadata_model = next(model for model in bundle.models if model.type == ModelManager.Type.metadata)
|
||||
metadata_path = f"{CUSTOM_MODEL_PATH}/{metadata_model.fileName}"
|
||||
else:
|
||||
metadata_path = METADATA_PATH
|
||||
|
||||
@@ -46,22 +46,12 @@ def load_metadata():
|
||||
return metadata
|
||||
|
||||
|
||||
def prepare_inputs() -> dict[str, np.ndarray]:
|
||||
model_metadata = load_metadata()
|
||||
|
||||
def prepare_inputs(metadata) -> dict[str, np.ndarray]:
|
||||
# img buffers are managed in openCL transform code
|
||||
inputs: dict[str, np.ndarray] = {
|
||||
key: np.zeros(shape, dtype=np.float32)
|
||||
for key, shape in model_metadata['input_shapes'].items()
|
||||
for key, shape in metadata['input_shapes'].items()
|
||||
if key not in ['input_imgs', 'big_input_imgs']
|
||||
}
|
||||
|
||||
return inputs
|
||||
|
||||
|
||||
def parse_runner_inputs(vec_desire, traffic_convention) -> dict[str, np.ndarray]:
|
||||
inputs: dict[str, np.ndarray] = {
|
||||
'desire': vec_desire,
|
||||
'traffic_convention': traffic_convention,
|
||||
}
|
||||
|
||||
return inputs
|
||||
|
||||
Reference in New Issue
Block a user