From 03122b4266a0b69f1871b73e2a378873d059810c Mon Sep 17 00:00:00 2001 From: r83575 Date: Tue, 11 Nov 2025 16:24:23 +0200 Subject: [PATCH 1/4] move ONNX inference logic to new structure and update imports --- scripts/onnx_validation.py | 3 +-- src/evaluation/run_evaluation.py | 2 +- .../base/classifier_inference_base_onnx.py} | 0 {scripts => src/inference/utils}/onnx_predict_images.py | 2 +- 4 files changed, 3 insertions(+), 4 deletions(-) rename src/{onnx_model.py => inference/base/classifier_inference_base_onnx.py} (100%) rename {scripts => src/inference/utils}/onnx_predict_images.py (95%) diff --git a/scripts/onnx_validation.py b/scripts/onnx_validation.py index 385cc7a5..6c9e0052 100644 --- a/scripts/onnx_validation.py +++ b/scripts/onnx_validation.py @@ -9,8 +9,7 @@ from src.optimal_class_mapping import MODEL_NAMES as ID_TO_NAME from src.inference.utils.inference_factory import InferenceFactory from src.inference.base.classifier_inference_base import ClassifierInferenceBase -from src.onnx_model import OnnxClassifierInferenceBase - +from src.inference.base.classifier_inference_base_onnx import OnnxClassifierInferenceBase NUM_CLASSES = 83 INPUT_SIZE = (256, 256) diff --git a/src/evaluation/run_evaluation.py b/src/evaluation/run_evaluation.py index 085ade50..82f08936 100644 --- a/src/evaluation/run_evaluation.py +++ b/src/evaluation/run_evaluation.py @@ -12,7 +12,7 @@ from src.optimal_class_mapping import MODEL_NAMES as class_mapping, map_prediction from src.inference.utils.inference_factory import InferenceFactory from src.evaluation.evaluator import EvaluationConfig, ModelEvaluator -from src.onnx_model import OnnxClassifierInferenceBase as OnnxModel +from src.inference.base.classifier_inference_base_onnx import OnnxClassifierInferenceBase as OnnxModel from dataset.optimal_class_mapping import map_prediction from src.path_utils import ensure_clean_directory diff --git a/src/onnx_model.py b/src/inference/base/classifier_inference_base_onnx.py similarity index 100% rename from src/onnx_model.py rename to src/inference/base/classifier_inference_base_onnx.py diff --git a/scripts/onnx_predict_images.py b/src/inference/utils/onnx_predict_images.py similarity index 95% rename from scripts/onnx_predict_images.py rename to src/inference/utils/onnx_predict_images.py index 2745a140..e4b53183 100644 --- a/scripts/onnx_predict_images.py +++ b/src/inference/utils/onnx_predict_images.py @@ -5,7 +5,7 @@ import torch from PIL import Image -from src.onnx_model import OnnxClassifierInferenceBase +from src.inference.base.classifier_inference_base_onnx import OnnxClassifierInferenceBase IMAGE_SIZE = (256, 256) COLOR_MODE = "RGB" From ae86da45b0690d68fbfa42d6696f07227b16a029 Mon Sep 17 00:00:00 2001 From: r83575 Date: Tue, 11 Nov 2025 18:21:23 +0200 Subject: [PATCH 2/4] integrate ONNX model creation into InferenceFactory --- scripts/onnx_validation.py | 9 +++++---- src/evaluation/run_evaluation.py | 9 +++++++-- src/inference/utils/inference_factory.py | 12 ++++++++---- src/inference/utils/onnx_predict_images.py | 9 ++++++--- 4 files changed, 26 insertions(+), 13 deletions(-) diff --git a/scripts/onnx_validation.py b/scripts/onnx_validation.py index 6c9e0052..f2e446af 100644 --- a/scripts/onnx_validation.py +++ b/scripts/onnx_validation.py @@ -6,7 +6,7 @@ import torch from PIL import Image -from src.optimal_class_mapping import MODEL_NAMES as ID_TO_NAME +from dataset.optimal_class_mapping import MODEL_NAMES as ID_TO_NAME from src.inference.utils.inference_factory import InferenceFactory from src.inference.base.classifier_inference_base import ClassifierInferenceBase from src.inference.base.classifier_inference_base_onnx import OnnxClassifierInferenceBase @@ -84,12 +84,13 @@ def main(model_name: str): class_mapping=ID_TO_NAME, ) - onnx_model = OnnxClassifierInferenceBase( + onnx_model = InferenceFactory.create( + model_type="onnx", + model_path=MODELS_DIR_PATH / "onnx" / f"{model_name}.onnx", device=DEVICE, - weights_path=MODELS_DIR_PATH / "onnx" / f"{model_name}.onnx", class_mapping=ID_TO_NAME, topk=1, - ) + ) assert_the_same_shapes(model_name, onnx_model) assert_numerical_accuracy(model_name, torch_model, onnx_model) diff --git a/src/evaluation/run_evaluation.py b/src/evaluation/run_evaluation.py index 82f08936..00f98b4e 100644 --- a/src/evaluation/run_evaluation.py +++ b/src/evaluation/run_evaluation.py @@ -12,7 +12,6 @@ from src.optimal_class_mapping import MODEL_NAMES as class_mapping, map_prediction from src.inference.utils.inference_factory import InferenceFactory from src.evaluation.evaluator import EvaluationConfig, ModelEvaluator -from src.inference.base.classifier_inference_base_onnx import OnnxClassifierInferenceBase as OnnxModel from dataset.optimal_class_mapping import map_prediction from src.path_utils import ensure_clean_directory @@ -44,7 +43,13 @@ def create_model(model_name: str, model_type: str, device: str): if model_type == "pytorch": base_model = InferenceFactory.create(model_name, weights_path.with_suffix(".pt"), device, class_mapping) elif model_type == "onnx": - base_model = OnnxModel(device=device, weights_path=weights_path.with_suffix(".onnx"), topk=1) + base_model = InferenceFactory.create( + model_type="onnx", + model_path=weights_path.with_suffix(".onnx"), + device=device, + class_mapping=class_mapping, + topk=1, + ) else: raise ValueError(f"Unknown model_type: {model_type}") diff --git a/src/inference/utils/inference_factory.py b/src/inference/utils/inference_factory.py index 4a5660ed..2c91367d 100644 --- a/src/inference/utils/inference_factory.py +++ b/src/inference/utils/inference_factory.py @@ -6,16 +6,20 @@ from src.inference.pytorch.resnet_inference import ResNetInference from src.inference.pytorch.mobilenet_inference import MobileNetInference +from src.inference.base.classifier_inference_base_onnx import OnnxClassifierInferenceBase + class InferenceFactory: @staticmethod - def create(model_type: str, model_path: str, device: str, class_mapping=None): + def create(model_type: str, model_path: str, device: str, class_mapping=None, **kwargs): """Create and return an inference model instance based on type.""" model_type = model_type.lower() if "resnet" in model_type: - return ResNetInference(device=device, weights_path=model_path, class_mapping=class_mapping) + return ResNetInference(device=device, weights_path=model_path, class_mapping=class_mapping, **kwargs) elif "mobilenet" in model_type: - return MobileNetInference(device=device, weights_path=model_path, class_mapping=class_mapping) + return MobileNetInference(device=device, weights_path=model_path, class_mapping=class_mapping, **kwargs) + elif "onnx" in model_type or model_path.endswith(".onnx"): + return OnnxClassifierInferenceBase(device=device, weights_path=model_path, class_mapping=class_mapping, **kwargs) else: - raise ValueError(f"Unsupported model type: {model_type}") + raise ValueError(f"Unsupported model type: {model_type}") \ No newline at end of file diff --git a/src/inference/utils/onnx_predict_images.py b/src/inference/utils/onnx_predict_images.py index e4b53183..ec931a9a 100644 --- a/src/inference/utils/onnx_predict_images.py +++ b/src/inference/utils/onnx_predict_images.py @@ -5,7 +5,7 @@ import torch from PIL import Image -from src.inference.base.classifier_inference_base_onnx import OnnxClassifierInferenceBase +from src.inference.utils.inference_factory import InferenceFactory IMAGE_SIZE = (256, 256) COLOR_MODE = "RGB" @@ -29,8 +29,11 @@ def main(): parser.add_argument("--out_dir", required=True, help="Directory to save predictions.") args = parser.parse_args() - model = OnnxClassifierInferenceBase(weights_path=args.onnx_model_weights_path) - model._initialize_model() + model = InferenceFactory.create( + model_type="onnx", + model_path=args.onnx_model_weights_path, + device="cpu", + ) pathlib.Path(args.out_dir).mkdir(parents=True, exist_ok=True) images = sorted([f for f in os.listdir(args.images_dir) if f.lower().endswith((".jpg", ".jpeg", ".png"))]) From 63a2964b6c7f1557658d4b2d7ea8ee9ccac83c93 Mon Sep 17 00:00:00 2001 From: r83575 Date: Tue, 11 Nov 2025 23:53:54 +0200 Subject: [PATCH 3/4] =?UTF-8?q?Refactor=20according=20to=20code=20review?= =?UTF-8?q?=20feedback=20=E2=80=94=20verified=20successful=20run?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- scripts/export_to_onnx.py | 11 +++- scripts/onnx_validation.py | 19 +++--- src/evaluation/run_evaluation.py | 20 ++++-- src/evaluation/run_hierarchical_evaluation.py | 25 ++++++-- src/inference/utils/inference_factory.py | 63 ++++++++++++++----- src/inference/utils/onnx_predict_images.py | 8 ++- src/inference/utils/run_inference_main.py | 12 +++- 7 files changed, 118 insertions(+), 40 deletions(-) diff --git a/scripts/export_to_onnx.py b/scripts/export_to_onnx.py index 6ae1fa98..fa1675de 100644 --- a/scripts/export_to_onnx.py +++ b/scripts/export_to_onnx.py @@ -8,7 +8,7 @@ # Add workspace root to sys.path sys.path.append("/workspace") -from src.inference.utils.inference_factory import InferenceFactory +from src.inference.utils.inference_factory import InferenceFactory, InferenceConfig, Backend, ModelArch from src.path_utils import ensure_clean_directory NUM_CLASSES = 83 @@ -17,12 +17,17 @@ MODELS_DIR_PATH = Path("models") def main(model_name: str): - pytorch_model = InferenceFactory.create( - model_type=model_name, + arch = ModelArch.RESNET if "resnet" in model_name else ModelArch.MOBILENET + + config = InferenceConfig( + arch=arch, + backend=Backend.PYTORCH, model_path=MODELS_DIR_PATH / "pytorch" / f"{model_name}.pt", device=DEVICE, ) + pytorch_model = InferenceFactory.create(config) + # Detect parameter dtype (fp16/fp32) and match input accordingly param_dtype = next( (p.dtype for p in pytorch_model.model.parameters() if p.is_floating_point()), diff --git a/scripts/onnx_validation.py b/scripts/onnx_validation.py index f2e446af..5d4aaef6 100644 --- a/scripts/onnx_validation.py +++ b/scripts/onnx_validation.py @@ -7,7 +7,7 @@ from PIL import Image from dataset.optimal_class_mapping import MODEL_NAMES as ID_TO_NAME -from src.inference.utils.inference_factory import InferenceFactory +from src.inference.utils.inference_factory import InferenceFactory, InferenceConfig, Backend, ModelArch from src.inference.base.classifier_inference_base import ClassifierInferenceBase from src.inference.base.classifier_inference_base_onnx import OnnxClassifierInferenceBase @@ -77,21 +77,26 @@ def assert_same_predictions( def main(model_name: str): - torch_model = InferenceFactory.create( - model_type=model_name, + torch_config = InferenceConfig( + arch=ModelArch.RESNET if "resnet" in model_name else ModelArch.MOBILENET, + backend=Backend.PYTORCH, model_path=MODELS_DIR_PATH / "pytorch" / f"{model_name}.pt", device=DEVICE, class_mapping=ID_TO_NAME, ) - onnx_model = InferenceFactory.create( - model_type="onnx", + onnx_config = InferenceConfig( + arch=None, + backend=Backend.ONNX, model_path=MODELS_DIR_PATH / "onnx" / f"{model_name}.onnx", device=DEVICE, class_mapping=ID_TO_NAME, - topk=1, - ) + extra_args={"topk": 1}, + ) + torch_model = InferenceFactory.create(torch_config) + onnx_model = InferenceFactory.create(onnx_config) + assert_the_same_shapes(model_name, onnx_model) assert_numerical_accuracy(model_name, torch_model, onnx_model) assert_same_predictions(model_name, torch_model, onnx_model) diff --git a/src/evaluation/run_evaluation.py b/src/evaluation/run_evaluation.py index 00f98b4e..13d4c785 100644 --- a/src/evaluation/run_evaluation.py +++ b/src/evaluation/run_evaluation.py @@ -10,7 +10,7 @@ sys.path.append("/workspace") from src.optimal_class_mapping import MODEL_NAMES as class_mapping, map_prediction -from src.inference.utils.inference_factory import InferenceFactory +from src.inference.utils.inference_factory import InferenceFactory, InferenceConfig, Backend, ModelArch from src.evaluation.evaluator import EvaluationConfig, ModelEvaluator from dataset.optimal_class_mapping import map_prediction from src.path_utils import ensure_clean_directory @@ -41,18 +41,28 @@ def create_model(model_name: str, model_type: str, device: str): weights_path = MODELS_DIR_PATH / model_type / f"{model_name}" if model_type == "pytorch": - base_model = InferenceFactory.create(model_name, weights_path.with_suffix(".pt"), device, class_mapping) + config = InferenceConfig( + arch=ModelArch.RESNET if "resnet" in model_name else ModelArch.MOBILENET, + backend=Backend.PYTORCH, + model_path=weights_path.with_suffix(".pt"), + device=device, + class_mapping=class_mapping, + ) elif model_type == "onnx": - base_model = InferenceFactory.create( - model_type="onnx", + config = InferenceConfig( + arch=None, + backend=Backend.ONNX, model_path=weights_path.with_suffix(".onnx"), device=device, class_mapping=class_mapping, - topk=1, + extra_args={"topk": 1}, ) else: raise ValueError(f"Unknown model_type: {model_type}") + base_model = InferenceFactory.create(config) + + return MappedModelWrapper(base_model) diff --git a/src/evaluation/run_hierarchical_evaluation.py b/src/evaluation/run_hierarchical_evaluation.py index 6bbe72d4..4909a7db 100644 --- a/src/evaluation/run_hierarchical_evaluation.py +++ b/src/evaluation/run_hierarchical_evaluation.py @@ -10,9 +10,8 @@ from PIL import Image from dataset.utilities.datasets import DATASETS from metrics.metrics_api import compute_metrics -from dataset_preparation.utilities.datasets import DATASETS from src.evaluation.evaluator import EvaluationConfig, ModelEvaluator -from src.inference.utils.inference_factory import InferenceFactory +from src.inference.utils.inference_factory import InferenceFactory, InferenceConfig, Backend, ModelArch from src.optimal_class_mapping import MODEL_NAMES as class_mapping, map_prediction @@ -71,7 +70,16 @@ def main(): # ResNet print("🔄 ResNet...") - resnet = InferenceFactory.create("resnet", "models/pytorch/resnet18.pt", "cpu") + resnet_config = InferenceConfig( + arch=ModelArch.RESNET, + backend=Backend.PYTORCH, + model_path=Path("models/pytorch/resnet18.pt"), + device="cpu", + class_mapping=class_mapping, + ) + + resnet = InferenceFactory.create(resnet_config) + wrapped_resnet = MappedModelWrapper(resnet) predictions, latencies = [], [] @@ -88,7 +96,16 @@ def main(): # MobileNet print("🔄 MobileNet...") - mobilenet = InferenceFactory.create("mobilenet", "models/pytorch/mobilenet.pt", "cpu") + mobilenet_config = InferenceConfig( + arch=ModelArch.MOBILENET, + backend=Backend.PYTORCH, + model_path=Path("models/pytorch/mobilenet.pt"), + device="cpu", + class_mapping=class_mapping, + ) + + mobilenet = InferenceFactory.create(mobilenet_config) + wrapped_mobilenet = MappedModelWrapper(mobilenet) predictions, latencies = [], [] diff --git a/src/inference/utils/inference_factory.py b/src/inference/utils/inference_factory.py index 2c91367d..4c397f21 100644 --- a/src/inference/utils/inference_factory.py +++ b/src/inference/utils/inference_factory.py @@ -1,25 +1,56 @@ -import os -import sys - -# Add workspace root to sys.path -sys.path.append("/workspace") +from enum import Enum +from pydantic import BaseModel +from pathlib import Path +from typing import Optional, Type, Union from src.inference.pytorch.resnet_inference import ResNetInference from src.inference.pytorch.mobilenet_inference import MobileNetInference from src.inference.base.classifier_inference_base_onnx import OnnxClassifierInferenceBase +class Backend(str, Enum): + PYTORCH = "pytorch" + ONNX = "onnx" + + +class ModelArch(str, Enum): + RESNET = "resnet" + MOBILENET = "mobilenet" + + +class InferenceConfig(BaseModel): + arch: Optional[ModelArch] + backend: Backend + model_path: Union[str, Path] + device: str + class_mapping: Optional[dict] = None + extra_args: dict = {} + + @property + def is_onnx(self) -> bool: + return self.backend == Backend.ONNX or str(self.model_path).endswith(".onnx") + + class InferenceFactory: + _PYTORCH_MAP: dict[ModelArch, Type] = { + ModelArch.RESNET: ResNetInference, + ModelArch.MOBILENET: MobileNetInference, + } + + _ONNX_CLASS: Type = OnnxClassifierInferenceBase + @staticmethod - def create(model_type: str, model_path: str, device: str, class_mapping=None, **kwargs): - """Create and return an inference model instance based on type.""" - model_type = model_type.lower() - - if "resnet" in model_type: - return ResNetInference(device=device, weights_path=model_path, class_mapping=class_mapping, **kwargs) - elif "mobilenet" in model_type: - return MobileNetInference(device=device, weights_path=model_path, class_mapping=class_mapping, **kwargs) - elif "onnx" in model_type or model_path.endswith(".onnx"): - return OnnxClassifierInferenceBase(device=device, weights_path=model_path, class_mapping=class_mapping, **kwargs) + def create(config: InferenceConfig): + if config.is_onnx: + model_cls = InferenceFactory._ONNX_CLASS else: - raise ValueError(f"Unsupported model type: {model_type}") \ No newline at end of file + model_cls = InferenceFactory._PYTORCH_MAP.get(config.arch) + if not model_cls: + raise ValueError(f"Unsupported PyTorch architecture: {config.arch}") + + return model_cls( + device=config.device, + weights_path=config.model_path, + class_mapping=config.class_mapping, + **config.extra_args, + ) diff --git a/src/inference/utils/onnx_predict_images.py b/src/inference/utils/onnx_predict_images.py index ec931a9a..9bda588d 100644 --- a/src/inference/utils/onnx_predict_images.py +++ b/src/inference/utils/onnx_predict_images.py @@ -5,7 +5,7 @@ import torch from PIL import Image -from src.inference.utils.inference_factory import InferenceFactory +from src.inference.utils.inference_factory import InferenceFactory, InferenceConfig, Backend IMAGE_SIZE = (256, 256) COLOR_MODE = "RGB" @@ -29,11 +29,13 @@ def main(): parser.add_argument("--out_dir", required=True, help="Directory to save predictions.") args = parser.parse_args() - model = InferenceFactory.create( - model_type="onnx", + config = InferenceConfig( + arch=None, + backend=Backend.ONNX, model_path=args.onnx_model_weights_path, device="cpu", ) + model = InferenceFactory.create(config) pathlib.Path(args.out_dir).mkdir(parents=True, exist_ok=True) images = sorted([f for f in os.listdir(args.images_dir) if f.lower().endswith((".jpg", ".jpeg", ".png"))]) diff --git a/src/inference/utils/run_inference_main.py b/src/inference/utils/run_inference_main.py index f099bd3f..7e321509 100644 --- a/src/inference/utils/run_inference_main.py +++ b/src/inference/utils/run_inference_main.py @@ -11,7 +11,7 @@ sys.path.append("/workspace") from src.optimal_class_mapping import MODEL_NAMES as ID_TO_NAME -from inference_factory import InferenceFactory +from src.inference.utils.inference_factory import InferenceFactory, InferenceConfig, Backend, ModelArch # Main image processing logic @@ -20,7 +20,15 @@ def run_inference(model_type: str, model_path: str, image_dir: str, device: str) class_mapping = ID_TO_NAME if ID_TO_NAME else None # create model using factory - clf = InferenceFactory.create(model_type, model_path, device, class_mapping) + config = InferenceConfig( + arch=ModelArch.RESNET if "resnet" in model_type else ModelArch.MOBILENET, + backend=Backend.PYTORCH if model_path.endswith(".pt") else Backend.ONNX, + model_path=Path(model_path), + device=device, + class_mapping=class_mapping, + ) + + clf = InferenceFactory.create(config) results = {} total_time = 0 From 74c3bc094027191e601bcf928ada7bda2d0be6b9 Mon Sep 17 00:00:00 2001 From: r83575 Date: Wed, 12 Nov 2025 00:16:13 +0200 Subject: [PATCH 4/4] use InferenceConfig and Backend enums cleanly for ONNX prediction script --- scripts/export_to_onnx.py | 5 ++-- scripts/onnx_validation.py | 32 ++++++++++------------- src/inference/utils/run_inference_main.py | 8 ++++-- 3 files changed, 23 insertions(+), 22 deletions(-) diff --git a/scripts/export_to_onnx.py b/scripts/export_to_onnx.py index fa1675de..3bba5677 100644 --- a/scripts/export_to_onnx.py +++ b/scripts/export_to_onnx.py @@ -17,11 +17,12 @@ MODELS_DIR_PATH = Path("models") def main(model_name: str): - arch = ModelArch.RESNET if "resnet" in model_name else ModelArch.MOBILENET + arch = ModelArch.RESNET if "resnet" in model_name.lower() else ModelArch.MOBILENET + backend = Backend.PYTORCH config = InferenceConfig( arch=arch, - backend=Backend.PYTORCH, + backend=backend, model_path=MODELS_DIR_PATH / "pytorch" / f"{model_name}.pt", device=DEVICE, ) diff --git a/scripts/onnx_validation.py b/scripts/onnx_validation.py index 5d4aaef6..6fe64c81 100644 --- a/scripts/onnx_validation.py +++ b/scripts/onnx_validation.py @@ -77,17 +77,25 @@ def assert_same_predictions( def main(model_name: str): + # define enums separately for clarity + torch_arch = ModelArch.RESNET if "resnet" in model_name.lower() else ModelArch.MOBILENET + torch_backend = Backend.PYTORCH + + onnx_arch = None # ONNX doesn't need a specific architecture + onnx_backend = Backend.ONNX + + # PyTorch model config torch_config = InferenceConfig( - arch=ModelArch.RESNET if "resnet" in model_name else ModelArch.MOBILENET, - backend=Backend.PYTORCH, + arch=torch_arch, + backend=torch_backend, model_path=MODELS_DIR_PATH / "pytorch" / f"{model_name}.pt", device=DEVICE, class_mapping=ID_TO_NAME, ) - + onnx_config = InferenceConfig( - arch=None, - backend=Backend.ONNX, + arch=onnx_arch, + backend=onnx_backend, model_path=MODELS_DIR_PATH / "onnx" / f"{model_name}.onnx", device=DEVICE, class_mapping=ID_TO_NAME, @@ -96,19 +104,7 @@ def main(model_name: str): torch_model = InferenceFactory.create(torch_config) onnx_model = InferenceFactory.create(onnx_config) - + assert_the_same_shapes(model_name, onnx_model) assert_numerical_accuracy(model_name, torch_model, onnx_model) assert_same_predictions(model_name, torch_model, onnx_model) - - -if __name__ == "__main__": - p = argparse.ArgumentParser("Validate ONNX model against PyTorch") - p.add_argument( - "model_name", - type=str, - choices=["resnet18", "mobilenet"], - help="Name of the model to validate", - ) - args = p.parse_args() - main(args.model_name) diff --git a/src/inference/utils/run_inference_main.py b/src/inference/utils/run_inference_main.py index 7e321509..7eb46a82 100644 --- a/src/inference/utils/run_inference_main.py +++ b/src/inference/utils/run_inference_main.py @@ -20,14 +20,18 @@ def run_inference(model_type: str, model_path: str, image_dir: str, device: str) class_mapping = ID_TO_NAME if ID_TO_NAME else None # create model using factory + arch = ModelArch.RESNET if "resnet" in model_type.lower() else ModelArch.MOBILENET + backend = Backend.ONNX if model_path.endswith(".onnx") else Backend.PYTORCH + config = InferenceConfig( - arch=ModelArch.RESNET if "resnet" in model_type else ModelArch.MOBILENET, - backend=Backend.PYTORCH if model_path.endswith(".pt") else Backend.ONNX, + arch=arch, + backend=backend, model_path=Path(model_path), device=device, class_mapping=class_mapping, ) + clf = InferenceFactory.create(config) results = {}