Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 9 additions & 3 deletions scripts/export_to_onnx.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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()),
Expand Down
45 changes: 23 additions & 22 deletions scripts/onnx_validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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)
173 changes: 122 additions & 51 deletions src/evaluation/run_evaluation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
30 changes: 24 additions & 6 deletions src/evaluation/run_hierarchical_evaluation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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 = [], []
Expand All @@ -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 = [], []
Expand Down

Copy link
Copy Markdown
Collaborator

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 onnx directory and rename it to onnx_inference.py (similar to mobilenet_inference.py and resnet_inference.py).

Copy link
Copy Markdown
Collaborator Author

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 to onnx/ is out of scope for this ticket.

File renamed without changes.
63 changes: 49 additions & 14 deletions src/inference/utils/inference_factory.py
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:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I like the fact that now the InferenceFactory supports onnx - that makes sense. However, I think the implementation could be more transparent.

Currently, the main method is def create(model_type: str, model_path: str, device: str, class_mapping=None, **kwargs):, and we have many if/elif/else statements. In general, the best practice is to use enums for that.

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).

from enum import Enum
from pydantic import BaseModel, Field, validator
from typing import Optional, Type, Union
from pathlib import Path

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 ModelArch(str, Enum):
    RESNET = "resnet"
    MOBILENET = "mobilenet"


class Backend(str, Enum):
    PYTORCH = "pytorch"
    ONNX = "onnx"


class InferenceConfig(BaseModel):
    arch: ModelArch = Field(..., description="Model architecture (resnet or mobilenet).")
    backend: Backend = Field(..., description="Backend (pytorch or onnx).")
    model_path: Union[str, Path] = Field(..., description="Model checkpoint or ONNX path.")
    device: str = Field(..., description="Device string (e.g. 'cuda', 'cpu').")
    class_mapping: Optional[dict] = None
    extra_args: dict = Field(default_factory=dict)

    @validator("backend", pre=True, always=True)
    def infer_backend(cls, v, values):
        """If backend not given, infer ONNX from file extension."""
        if v:
            return v
        model_path = Path(values.get("model_path", ""))
        if model_path.suffix == ".onnx":
            return Backend.ONNX
        return Backend.PYTORCH


class InferenceFactory:
    """Factory to create model inference instances from architecture/backend pairs."""

    _PYTORCH_CLASSES : dict[ModelArch, Type] = {
        ModelArch.RESNET: ResNetInference,
        ModelArch.MOBILENET: MobileNetInference,
    }

    _ONNX_CLASSES : Type = OnnxClassifierInferenceBase

    @staticmethod
    def create(config: InferenceConfig):
        if config.backend == Backend.ONNX:
            model_cls = InferenceFactory._ONNX_IMPL
        elif config.backend == Backend.PYTORCH:
            model_cls = InferenceFactory._PYTORCH_IMPLS.get(config.arch)
            if model_cls is None:
                raise ValueError(f"Unsupported PyTorch architecture: {config.arch}")
        else:
            raise ValueError(f"Unsupported backend: {config.backend}")

        return model_cls(
            device=config.device,
            weights_path=config.model_path,
            class_mapping=config.class_mapping,
            **config.extra_args,
        )

_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,
)
Loading
Loading