diff --git a/scripts/export_to_onnx.py b/scripts/export_to_onnx.py index 6ae1fa98..3bba5677 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,18 @@ 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.lower() else ModelArch.MOBILENET + backend = Backend.PYTORCH + + config = InferenceConfig( + arch=arch, + backend=backend, 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 385cc7a5..6fe64c81 100644 --- a/scripts/onnx_validation.py +++ b/scripts/onnx_validation.py @@ -6,11 +6,10 @@ import torch from PIL import Image -from src.optimal_class_mapping import MODEL_NAMES as ID_TO_NAME -from src.inference.utils.inference_factory import InferenceFactory +from dataset.optimal_class_mapping import MODEL_NAMES as ID_TO_NAME +from src.inference.utils.inference_factory import InferenceFactory, InferenceConfig, Backend, ModelArch 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) @@ -78,32 +77,34 @@ def assert_same_predictions( def main(model_name: str): - torch_model = InferenceFactory.create( - model_type=model_name, + # 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=torch_arch, + backend=torch_backend, model_path=MODELS_DIR_PATH / "pytorch" / f"{model_name}.pt", device=DEVICE, class_mapping=ID_TO_NAME, ) - - onnx_model = OnnxClassifierInferenceBase( + + onnx_config = InferenceConfig( + arch=onnx_arch, + backend=onnx_backend, + 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, + 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) - - -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/evaluation/run_evaluation.py b/src/evaluation/run_evaluation.py index 983e07eb..86c25098 100644 --- a/src/evaluation/run_evaluation.py +++ b/src/evaluation/run_evaluation.py @@ -8,55 +8,126 @@ # Add workspace root to sys.path sys.path.append("/workspace") -import torch -from PIL import Image - -from src.evaluation.metrics.metrics_api import ClassificationReport, compute_metrics -from src.inference import ClassifierInferenceBase as InferenceModel - -@dataclass -class EvaluationConfig: - dataset_path: Path - model_weights_path: Optional[Path] = None - device: str = "cpu" - batch_size: int = 32 - num_workers: int = 4 - - -class ModelEvaluator: - def __init__(self, config: EvaluationConfig): - self.config = config - self.device = torch.device(config.device) - - def load_dataset(self) -> Tuple[List[Path], List[int]]: - image_paths = [] - labels = [] - dataset_path = self.config.dataset_path / "images" - for class_dir in sorted(dataset_path.iterdir()): - if class_dir.is_dir(): - class_id = int(class_dir.name) - for image_path in class_dir.glob("*.png"): - image_paths.append(image_path) - labels.append(class_id) - print(f"āœ… Loaded {len(image_paths)} images from {len(set(labels))} classes") - return image_paths, labels - - def evaluate_model(self, model: InferenceModel) -> ClassificationReport: - print("šŸ”„ Starting evaluation...") - image_paths, true_labels = self.load_dataset() - predictions = [] - latencies = [] - for i, image_path in enumerate(image_paths): - if i % 100 == 0: - print(f"Progress: {i}/{len(image_paths)}") - image = Image.open(image_path).convert("RGB") - start_time = time.perf_counter() - result = model.infer(image) - end_time = time.perf_counter() - latencies.append(end_time - start_time) - predictions.append(result[0]["class_id"]) - report = compute_metrics( - y_true=true_labels, y_pred=predictions, latencies=latencies +from src.optimal_class_mapping import MODEL_NAMES as class_mapping, map_prediction +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 + + +MODELS_DIR_PATH = Path("models") + +class MappedModelWrapper: + """Wrapper that adds mapping between 83 model classes to 76 dataset classes""" + + def __init__(self, model): + self.model = model + + def infer(self, image): + predictions = self.model.infer(image) + # Map the first prediction + mapped_class_id = map_prediction(predictions[0]["class_id"]) + return [ + { + "class_id": mapped_class_id, + "class_name": f"class_{mapped_class_id}", + "probability": predictions[0]["probability"], + } + ] + + +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": + 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": + config = InferenceConfig( + arch=None, + backend=Backend.ONNX, + model_path=weights_path.with_suffix(".onnx"), + device=device, + class_mapping=class_mapping, + extra_args={"topk": 1}, ) - print("āœ… Evaluation completed!") - return report + else: + raise ValueError(f"Unknown model_type: {model_type}") + + base_model = InferenceFactory.create(config) + + + return MappedModelWrapper(base_model) + + +def parse_arguments(): + parser = argparse.ArgumentParser( + description="Evaluate models on classification dataset" + ) + parser.add_argument("--dataset", type=str, required=True) + parser.add_argument("--models", type=str, default="resnet18,mobilenet") + parser.add_argument( + "--model_type", type=str, choices=["pytorch", "onnx"], required=True + ) + parser.add_argument("--device", type=str, default="cpu") + parser.add_argument("--output", type=str, default="outputs/evaluation_results.json") + return parser.parse_args() + + +def main(): + args = parse_arguments() + models_to_eval = [m.strip() for m in args.models.split(",")] + config = EvaluationConfig(dataset_path=Path(args.dataset), device=args.device) + evaluator = ModelEvaluator(config) + results = {} + print("šŸ“Š EVALUATION STARTING") + print(f"Dataset: {args.dataset}, Models: {models_to_eval}, Device: {args.device}") + print("=" * 50) + + for model_name in models_to_eval: + try: + model = create_model(model_name, args.model_type, args.device) + report = evaluator.evaluate_model(model) + result_key = f"{model_name}_{args.model_type}" + results[result_key] = report.dict() + print(f"\nāœ… {model_name.upper()} Results:") + print(f" Accuracy: {report.accuracy:.4f}") + print(f" Precision (Micro): {report.precision_micro:.4f}") + print(f" Recall (Micro): {report.recall_micro:.4f}") + print(f" F1-Score (Micro): {report.f1_micro:.4f}") + print(f" Latency: {report.latency_mean:.4f}s ± {report.latency_std:.4f}s") + except Exception as e: + print(f"āŒ Error evaluating {model_name}: {e}") + continue + + ensure_clean_directory(Path(args.output).parent) + with open(args.output, "a") as f: + json.dump(results, f, indent=2) + print(f"\nšŸ’¾ Results saved to: {args.output}") + + if results: + df_data = [] + for model_name, report in results.items(): + df_data.append( + { + "Model": model_name.upper(), + "Accuracy": f"{report['accuracy']:.4f}", + "Precision": f"{report['precision_micro']:.4f}", + "Recall": f"{report['recall_micro']:.4f}", + "F1-Score": f"{report['f1_micro']:.4f}", + "Latency (s)": f"{report['latency_mean']:.4f} ± {report['latency_std']:.4f}", + } + ) + df = pd.DataFrame(df_data) + print("\nšŸ“Š SUMMARY TABLE:") + print("=" * 80) + print(df.to_string(index=False)) + + +if __name__ == "__main__": + main() diff --git a/src/evaluation/run_hierarchical_evaluation.py b/src/evaluation/run_hierarchical_evaluation.py index 7d114f5e..7b5b7e87 100644 --- a/src/evaluation/run_hierarchical_evaluation.py +++ b/src/evaluation/run_hierarchical_evaluation.py @@ -9,10 +9,10 @@ from PIL import Image from dataset.utilities.datasets import DATASETS -from src.evaluation.metrics.metrics_api import compute_metrics -from src.evaluation.base.evaluator import EvaluationConfig, ModelEvaluator -from src.inference import MobileNetInference, ResNetInference -from dataset.optimal_class_mapping import map_prediction +from metrics.metrics_api import compute_metrics +from src.evaluation.evaluator import EvaluationConfig, ModelEvaluator +from src.inference.utils.inference_factory import InferenceFactory, InferenceConfig, Backend, ModelArch +from src.optimal_class_mapping import MODEL_NAMES as class_mapping, map_prediction class MappedModelWrapper: @@ -70,7 +70,16 @@ def main(): # ResNet print("šŸ”„ ResNet...") - resnet = ResNetInference(weights_path="models/pytorch/resnet18.pt", num_classes=83) + 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 = [], [] @@ -87,7 +96,16 @@ def main(): # MobileNet print("šŸ”„ MobileNet...") - mobilenet = MobileNetInference(weights_path="models/pytorch/mobilenet.pt", num_classes=83) + 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/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/src/inference/utils/inference_factory.py b/src/inference/utils/inference_factory.py index 4a5660ed..4c397f21 100644 --- a/src/inference/utils/inference_factory.py +++ b/src/inference/utils/inference_factory.py @@ -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: + _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, + ) diff --git a/scripts/onnx_predict_images.py b/src/inference/utils/onnx_predict_images.py similarity index 84% rename from scripts/onnx_predict_images.py rename to src/inference/utils/onnx_predict_images.py index 2745a140..9bda588d 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.utils.inference_factory import InferenceFactory, InferenceConfig, Backend IMAGE_SIZE = (256, 256) COLOR_MODE = "RGB" @@ -29,8 +29,13 @@ 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() + 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..7eb46a82 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,19 @@ 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) + 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=arch, + backend=backend, + model_path=Path(model_path), + device=device, + class_mapping=class_mapping, + ) + + + clf = InferenceFactory.create(config) results = {} total_time = 0