mirror of
https://github.com/infiniteCable2/openpilot.git
synced 2026-09-06 08:13:41 +08:00
Merge branch 'master-new' of https://github.com/sunnypilot/sunnypilot into master-new
This commit is contained in:
+21
-14
@@ -39,20 +39,6 @@ struct ModelManagerSP @0xaedffd8f31e7b55d {
|
||||
sha256 @1 :Text;
|
||||
}
|
||||
|
||||
enum Type {
|
||||
drive @0;
|
||||
navigation @1;
|
||||
metadata @2;
|
||||
}
|
||||
|
||||
struct Model {
|
||||
fullName @0 :Text;
|
||||
fileName @1 :Text;
|
||||
downloadUri @2 :DownloadUri;
|
||||
downloadProgress @3 :DownloadProgress;
|
||||
type @4 :Type;
|
||||
}
|
||||
|
||||
enum DownloadStatus {
|
||||
notDownloading @0;
|
||||
downloading @1;
|
||||
@@ -67,6 +53,25 @@ struct ModelManagerSP @0xaedffd8f31e7b55d {
|
||||
eta @2 :UInt32;
|
||||
}
|
||||
|
||||
struct Artifact {
|
||||
fileName @0 :Text;
|
||||
downloadUri @1 :DownloadUri;
|
||||
downloadProgress @2 :DownloadProgress;
|
||||
}
|
||||
|
||||
struct Model {
|
||||
type @0 :Type;
|
||||
artifact @1 :Artifact; # Main artifact
|
||||
metadata @2 :Artifact; # Metadata artifact
|
||||
|
||||
enum Type {
|
||||
supercombo @0;
|
||||
navigation @1;
|
||||
vision @2;
|
||||
policy @3;
|
||||
}
|
||||
}
|
||||
|
||||
enum Runner {
|
||||
snpe @0;
|
||||
tinygrad @1;
|
||||
@@ -83,6 +88,8 @@ struct ModelManagerSP @0xaedffd8f31e7b55d {
|
||||
environment @6 :Text;
|
||||
runner @7 :Runner;
|
||||
is20hz @8 :Bool;
|
||||
ref @9 :Text; # New field
|
||||
minimumSelectorVersion @10 :UInt32;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -47,24 +47,24 @@ void SoftwarePanelSP::handleBundleDownloadProgress() {
|
||||
// Get status for each model type in order
|
||||
for (const auto &model: models) {
|
||||
QString typeName;
|
||||
QString modelName;
|
||||
QString modelName = QString::fromStdString(bundle.getDisplayName());
|
||||
|
||||
switch (model.getType()) {
|
||||
case cereal::ModelManagerSP::Type::DRIVE:
|
||||
case cereal::ModelManagerSP::Model::Type::SUPERCOMBO:
|
||||
typeName = tr("Driving");
|
||||
modelName = QString::fromStdString(bundle.getDisplayName());
|
||||
break;
|
||||
case cereal::ModelManagerSP::Type::NAVIGATION:
|
||||
case cereal::ModelManagerSP::Model::Type::NAVIGATION:
|
||||
typeName = tr("Navigation");
|
||||
modelName = QString::fromStdString(model.getFullName());
|
||||
break;
|
||||
case cereal::ModelManagerSP::Type::METADATA:
|
||||
typeName = tr("Metadata");
|
||||
modelName = QString::fromStdString(model.getFullName());
|
||||
case cereal::ModelManagerSP::Model::Type::VISION:
|
||||
typeName = tr("Vision");
|
||||
break;
|
||||
case cereal::ModelManagerSP::Model::Type::POLICY:
|
||||
typeName = tr("Policy");
|
||||
break;
|
||||
}
|
||||
|
||||
const auto &progress = model.getDownloadProgress();
|
||||
const auto &progress = model.getArtifact().getDownloadProgress();
|
||||
QString line;
|
||||
|
||||
if (progress.getStatus() == cereal::ModelManagerSP::DownloadStatus::DOWNLOADING) {
|
||||
|
||||
@@ -25,13 +25,19 @@ METADATA_PATH = Path(__file__).parent / '../models/supercombo_metadata.pkl'
|
||||
ModelManager = custom.ModelManagerSP
|
||||
|
||||
|
||||
def _get_model():
|
||||
if bundle := get_active_bundle():
|
||||
drive_model = next(model for model in bundle.models if model.type == ModelManager.Model.Type.supercombo)
|
||||
return drive_model
|
||||
|
||||
return None
|
||||
|
||||
def get_model_path():
|
||||
if USE_ONNX:
|
||||
return {ModelRunner.ONNX: Path(__file__).parent / '../models/supercombo.onnx'}
|
||||
|
||||
if bundle := get_active_bundle():
|
||||
drive_model = next(model for model in bundle.models if model.type == ModelManager.Type.drive)
|
||||
return {ModelRunner.THNEED: f"{CUSTOM_MODEL_PATH}/{drive_model.fileName}"}
|
||||
if model := _get_model():
|
||||
return {ModelRunner.THNEED: f"{CUSTOM_MODEL_PATH}/{model.fileName}"}
|
||||
|
||||
return {ModelRunner.THNEED: Path(__file__).parent / '../models/supercombo.thneed'}
|
||||
|
||||
@@ -39,9 +45,8 @@ def get_model_path():
|
||||
def load_metadata():
|
||||
metadata_path = METADATA_PATH
|
||||
|
||||
if bundle := get_active_bundle():
|
||||
metadata_model = next(model for model in bundle.models if model.type == ModelManager.Type.metadata)
|
||||
metadata_path = f"{CUSTOM_MODEL_PATH}/{metadata_model.fileName}"
|
||||
if model := _get_model():
|
||||
metadata_path = f"{CUSTOM_MODEL_PATH}/{model.metadata.fileName}"
|
||||
|
||||
with open(metadata_path, 'rb') as f:
|
||||
return pickle.load(f)
|
||||
|
||||
@@ -35,8 +35,8 @@ class ModelRunner(ABC):
|
||||
|
||||
if bundle := get_active_bundle():
|
||||
bundle_models = {model.type.raw: model for model in bundle.models}
|
||||
self._drive_model = bundle_models.get(ModelManager.Type.drive)
|
||||
self._metadata_model = bundle_models.get(ModelManager.Type.metadata)
|
||||
self._drive_model = bundle_models.get(ModelManager.Model.Type.supercombo)
|
||||
self._metadata_model = self._drive_model.metadata
|
||||
self.is_20hz = bundle.is20hz
|
||||
|
||||
# Override the metadata path if a metadata model is found in the active bundle
|
||||
@@ -82,7 +82,7 @@ class TinygradRunner(ModelRunner):
|
||||
|
||||
model_pkl_path = MODEL_PKL_PATH
|
||||
if self._drive_model:
|
||||
model_pkl_path = f"{CUSTOM_MODEL_PATH}/{self._drive_model.fileName}"
|
||||
model_pkl_path = f"{CUSTOM_MODEL_PATH}/{self._drive_model.artifact.fileName}"
|
||||
assert model_pkl_path.endswith('_tinygrad.pkl'), f"Invalid model file: {model_pkl_path} for TinygradRunner"
|
||||
|
||||
# Load Tinygrad model
|
||||
|
||||
@@ -11,6 +11,7 @@ import time
|
||||
import requests
|
||||
from openpilot.common.params import Params
|
||||
from openpilot.common.swaglog import cloudlog
|
||||
from sunnypilot.models.helpers import is_bundle_version_compatible
|
||||
|
||||
from cereal import custom
|
||||
|
||||
@@ -19,68 +20,49 @@ class ModelParser:
|
||||
"""Handles parsing of model data into cereal objects"""
|
||||
|
||||
@staticmethod
|
||||
def _parse_model(full_name: str, file_name: str, uri_data: dict,
|
||||
model_type: custom.ModelManagerSP.Type) -> custom.ModelManagerSP.Model:
|
||||
model = custom.ModelManagerSP.Model()
|
||||
def _parse_download_uri(download_uri_data) -> custom.ModelManagerSP.DownloadUri:
|
||||
download_uri = custom.ModelManagerSP.DownloadUri()
|
||||
download_uri.uri = download_uri_data.get("url")
|
||||
download_uri.sha256 = download_uri_data.get("sha256")
|
||||
return download_uri
|
||||
|
||||
download_uri.uri = uri_data["url"]
|
||||
download_uri.sha256 = uri_data["sha256"]
|
||||
@staticmethod
|
||||
def _parse_artifact(artifact_data) -> custom.ModelManagerSP.Artifact:
|
||||
artifact = custom.ModelManagerSP.Artifact()
|
||||
artifact.fileName = artifact_data.get("file_name")
|
||||
artifact.downloadUri = ModelParser._parse_download_uri(artifact_data.get("download_uri", {}))
|
||||
return artifact
|
||||
|
||||
model.fullName = full_name
|
||||
model.fileName = file_name
|
||||
model.downloadUri = download_uri
|
||||
model.type = model_type
|
||||
@staticmethod
|
||||
def _parse_model(model_data) -> custom.ModelManagerSP.Model:
|
||||
model = custom.ModelManagerSP.Model()
|
||||
|
||||
model.type = model_data.get("type")
|
||||
model.artifact = ModelParser._parse_artifact(model_data.get("artifact", {}))
|
||||
if metadata := model_data.get("metadata"):
|
||||
model.metadata = ModelParser._parse_artifact(metadata)
|
||||
return model
|
||||
|
||||
@staticmethod
|
||||
def _parse_bundle(key: str, value: dict) -> custom.ModelManagerSP.ModelBundle:
|
||||
def _parse_bundle(bundle) -> custom.ModelManagerSP.ModelBundle:
|
||||
model_bundle = custom.ModelManagerSP.ModelBundle()
|
||||
|
||||
# Parse main driving model
|
||||
models = [
|
||||
ModelParser._parse_model(
|
||||
value["full_name"],
|
||||
value["file_name"],
|
||||
value["download_uri"],
|
||||
custom.ModelManagerSP.Type.drive
|
||||
)
|
||||
]
|
||||
|
||||
# Parse navigation model if exists
|
||||
if value.get("download_uri_nav"):
|
||||
models.append(ModelParser._parse_model(
|
||||
value["full_name_nav"],
|
||||
value["file_name_nav"],
|
||||
value["download_uri_nav"],
|
||||
custom.ModelManagerSP.Type.navigation
|
||||
))
|
||||
|
||||
# Parse metadata model if exists
|
||||
if value.get("download_uri_metadata"):
|
||||
models.append(ModelParser._parse_model(
|
||||
value["full_name_metadata"],
|
||||
value["file_name_metadata"],
|
||||
value["download_uri_metadata"],
|
||||
custom.ModelManagerSP.Type.metadata
|
||||
))
|
||||
|
||||
model_bundle.index = int(value["index"])
|
||||
model_bundle.internalName = key
|
||||
model_bundle.displayName = value["display_name"]
|
||||
model_bundle.models = models
|
||||
model_bundle.index = int(bundle["index"])
|
||||
model_bundle.internalName = bundle["short_name"]
|
||||
model_bundle.displayName = bundle["display_name"]
|
||||
model_bundle.models = [ModelParser._parse_model(model) for model in bundle.get("models",[])]
|
||||
model_bundle.status = 0
|
||||
model_bundle.generation = int(value["generation"])
|
||||
model_bundle.environment = value["environment"]
|
||||
model_bundle.runner = value.get("runner", custom.ModelManagerSP.Runner.snpe)
|
||||
model_bundle.is20hz = value.get("is_20hz", False)
|
||||
model_bundle.generation = int(bundle["generation"])
|
||||
model_bundle.environment = bundle["environment"]
|
||||
model_bundle.runner = bundle.get("runner", custom.ModelManagerSP.Runner.snpe)
|
||||
model_bundle.is20hz = bundle.get("is_20hz", False)
|
||||
model_bundle.minimumSelectorVersion = int(bundle["minimum_selector_version"])
|
||||
|
||||
return model_bundle
|
||||
|
||||
@staticmethod
|
||||
def parse_models(json_data: dict) -> list[custom.ModelManagerSP.ModelBundle]:
|
||||
return [ModelParser._parse_bundle(key, value) for key, value in json_data.items()]
|
||||
found_bundles = [ModelParser._parse_bundle(bundle) for bundle in json_data.get("bundles", [])]
|
||||
return [bundle for bundle in found_bundles if is_bundle_version_compatible(bundle.to_dict())]
|
||||
|
||||
|
||||
class ModelCache:
|
||||
@@ -122,7 +104,7 @@ class ModelCache:
|
||||
|
||||
class ModelFetcher:
|
||||
"""Handles fetching and caching of model data from remote source"""
|
||||
MODEL_URL = "https://docs.sunnypilot.ai/driving_models_v2.json"
|
||||
MODEL_URL = "https://docs.sunnypilot.ai/driving_models_v3.json"
|
||||
|
||||
def __init__(self, params: Params):
|
||||
self.params = params
|
||||
@@ -143,7 +125,7 @@ class ModelFetcher:
|
||||
cloudlog.exception("Error fetching models")
|
||||
raise
|
||||
|
||||
def get_available_models(self) -> list[custom.ModelManagerSP.ModelBundle]:
|
||||
def get_available_bundles(self) -> list[custom.ModelManagerSP.ModelBundle]:
|
||||
"""Gets the list of available models, with smart cache handling"""
|
||||
cached_data, is_expired = self.model_cache.get()
|
||||
|
||||
@@ -160,3 +142,16 @@ class ModelFetcher:
|
||||
|
||||
cloudlog.warning("Failed to fetch fresh data. Using expired cache as fallback")
|
||||
return self.model_parser.parse_models(cached_data)
|
||||
|
||||
if __name__ == "__main__":
|
||||
params = Params()
|
||||
model_fetcher = ModelFetcher(params)
|
||||
bundles = model_fetcher.get_available_bundles()
|
||||
for bundle in bundles:
|
||||
for model in bundle.models:
|
||||
# Print model details
|
||||
print(f"Bundle: {bundle.internalName}, Type: {model.type}, Status: {bundle.status}")
|
||||
# Print artifact details
|
||||
print(f"Artifact: {model.artifact.fileName}, Download URI: {model.artifact.downloadUri.uri}")
|
||||
# Print metadata details
|
||||
print(f"Metadata: {model.metadata.fileName}, Download URI: {model.metadata.downloadUri.uri}")
|
||||
|
||||
@@ -8,7 +8,11 @@ See the LICENSE.md file in the root directory for more details.
|
||||
import hashlib
|
||||
import os
|
||||
from openpilot.common.params import Params
|
||||
from cereal import custom, messaging
|
||||
from cereal import custom
|
||||
import json
|
||||
|
||||
CURRENT_SELECTOR_VERSION = 2
|
||||
REQUIRED_MIN_SELECTOR_VERSION = 2
|
||||
|
||||
|
||||
async def verify_file(file_path: str, expected_hash: str) -> bool:
|
||||
@@ -24,13 +28,36 @@ async def verify_file(file_path: str, expected_hash: str) -> bool:
|
||||
return sha256_hash.hexdigest().lower() == expected_hash.lower()
|
||||
|
||||
|
||||
def is_bundle_version_compatible(bundle: dict) -> bool:
|
||||
"""
|
||||
Checks whether the model bundle is compatible with the current selector version constraints.
|
||||
|
||||
The bundle specifies a `minimum_selector_version`, which defines the minimum selector version
|
||||
required to load the model. This function ensures that:
|
||||
|
||||
1. The model is not too old: the bundle must require at least `REQUIRED_MIN_SELECTOR_VERSION`.
|
||||
2. The model is not too new: it must support the current selector version (`CURRENT_SELECTOR_VERSION`).
|
||||
|
||||
This allows the selector to enforce both a minimum and maximum range of supported models,
|
||||
even if a model would otherwise be compatible.
|
||||
|
||||
:param bundle: Dictionary containing `minimum_selector_version`, as defined by the model bundle.
|
||||
:type bundle: Dict
|
||||
:return: True if the selector version is within the accepted range for the bundle; otherwise False.
|
||||
:rtype: Bool
|
||||
"""
|
||||
return bool(REQUIRED_MIN_SELECTOR_VERSION <= bundle.get("minimumSelectorVersion", 0) <= CURRENT_SELECTOR_VERSION)
|
||||
|
||||
def get_active_bundle(params: Params = None) -> custom.ModelManagerSP.ModelBundle:
|
||||
"""Gets the active model bundle from cache"""
|
||||
if params is None:
|
||||
params = Params()
|
||||
|
||||
if active_bundle := params.get("ModelManager_ActiveBundle"):
|
||||
return messaging.log_from_bytes(active_bundle, custom.ModelManagerSP.ModelBundle)
|
||||
try:
|
||||
if (active_bundle := json.loads(params.get("ModelManager_ActiveBundle") or "{}")) and is_bundle_version_compatible(active_bundle):
|
||||
return custom.ModelManagerSP.ModelBundle(**active_bundle)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return None
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ See the LICENSE.md file in the root directory for more details.
|
||||
import asyncio
|
||||
import os
|
||||
import time
|
||||
import json
|
||||
|
||||
import aiohttp
|
||||
from openpilot.common.params import Params
|
||||
@@ -73,43 +74,54 @@ class ModelManagerSP:
|
||||
# Clean up start time after download completes
|
||||
del self._download_start_times[model.fileName]
|
||||
|
||||
async def _process_model(self, model, destination_path: str) -> None:
|
||||
async def _process_artifact(self, artifact, destination_path: str) -> None:
|
||||
"""Processes a single model download including verification"""
|
||||
url = model.downloadUri.uri
|
||||
expected_hash = model.downloadUri.sha256
|
||||
filename = model.fileName
|
||||
if not artifact.downloadUri.uri:
|
||||
return None
|
||||
|
||||
url = artifact.downloadUri.uri
|
||||
expected_hash = artifact.downloadUri.sha256
|
||||
filename = artifact.fileName
|
||||
full_path = os.path.join(destination_path, filename)
|
||||
|
||||
try:
|
||||
# Check existing file
|
||||
if os.path.exists(full_path) and await verify_file(full_path, expected_hash):
|
||||
model.downloadProgress.status = custom.ModelManagerSP.DownloadStatus.cached
|
||||
model.downloadProgress.progress = 100
|
||||
model.downloadProgress.eta = 0
|
||||
artifact.downloadProgress.status = custom.ModelManagerSP.DownloadStatus.cached
|
||||
artifact.downloadProgress.progress = 100
|
||||
artifact.downloadProgress.eta = 0
|
||||
self._report_status()
|
||||
return
|
||||
|
||||
# Download and verify
|
||||
await self._download_file(url, full_path, model)
|
||||
await self._download_file(url, full_path, artifact)
|
||||
if not await verify_file(full_path, expected_hash):
|
||||
raise ValueError(f"Hash validation failed for {filename}")
|
||||
|
||||
model.downloadProgress.status = custom.ModelManagerSP.DownloadStatus.downloaded
|
||||
model.downloadProgress.eta = 0
|
||||
artifact.downloadProgress.status = custom.ModelManagerSP.DownloadStatus.downloaded
|
||||
artifact.downloadProgress.eta = 0
|
||||
self._report_status()
|
||||
|
||||
except Exception as e:
|
||||
cloudlog.error(f"Error downloading {filename}: {str(e)}")
|
||||
if os.path.exists(full_path):
|
||||
os.remove(full_path)
|
||||
model.downloadProgress.status = custom.ModelManagerSP.DownloadStatus.failed
|
||||
model.downloadProgress.eta = 0
|
||||
artifact.downloadProgress.status = custom.ModelManagerSP.DownloadStatus.failed
|
||||
artifact.downloadProgress.eta = 0
|
||||
self.selected_bundle.status = custom.ModelManagerSP.DownloadStatus.failed
|
||||
self._report_status()
|
||||
# Clean up start time if it exists
|
||||
self._download_start_times.pop(model.fileName, None)
|
||||
self._download_start_times.pop(artifact.fileName, None)
|
||||
raise
|
||||
|
||||
async def _process_model(self, model, destination_path: str) -> None:
|
||||
"""Processes a single model download including verification"""
|
||||
model_artifact = model.artifact
|
||||
metadata_artifact = model.metadata
|
||||
|
||||
await self._process_artifact(metadata_artifact, destination_path)
|
||||
await self._process_artifact(model_artifact, destination_path)
|
||||
|
||||
def _report_status(self) -> None:
|
||||
"""Reports current status through messaging system"""
|
||||
msg = messaging.new_message('modelManagerSP', valid=True)
|
||||
@@ -134,7 +146,7 @@ class ModelManagerSP:
|
||||
await asyncio.gather(*tasks)
|
||||
self.active_bundle = self.selected_bundle
|
||||
self.active_bundle.status = custom.ModelManagerSP.DownloadStatus.downloaded
|
||||
self.params.put("ModelManager_ActiveBundle", self.active_bundle.to_bytes())
|
||||
self.params.put("ModelManager_ActiveBundle", json.dumps(self.active_bundle.to_dict()))
|
||||
self.selected_bundle = None
|
||||
|
||||
except Exception:
|
||||
@@ -154,7 +166,7 @@ class ModelManagerSP:
|
||||
|
||||
while True:
|
||||
try:
|
||||
self.available_models = self.model_fetcher.get_available_models()
|
||||
self.available_models = self.model_fetcher.get_available_bundles()
|
||||
self.active_bundle = get_active_bundle(self.params)
|
||||
|
||||
if index_to_download := self.params.get("ModelManager_DownloadIndex", block=False, encoding="utf-8"):
|
||||
|
||||
Reference in New Issue
Block a user