46 lines
1.2 KiB
Python
46 lines
1.2 KiB
Python
"""Base Model Abstractions and Registry."""
|
|
|
|
from abc import ABC, abstractmethod
|
|
from typing import Any, Dict, Optional
|
|
|
|
|
|
class BaseModelWrapper(ABC):
|
|
"""Abstract base class for all AI model wrappers (ONNX, PyTorch, Custom)."""
|
|
|
|
def __init__(self, model_name: str, weights_path: Optional[str] = None):
|
|
self.model_name = model_name
|
|
self.weights_path = weights_path
|
|
self._is_loaded = False
|
|
|
|
@abstractmethod
|
|
def load(self) -> bool:
|
|
"""Load model weights into memory."""
|
|
pass
|
|
|
|
@abstractmethod
|
|
def predict(self, inputs: Any) -> Any:
|
|
"""Execute model inference."""
|
|
pass
|
|
|
|
@property
|
|
def is_loaded(self) -> bool:
|
|
return self._is_loaded
|
|
|
|
|
|
class ModelRegistry:
|
|
"""Central registry holding loaded model instances."""
|
|
|
|
_models: Dict[str, BaseModelWrapper] = {}
|
|
|
|
@classmethod
|
|
def register(cls, name: str, model: BaseModelWrapper) -> None:
|
|
cls._models[name] = model
|
|
|
|
@classmethod
|
|
def get(cls, name: str) -> Optional[BaseModelWrapper]:
|
|
return cls._models.get(name)
|
|
|
|
@classmethod
|
|
def is_registered(cls, name: str) -> bool:
|
|
return name in cls._models
|