-
Notifications
You must be signed in to change notification settings - Fork 0
Refactor ONNX Inference Logic #134
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
03122b4
ae86da4
63a2964
74c3bc0
fdfbc78
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -1,21 +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: | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I like the fact that now the Currently, the main method is I asked chatgpt to refactor current factory using pydantic's BaseModel and Enums, and I actually like this version - it's transparent, clean and also scalable (imagine, adding 3rd backend like TRT). |
||
| _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): | ||
| """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) | ||
| elif "mobilenet" in model_type: | ||
| return MobileNetInference(device=device, weights_path=model_path, class_mapping=class_mapping) | ||
| def create(config: InferenceConfig): | ||
| if config.is_onnx: | ||
| model_cls = InferenceFactory._ONNX_CLASS | ||
| else: | ||
| raise ValueError(f"Unsupported model type: {model_type}") | ||
| 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, | ||
| ) | ||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
I think you should move this file into the
onnxdirectory and rename it toonnx_inference.py(similar tomobilenet_inference.pyandresnet_inference.py).There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
According to the task description, the ONNX file should stay under
src/inference/base/- moving it toonnx/is out of scope for this ticket.