From 6105aab633191edaa0f2f4a907c0fd7430024b77 Mon Sep 17 00:00:00 2001 From: Sarah Gershuni Date: Tue, 11 Nov 2025 02:42:18 +0200 Subject: [PATCH 1/5] refactor pytorch-related inference logic --- src/inference.py | 314 ------------------ .../base/classifier_inference_base.py | 96 ++++++ src/inference/pytorch/mobilenet_inference.py | 68 ++++ src/inference/pytorch/resnet_inference.py | 67 ++++ src/inference/utils/inference_factory.py | 28 ++ src/inference/utils/run_inference_main.py | 94 ++++++ 6 files changed, 353 insertions(+), 314 deletions(-) delete mode 100644 src/inference.py create mode 100644 src/inference/base/classifier_inference_base.py create mode 100644 src/inference/pytorch/mobilenet_inference.py create mode 100644 src/inference/pytorch/resnet_inference.py create mode 100644 src/inference/utils/inference_factory.py create mode 100644 src/inference/utils/run_inference_main.py diff --git a/src/inference.py b/src/inference.py deleted file mode 100644 index 6eedb173..00000000 --- a/src/inference.py +++ /dev/null @@ -1,314 +0,0 @@ -import argparse -import json -import os -import time -from abc import ABC, abstractmethod -from pathlib import Path -from typing import Dict, List, Optional, Union - -import numpy as np -import torch -import torch.nn.functional as F -from PIL import Image -from torch import Tensor -from torchvision import models, transforms - -# Import the complete mapping -try: - from src.optimal_class_mapping import MODEL_NAMES as ID_TO_NAME - - print(f"Loaded {len(ID_TO_NAME)} classes from CropWeed dataset") -except ImportError: - print("Warning: optimal_class_mapping.py not found, using fallback") - ID_TO_NAME = {0: "unknown"} - - -def _to_pil_image(image: Union[np.ndarray, Image.Image]) -> Image.Image: - """Convert input to PIL Image.""" - if isinstance(image, np.ndarray): - return Image.fromarray(image).convert("RGB") - elif isinstance(image, Image.Image): - return image.convert("RGB") - else: - raise TypeError(f"Unsupported image type: {type(image)}") - - -class ClassifierInferenceBase(ABC): - """ - Abstract base for image classification inference. - Subclasses must implement `_initialize_model` and `_forward`. - """ - - def __init__( - self, - device: Union[str, torch.device] = "cpu", - weights_path: Optional[Union[str, Path]] = None, - class_mapping: Optional[Dict[int, str]] = None, - topk: int = 5, - transform: Optional[transforms.Compose] = None, - ) -> None: - self.device = torch.device(device) - self.weights_path = ( - Path(str(weights_path)) if weights_path is not None else None - ) - self.class_mapping = class_mapping - self.topk = max(1, int(topk)) - self._transform = ( - transform if transform is not None else self._build_preprocess() - ) - self._initialize_model() - - @abstractmethod - def _initialize_model(self) -> None: - """Create/load the model and put it into eval() on the right device.""" - raise NotImplementedError - - @abstractmethod - def _forward(self, x: Tensor) -> Tensor: - """Return raw logits of shape [N, C].""" - raise NotImplementedError - - def _build_preprocess(self) -> transforms.Compose: - """Default preprocessing - תואם לאימון המודל""" - return transforms.Compose( - [ - transforms.Resize((256, 256)), - transforms.ToTensor(), - transforms.Normalize( - mean=[0.5, 0.5, 0.5], - std=[0.25, 0.25, 0.25], - ), - ] - ) - - @torch.inference_mode() - def infer(self, image: Union[np.ndarray, Image.Image]) -> List[Dict]: - """ - Run single-image inference and return top-k predictions: - [{'class_id': int, 'class_name': str, 'probability': float}, ...] - """ - pil_img = _to_pil_image(image) - x = self._transform(pil_img).unsqueeze(0).to(self.device) - logits = self._forward(x) - probs = F.softmax(logits, dim=1)[0] - k = min(self.topk, probs.numel()) - top_probs, top_idx = torch.topk(probs, k) - out: List[Dict] = [] - for p, idx in zip(top_probs.tolist(), top_idx.tolist()): - name = self._class_name(idx) - out.append( - { - "class_id": int(idx), - "class_name": name, - "probability": float(p), - } - ) - return out - - def _class_name(self, class_id: int) -> str: - if self.class_mapping and class_id in self.class_mapping: - return self.class_mapping[class_id] - return f"class_{class_id}" - - -class ResNetInference(ClassifierInferenceBase): - """ResNet18 classifier that strictly requires a - checkpoint at `weights_path`.""" - - def __init__( - self, - device: Union[str, torch.device] = "cpu", - weights_path: Union[str, Path] = "", - num_classes: int = 83, - class_mapping: Optional[Dict[int, str]] = None, - topk: int = 5, - transform: Optional[transforms.Compose] = None, - strict: bool = True, - ) -> None: - self.num_classes = int(num_classes) - self.strict = bool(strict) - super().__init__( - device=device, - weights_path=weights_path, - class_mapping=class_mapping, - topk=topk, - transform=transform, - ) - - def _initialize_model(self) -> None: - if self.weights_path is None or not self.weights_path.exists(): - raise FileNotFoundError(f"Checkpoint not found: {self.weights_path}") - - model = models.resnet18(pretrained=False) - model.fc = torch.nn.Linear(model.fc.in_features, self.num_classes) - - ckpt = torch.load( - self.weights_path.as_posix(), - map_location="cpu", - weights_only=False, - ) - if isinstance(ckpt, dict): - state_dict = ckpt.get("model_state_dict") or ckpt.get("state_dict") or ckpt - else: - raise RuntimeError( - "Unrecognized checkpoint format (expected dict/state_dict)" - ) - - if any(k.startswith("model.") for k in state_dict.keys()): - state_dict = {k.replace("model.", "", 1): v for k, v in state_dict.items()} - - model.load_state_dict(state_dict, strict=self.strict) - self.model = model.to(self.device).eval() - - @torch.inference_mode() - def _forward(self, x: Tensor) -> Tensor: - return self.model(x) - - -class MobileNetInference(ClassifierInferenceBase): - """MobileNetV2 classifier that strictly requires a checkpoint at `weights_path`.""" - - def __init__( - self, - device: Union[str, torch.device] = "cpu", - weights_path: Union[str, Path] = "", - num_classes: int = 83, - class_mapping: Optional[Dict[int, str]] = None, - topk: int = 5, - transform: Optional[transforms.Compose] = None, - strict: bool = True, - ) -> None: - self.num_classes = int(num_classes) - self.strict = bool(strict) - super().__init__( - device=device, - weights_path=weights_path, - class_mapping=class_mapping, - topk=topk, - transform=transform, - ) - - def _initialize_model(self) -> None: - if self.weights_path is None or not self.weights_path.exists(): - raise FileNotFoundError(f"Checkpoint not found: {self.weights_path}") - - model = models.mobilenet_v2(pretrained=False) - model.classifier[1] = torch.nn.Linear( - model.classifier[1].in_features, - self.num_classes, - ) - - ckpt = torch.load( - self.weights_path.as_posix(), - map_location="cpu", - weights_only=False, - ) - if isinstance(ckpt, dict): - state_dict = ckpt.get("model_state_dict") or ckpt.get("state_dict") or ckpt - else: - raise RuntimeError( - "Unrecognized checkpoint format (expected dict/state_dict)" - ) - - if any(k.startswith("model.") for k in state_dict.keys()): - state_dict = {k.replace("model.", "", 1): v for k, v in state_dict.items()} - - model.load_state_dict(state_dict, strict=self.strict) - self.model = model.to(self.device).eval() - - @torch.inference_mode() - def _forward(self, x: Tensor) -> Tensor: - return self.model(x) - - -def process_images_with_oop( - model_type: str, model_path: str, image_dir: str, device: str -): - """Process images using OOP approach for backward compatibility.""" - class_mapping = ID_TO_NAME if ID_TO_NAME else None - - if model_type == "resnet": - clf = ResNetInference( - device=device, - weights_path=model_path, - class_mapping=class_mapping, - ) - elif model_type == "mobilenet": - clf = MobileNetInference( - device=device, - weights_path=model_path, - class_mapping=class_mapping, - ) - else: - raise ValueError(f"Unsupported model type: {model_type}") - - results = {} - total_time = 0 - image_count = 0 - img_extensions = {".jpg", ".jpeg", ".png", ".bmp"} - - for filename in os.listdir(image_dir): - file_ext = os.path.splitext(filename)[1].lower() - if file_ext in img_extensions: - image_path = os.path.join(image_dir, filename) - print(f"Processing {filename}...") - - try: - start_time = time.time() - image = Image.open(image_path) - predictions = clf.infer(image) - inference_time = (time.time() - start_time) * 1000 - - results[filename] = { - "predictions": predictions, - "inference_time_ms": inference_time, - } - print( - f" {predictions[0]['class_name']} " - f"({predictions[0]['probability']:.1%})" - ) - - total_time += inference_time - image_count += 1 - - except Exception as e: - print(f" Error: {e}") - results[filename] = {"error": str(e)} - - avg_time = total_time / image_count if image_count > 0 else 0 - print(f"\nAverage inference time: {avg_time:.1f}ms") - return results, avg_time - - -def main(): - """Main function to run inference on images.""" - parser = argparse.ArgumentParser(description="Run inference on images") - parser.add_argument( - "--model", - choices=["mobilenet", "resnet"], - default="mobilenet", - ) - parser.add_argument("--model_path", default="../models/pytorch/mobilenet.pt") - parser.add_argument("--image_dir", default="../assets/test_images") - parser.add_argument("--output_file", default="outputs/predictions.json") - - args = parser.parse_args() - device = torch.device("cuda" if torch.cuda.is_available() else "cpu") - print(f"Starting inference on device: {device}") - - results, avg_time = process_images_with_oop( - args.model, - args.model_path, - args.image_dir, - str(device), - ) - - os.makedirs(os.path.dirname(args.output_file), exist_ok=True) - with open(args.output_file, "w") as f: - json.dump(results, f, indent=2) - - print(f"Results saved to {args.output_file}") - - -if __name__ == "__main__": - main() diff --git a/src/inference/base/classifier_inference_base.py b/src/inference/base/classifier_inference_base.py new file mode 100644 index 00000000..4176d284 --- /dev/null +++ b/src/inference/base/classifier_inference_base.py @@ -0,0 +1,96 @@ +from abc import ABC, abstractmethod +from pathlib import Path +from typing import Dict, List, Optional, Union + +import numpy as np +import torch +import torch.nn.functional as F +from PIL import Image +from torch import Tensor +from torchvision import transforms + +class ClassifierInferenceBase(ABC): + """ + Abstract base for image classification inference. + Subclasses must implement `_initialize_model` and `_forward`. + """ + + def __init__( + self, + device: Union[str, torch.device] = "cpu", + weights_path: Optional[Union[str, Path]] = None, + class_mapping: Optional[Dict[int, str]] = None, + topk: int = 5, + transform: Optional[transforms.Compose] = None, + ) -> None: + self.device = torch.device(device) + self.weights_path = ( + Path(str(weights_path)) if weights_path is not None else None + ) + self.class_mapping = class_mapping + self.topk = max(1, int(topk)) + self._transform = ( + transform if transform is not None else self._build_preprocess() + ) + self._initialize_model() + + @abstractmethod + def _initialize_model(self) -> None: + """Create/load the model and put it into eval() on the right device.""" + raise NotImplementedError + + @abstractmethod + def _forward(self, x: Tensor) -> Tensor: + """Return raw logits of shape [N, C].""" + raise NotImplementedError + + def _build_preprocess(self) -> transforms.Compose: + """Default preprocessing - תואם לאימון המודל""" + return transforms.Compose( + [ + transforms.Resize((256, 256)), + transforms.ToTensor(), + transforms.Normalize( + mean=[0.5, 0.5, 0.5], + std=[0.25, 0.25, 0.25], + ), + ] + ) + + def _to_pil_image(self, image: Union[np.ndarray, Image.Image]) -> Image.Image: + """Convert input to PIL Image.""" + if isinstance(image, np.ndarray): + return Image.fromarray(image).convert("RGB") + elif isinstance(image, Image.Image): + return image.convert("RGB") + else: + raise TypeError(f"Unsupported image type: {type(image)}") + + @torch.inference_mode() + def infer(self, image: Union[np.ndarray, Image.Image]) -> List[Dict]: + """ + Run single-image inference and return top-k predictions: + [{'class_id': int, 'class_name': str, 'probability': float}, ...] + """ + pil_img = self._to_pil_image(image) + x = self._transform(pil_img).unsqueeze(0).to(self.device) + logits = self._forward(x) + probs = F.softmax(logits, dim=1)[0] + k = min(self.topk, probs.numel()) + top_probs, top_idx = torch.topk(probs, k) + out: List[Dict] = [] + for p, idx in zip(top_probs.tolist(), top_idx.tolist()): + name = self._class_name(idx) + out.append( + { + "class_id": int(idx), + "class_name": name, + "probability": float(p), + } + ) + return out + + def _class_name(self, class_id: int) -> str: + if self.class_mapping and class_id in self.class_mapping: + return self.class_mapping[class_id] + return f"class_{class_id}" diff --git a/src/inference/pytorch/mobilenet_inference.py b/src/inference/pytorch/mobilenet_inference.py new file mode 100644 index 00000000..3629a75a --- /dev/null +++ b/src/inference/pytorch/mobilenet_inference.py @@ -0,0 +1,68 @@ +import os +import sys +from pathlib import Path +from typing import Dict, Optional, Union + +# Add workspace root to sys.path +sys.path.append("/workspace") + +import torch +from torch import Tensor +from torchvision import models, transforms + +from src.inference.base.classifier_inference_base import ClassifierInferenceBase + +class MobileNetInference(ClassifierInferenceBase): + """MobileNetV2 classifier that strictly requires a checkpoint at `weights_path`.""" + + def __init__( + self, + device: Union[str, torch.device] = "cpu", + weights_path: Union[str, Path] = "", + num_classes: int = 83, + class_mapping: Optional[Dict[int, str]] = None, + topk: int = 5, + transform: Optional[transforms.Compose] = None, + strict: bool = True, + ) -> None: + self.num_classes = int(num_classes) + self.strict = bool(strict) + super().__init__( + device=device, + weights_path=weights_path, + class_mapping=class_mapping, + topk=topk, + transform=transform, + ) + + def _initialize_model(self) -> None: + if self.weights_path is None or not self.weights_path.exists(): + raise FileNotFoundError(f"Checkpoint not found: {self.weights_path}") + + model = models.mobilenet_v2(pretrained=False) + model.classifier[1] = torch.nn.Linear( + model.classifier[1].in_features, + self.num_classes, + ) + + ckpt = torch.load( + self.weights_path.as_posix(), + map_location="cpu", + weights_only=False, + ) + if isinstance(ckpt, dict): + state_dict = ckpt.get("model_state_dict") or ckpt.get("state_dict") or ckpt + else: + raise RuntimeError( + "Unrecognized checkpoint format (expected dict/state_dict)" + ) + + if any(k.startswith("model.") for k in state_dict.keys()): + state_dict = {k.replace("model.", "", 1): v for k, v in state_dict.items()} + + model.load_state_dict(state_dict, strict=self.strict) + self.model = model.to(self.device).eval() + + @torch.inference_mode() + def _forward(self, x: Tensor) -> Tensor: + return self.model(x) diff --git a/src/inference/pytorch/resnet_inference.py b/src/inference/pytorch/resnet_inference.py new file mode 100644 index 00000000..8054d83d --- /dev/null +++ b/src/inference/pytorch/resnet_inference.py @@ -0,0 +1,67 @@ +import os +import sys +from pathlib import Path +from typing import Dict, Optional, Union + +# Add workspace root to sys.path +sys.path.append("/workspace") + +import torch +from torch import Tensor +from torchvision import models, transforms + +from src.inference.base.classifier_inference_base import ClassifierInferenceBase + + +class ResNetInference(ClassifierInferenceBase): + """ResNet18 classifier that strictly requires a + checkpoint at `weights_path`.""" + + def __init__( + self, + device: Union[str, torch.device] = "cpu", + weights_path: Union[str, Path] = "", + num_classes: int = 83, + class_mapping: Optional[Dict[int, str]] = None, + topk: int = 5, + transform: Optional[transforms.Compose] = None, + strict: bool = True, + ) -> None: + self.num_classes = int(num_classes) + self.strict = bool(strict) + super().__init__( + device=device, + weights_path=weights_path, + class_mapping=class_mapping, + topk=topk, + transform=transform, + ) + + def _initialize_model(self) -> None: + if self.weights_path is None or not self.weights_path.exists(): + raise FileNotFoundError(f"Checkpoint not found: {self.weights_path}") + + model = models.resnet18(pretrained=False) + model.fc = torch.nn.Linear(model.fc.in_features, self.num_classes) + + ckpt = torch.load( + self.weights_path.as_posix(), + map_location="cpu", + weights_only=False, + ) + if isinstance(ckpt, dict): + state_dict = ckpt.get("model_state_dict") or ckpt.get("state_dict") or ckpt + else: + raise RuntimeError( + "Unrecognized checkpoint format (expected dict/state_dict)" + ) + + if any(k.startswith("model.") for k in state_dict.keys()): + state_dict = {k.replace("model.", "", 1): v for k, v in state_dict.items()} + + model.load_state_dict(state_dict, strict=self.strict) + self.model = model.to(self.device).eval() + + @torch.inference_mode() + def _forward(self, x: Tensor) -> Tensor: + return self.model(x) diff --git a/src/inference/utils/inference_factory.py b/src/inference/utils/inference_factory.py new file mode 100644 index 00000000..37da19ed --- /dev/null +++ b/src/inference/utils/inference_factory.py @@ -0,0 +1,28 @@ +import os +import sys + +# Add workspace root to sys.path +sys.path.append("/workspace") + +from src.inference.pytorch.resnet_inference import ResNetInference +from src.inference.pytorch.mobilenet_inference import MobileNetInference + +class InferenceFactory: + @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) + else: + raise ValueError(f"Unsupported model type: {model_type}") + + model_class = mapping[model_type] + return model_class( + device=device, + weights_path=model_path, + class_mapping=class_mapping, + ) diff --git a/src/inference/utils/run_inference_main.py b/src/inference/utils/run_inference_main.py new file mode 100644 index 00000000..7145e056 --- /dev/null +++ b/src/inference/utils/run_inference_main.py @@ -0,0 +1,94 @@ +import os +import sys +import time +import json +import argparse + +from PIL import Image +import torch + +# Add workspace root to sys.path +sys.path.append("/workspace") + +from src.optimal_class_mapping import MODEL_NAMES as ID_TO_NAME +from inference_factory import InferenceFactory + + +# Main image processing logic +def process_images_with_oop(model_type: str, model_path: str, image_dir: str, device: str): + """Process images using OOP approach with Factory pattern.""" + 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) + + results = {} + total_time = 0 + image_count = 0 + img_extensions = {".jpg", ".jpeg", ".png", ".bmp"} + + for filename in os.listdir(image_dir): + file_ext = os.path.splitext(filename)[1].lower() + if file_ext not in img_extensions: + continue + + image_path = os.path.join(image_dir, filename) + print(f"Processing {filename}...") + + try: + start_time = time.time() + image = Image.open(image_path) + predictions = clf.infer(image) + inference_time = (time.time() - start_time) * 1000 + + results[filename] = { + "predictions": predictions, + "inference_time_ms": inference_time, + } + + print( + f" {predictions[0]['class_name']} " + f"({predictions[0]['probability']:.1%})" + ) + + total_time += inference_time + image_count += 1 + + except Exception as e: + print(f" Error: {e}") + results[filename] = {"error": str(e)} + + avg_time = total_time / image_count if image_count > 0 else 0 + print(f"\nAverage inference time: {avg_time:.1f}ms") + return results, avg_time + + +# CLI entry point +def main(): + """Main function to run inference on images.""" + parser = argparse.ArgumentParser(description="Run inference on images") + parser.add_argument("--model", choices=["mobilenet", "resnet"], default="mobilenet") + parser.add_argument("--model_path", default="../models/pytorch/mobilenet.pt") + parser.add_argument("--image_dir", default="../assets/test_images") + parser.add_argument("--output_file", default="outputs/predictions.json") + + args = parser.parse_args() + device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + print(f"Starting inference on device: {device}") + + results, avg_time = process_images_with_oop( + args.model, + args.model_path, + args.image_dir, + str(device), + ) + + os.makedirs(os.path.dirname(args.output_file), exist_ok=True) + with open(args.output_file, "w") as f: + json.dump(results, f, indent=2) + + print(f"Results saved to {args.output_file}") + + +if __name__ == "__main__": + main() From cc6ac8e2b878cbf3a24075fc686960b7dbb8fe2a Mon Sep 17 00:00:00 2001 From: Sarah Gershuni Date: Tue, 11 Nov 2025 02:43:07 +0200 Subject: [PATCH 2/5] update all references to inference.py and integrate factory pattern changes --- scripts/export_to_onnx.py | 31 +++++++------------ scripts/onnx_validation.py | 30 +++++++----------- src/evaluation/evaluator.py | 8 +++-- src/evaluation/run_evaluation.py | 27 +++++++--------- src/evaluation/run_hierarchical_evaluation.py | 16 ++++++---- src/onnx_model.py | 10 +++--- 6 files changed, 56 insertions(+), 66 deletions(-) diff --git a/scripts/export_to_onnx.py b/scripts/export_to_onnx.py index 933dd6a9..2c0e2fa5 100644 --- a/scripts/export_to_onnx.py +++ b/scripts/export_to_onnx.py @@ -1,10 +1,15 @@ #!/usr/bin/env python3 +import os +import sys import argparse from pathlib import Path import torch -from src.inference import MobileNetInference, ResNetInference +# Add workspace root to sys.path +sys.path.append("/workspace") + +from src.inference.utils.inference_factory import InferenceFactory from src.path_utils import ensure_clean_directory NUM_CLASSES = 83 @@ -12,26 +17,12 @@ DEVICE = "cpu" MODELS_DIR_PATH = Path("models") - -def get_model(model_name: str): - if model_name == "resnet18": - return ResNetInference( - device=DEVICE, - weights_path=MODELS_DIR_PATH / "pytorch" / f"{model_name}.pt", - num_classes=NUM_CLASSES, - ) - elif model_name == "mobilenet": - return MobileNetInference( - device=DEVICE, - weights_path=MODELS_DIR_PATH / "pytorch" / f"{model_name}.pt", - num_classes=NUM_CLASSES, - ) - else: - raise ValueError(f"Unsupported model: {model_name}") - - def main(model_name: str): - pytorch_model = get_model(model_name) + pytorch_model = InferenceFactory.create( + model_type=model_name, + model_path=MODELS_DIR_PATH / "pytorch" / f"{model_name}.pt", + device=DEVICE, + ) # Detect parameter dtype (fp16/fp32) and match input accordingly param_dtype = next( diff --git a/scripts/onnx_validation.py b/scripts/onnx_validation.py index c0adf1f9..35f5dda8 100644 --- a/scripts/onnx_validation.py +++ b/scripts/onnx_validation.py @@ -1,15 +1,18 @@ #!/usr/bin/env python3 -import argparse import os +import argparse from pathlib import Path import numpy as np import torch from PIL import Image -from src.inference import ID_TO_NAME, ClassifierInferenceBase, MobileNetInference, ResNetInference +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 + NUM_CLASSES = 83 INPUT_SIZE = (256, 256) @@ -76,22 +79,13 @@ def assert_same_predictions( def main(model_name: str): - if model_name == "resnet18": - torch_model = ResNetInference( - device=DEVICE, - weights_path=MODELS_DIR_PATH / "pytorch" / f"{model_name}.pt", - num_classes=NUM_CLASSES, - strict=True, - ) - elif model_name == "mobilenet": - torch_model = MobileNetInference( - device=DEVICE, - weights_path=MODELS_DIR_PATH / "pytorch" / f"{model_name}.pt", - num_classes=NUM_CLASSES, - strict=True, - ) - else: - raise ValueError(f"Unsupported model: {model_name}") + torch_model = InferenceFactory.create( + model_type=model_name, + model_path=MODELS_DIR_PATH / "pytorch" / f"{model_name}.pt", + device=DEVICE, + class_mapping=ID_TO_NAME, + ) + onnx_model = OnnxClassifierInferenceBase( device=DEVICE, weights_path=MODELS_DIR_PATH / "onnx" / f"{model_name}.onnx", diff --git a/src/evaluation/evaluator.py b/src/evaluation/evaluator.py index 2276d780..757c4a7e 100644 --- a/src/evaluation/evaluator.py +++ b/src/evaluation/evaluator.py @@ -1,14 +1,18 @@ +import os +import sys import time from dataclasses import dataclass from pathlib import Path from typing import List, Optional, Tuple +# Add workspace root to sys.path +sys.path.append("/workspace") + import torch from PIL import Image +from src.inference.base.classifier_inference_base import ClassifierInferenceBase as InferenceModel from metrics.metrics_api import ClassificationReport, compute_metrics -from src.inference import ClassifierInferenceBase as InferenceModel - @dataclass class EvaluationConfig: diff --git a/src/evaluation/run_evaluation.py b/src/evaluation/run_evaluation.py index 10748038..b19ddd25 100644 --- a/src/evaluation/run_evaluation.py +++ b/src/evaluation/run_evaluation.py @@ -1,18 +1,22 @@ +import os +import sys import argparse import json from pathlib import Path import pandas as pd +# Add workspace root to sys.path +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.evaluation.evaluator import EvaluationConfig, ModelEvaluator -from src.inference import MobileNetInference, ResNetInference from src.onnx_model import OnnxClassifierInferenceBase as OnnxModel -from src.optimal_class_mapping import map_prediction from src.path_utils import ensure_clean_directory -MODELS_DIR_PATH = Path("models") -NUM_CLASSES = 83 +MODELS_DIR_PATH = Path("models") class MappedModelWrapper: """Wrapper that adds mapping between 83 model classes to 76 dataset classes""" @@ -34,21 +38,12 @@ def infer(self, image): def create_model(model_name: str, model_type: str, device: str): - weights_path = MODELS_DIR_PATH / model_type / f"{model_name}.pt" + weights_path = MODELS_DIR_PATH / model_type / f"{model_name}" if model_type == "pytorch": - if model_name == "resnet18": - base_model = ResNetInference( - device=device, weights_path=weights_path, num_classes=NUM_CLASSES - ) - elif model_name == "mobilenet": - base_model = MobileNetInference( - device=device, weights_path=weights_path, num_classes=NUM_CLASSES - ) - else: - raise ValueError(f"Unknown model: {model_name}") + 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, topk=1) + base_model = OnnxModel(device=device, weights_path=weights_path.with_suffix(".onnx"), topk=1) else: raise ValueError(f"Unknown model_type: {model_type}") diff --git a/src/evaluation/run_hierarchical_evaluation.py b/src/evaluation/run_hierarchical_evaluation.py index 5dee74e6..b07fb23f 100644 --- a/src/evaluation/run_hierarchical_evaluation.py +++ b/src/evaluation/run_hierarchical_evaluation.py @@ -1,14 +1,18 @@ +import os +import sys import json import time from pathlib import Path -from PIL import Image +# Add workspace root to sys.path +sys.path.append("/workspace") -from dataset_preparation.utilities.datasets import DATASETS +from PIL import Image from metrics.metrics_api import compute_metrics +from dataset_preparation.utilities.datasets import DATASETS from src.evaluation.evaluator import EvaluationConfig, ModelEvaluator -from src.inference import MobileNetInference, ResNetInference -from src.optimal_class_mapping import map_prediction +from src.inference.utils.inference_factory import InferenceFactory +from src.optimal_class_mapping import MODEL_NAMES as class_mapping, map_prediction class MappedModelWrapper: @@ -66,7 +70,7 @@ def main(): # ResNet print("🔄 ResNet...") - resnet = ResNetInference(weights_path="models/pytorch/resnet18.pt", num_classes=83) + resnet = InferenceFactory.create("resnet", "models/pytorch/resnet18.pt", "cpu") wrapped_resnet = MappedModelWrapper(resnet) predictions, latencies = [], [] @@ -83,7 +87,7 @@ def main(): # MobileNet print("🔄 MobileNet...") - mobilenet = MobileNetInference(weights_path="models/pytorch/mobilenet.pt", num_classes=83) + mobilenet = InferenceFactory.create("mobilenet", "models/pytorch/mobilenet.pt", "cpu") wrapped_mobilenet = MappedModelWrapper(mobilenet) predictions, latencies = [], [] diff --git a/src/onnx_model.py b/src/onnx_model.py index f0599524..2e02bd2e 100644 --- a/src/onnx_model.py +++ b/src/onnx_model.py @@ -1,16 +1,18 @@ -# onnx_inference.py +import os +import sys from pathlib import Path from typing import Dict, Optional, Union -import numpy as np +# Add workspace root to sys.path +sys.path.append("/workspace") -# --- ONNX Runtime --- +import numpy as np import onnxruntime as ort import torch from torch import Tensor from torchvision import transforms -from src.inference import ClassifierInferenceBase +from src.inference.base.classifier_inference_base import ClassifierInferenceBase class OnnxClassifierInferenceBase(ClassifierInferenceBase): From 3873e3992dfed8f31ca07f72605021d5d9ac15e6 Mon Sep 17 00:00:00 2001 From: Sarah Gershuni Date: Tue, 11 Nov 2025 10:04:13 +0200 Subject: [PATCH 3/5] remove unreachable code --- src/inference/utils/inference_factory.py | 7 ------- 1 file changed, 7 deletions(-) diff --git a/src/inference/utils/inference_factory.py b/src/inference/utils/inference_factory.py index 37da19ed..4a5660ed 100644 --- a/src/inference/utils/inference_factory.py +++ b/src/inference/utils/inference_factory.py @@ -19,10 +19,3 @@ def create(model_type: str, model_path: str, device: str, class_mapping=None): return MobileNetInference(device=device, weights_path=model_path, class_mapping=class_mapping) else: raise ValueError(f"Unsupported model type: {model_type}") - - model_class = mapping[model_type] - return model_class( - device=device, - weights_path=model_path, - class_mapping=class_mapping, - ) From 9aac83731e5fbb9f7d3d14ef8fc5c43a01d60abc Mon Sep 17 00:00:00 2001 From: Sarah Gershuni Date: Tue, 11 Nov 2025 14:59:08 +0200 Subject: [PATCH 4/5] remove unnecessary shebang and fix import style --- scripts/export_to_onnx.py | 1 - scripts/onnx_validation.py | 1 - src/evaluation/evaluator.py | 4 ++-- 3 files changed, 2 insertions(+), 4 deletions(-) diff --git a/scripts/export_to_onnx.py b/scripts/export_to_onnx.py index 2c0e2fa5..6ae1fa98 100644 --- a/scripts/export_to_onnx.py +++ b/scripts/export_to_onnx.py @@ -1,4 +1,3 @@ -#!/usr/bin/env python3 import os import sys import argparse diff --git a/scripts/onnx_validation.py b/scripts/onnx_validation.py index 35f5dda8..385cc7a5 100644 --- a/scripts/onnx_validation.py +++ b/scripts/onnx_validation.py @@ -1,4 +1,3 @@ -#!/usr/bin/env python3 import os import argparse from pathlib import Path diff --git a/src/evaluation/evaluator.py b/src/evaluation/evaluator.py index 757c4a7e..314be930 100644 --- a/src/evaluation/evaluator.py +++ b/src/evaluation/evaluator.py @@ -11,7 +11,7 @@ import torch from PIL import Image -from src.inference.base.classifier_inference_base import ClassifierInferenceBase as InferenceModel +from src.inference.base.classifier_inference_base import ClassifierInferenceBase from metrics.metrics_api import ClassificationReport, compute_metrics @dataclass @@ -41,7 +41,7 @@ def load_dataset(self) -> Tuple[List[Path], List[int]]: print(f"✅ Loaded {len(image_paths)} images from {len(set(labels))} classes") return image_paths, labels - def evaluate_model(self, model: InferenceModel) -> ClassificationReport: + def evaluate_model(self, model: ClassifierInferenceBase) -> ClassificationReport: print("🔄 Starting evaluation...") image_paths, true_labels = self.load_dataset() predictions = [] From b2d63f473b9ce00fa42f29b879abc8b3c6f82d15 Mon Sep 17 00:00:00 2001 From: Sarah Gershuni Date: Tue, 11 Nov 2025 14:59:29 +0200 Subject: [PATCH 5/5] rename process_images_with_oop to run_inference --- src/inference/utils/run_inference_main.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/inference/utils/run_inference_main.py b/src/inference/utils/run_inference_main.py index 7145e056..f099bd3f 100644 --- a/src/inference/utils/run_inference_main.py +++ b/src/inference/utils/run_inference_main.py @@ -15,7 +15,7 @@ # Main image processing logic -def process_images_with_oop(model_type: str, model_path: str, image_dir: str, device: str): +def run_inference(model_type: str, model_path: str, image_dir: str, device: str): """Process images using OOP approach with Factory pattern.""" class_mapping = ID_TO_NAME if ID_TO_NAME else None @@ -76,7 +76,7 @@ def main(): device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Starting inference on device: {device}") - results, avg_time = process_images_with_oop( + results, avg_time = run_inference( args.model, args.model_path, args.image_dir,