This commit is contained in:
Jason Wen
2025-01-05 06:28:46 -05:00
parent c54cc074e2
commit 59c551ac77
+5 -2
View File
@@ -21,6 +21,7 @@ async def verify_file(file_path: str, expected_hash: str) -> bool:
return sha256_hash.hexdigest().lower() == expected_hash.lower()
def get_active_bundle(params: Params) -> custom.ModelManagerSP.ModelBundle:
"""Gets the active model bundle from cache"""
if params is None:
@@ -31,13 +32,15 @@ def get_active_bundle(params: Params) -> custom.ModelManagerSP.ModelBundle:
return None
def get_model_runner_by_filename(filename: str) -> custom.ModelManagerSP.Runner:
if filename.endswith(".thneed"):
return custom.ModelManagerSP.Runner.snpe
if filename.endswith("_tinygrad.pkl"):
return custom.ModelManagerSP.Runner.tinygrad
def get_active_model_runner(params: Params) -> custom.ModelManagerSP.Runner:
"""Gets the model runner from the active model bundle. If no active bundle, returns tinygrad"""
if params is None:
@@ -46,5 +49,5 @@ def get_active_model_runner(params: Params) -> custom.ModelManagerSP.Runner:
if active_bundle := get_active_bundle(params):
drive_model = next(model for model in active_bundle.models if model.type == custom.ModelManagerSP.Type.drive)
return get_model_runner_by_filename(drive_model.fileName)
return custom.ModelManagerSP.Runner.tinygrad