ai-gateway / core /loader.py
basyx's picture
Upload 58 files
36333c5 verified
Raw
History Blame Contribute Delete
5.39 kB
"""Singleton lazy model loader with exclusive residency."""
from __future__ import annotations
from contextlib import contextmanager
from importlib import import_module
from threading import RLock
from typing import Iterator
from config import Settings, get_settings
from core.cache import LRUCache
from core.errors import ModelLoadError, OutOfMemoryError
from core.runtime import cleanup_memory, detect_device
from models.base import ModelAdapter
MODEL_CLASSES: dict[str, tuple[str, str]] = {
"flux": ("models.flux", "FluxModel"),
"wan": ("models.wan", "WanModel"),
"kokoro": ("models.kokoro", "KokoroModel"),
"musicgen": ("models.musicgen", "MusicGenModel"),
"whisper": ("models.whisper", "WhisperModel"),
"sfx": ("models.sfx", "SFXModel"),
}
class ModelLoader:
"""Load models lazily while allowing only one resident model at a time."""
_instance: "ModelLoader | None" = None
_instance_lock = RLock()
def __new__(cls, settings: Settings | None = None) -> "ModelLoader":
with cls._instance_lock:
if cls._instance is None:
cls._instance = super().__new__(cls)
return cls._instance
def __init__(self, settings: Settings | None = None) -> None:
with self._instance_lock:
if getattr(self, "_initialized", False):
return
self.settings = settings or get_settings()
self.device = detect_device(self.settings.device)
self._adapters: LRUCache[str, ModelAdapter] = LRUCache(self.settings.cache_size)
self._resident_name: str | None = None
self._active_name: str | None = None
self._lock = RLock()
self._initialized = True
@classmethod
def reset_instance(cls) -> None:
"""Dispose the singleton; intended for tests and process teardown."""
with cls._instance_lock:
if cls._instance is not None:
cls._instance.close()
cls._instance = None
@property
def active_model(self) -> str | None:
return self._active_name
@property
def resident_model(self) -> str | None:
return self._resident_name
@property
def loaded_models(self) -> list[str]:
"""Return names of adapters that currently retain model weights."""
with self._lock:
return sorted(adapter.name for adapter in self._adapters.values() if adapter.loaded)
def _adapter(self, name: str) -> ModelAdapter:
if name not in MODEL_CLASSES:
raise ValueError(f"Unknown model: {name}")
adapter = self._adapters.get(name)
if adapter is not None:
return adapter
module_name, class_name = MODEL_CLASSES[name]
module = import_module(module_name)
adapter_type = getattr(module, class_name)
adapter = adapter_type(self.settings)
evicted = self._adapters.put(name, adapter)
if evicted is not None:
_, old_adapter = evicted
old_adapter.unload()
return adapter
def _prepare(self, name: str) -> ModelAdapter:
# Re-detect in case an accelerator becomes visible after process startup.
self.device = detect_device(self.settings.device)
if self._resident_name and self._resident_name != name:
previous = self._adapters.get(self._resident_name)
if previous is not None:
previous.unload()
self._resident_name = None
cleanup_memory()
adapter = self._adapter(name)
try:
adapter.ensure_loaded(self.device)
except Exception as exc:
adapter.unload()
if "out of memory" in str(exc).lower():
raise OutOfMemoryError(name) from exc
raise ModelLoadError(name, str(exc)) from exc
self._resident_name = name
return adapter
@contextmanager
def use_model(self, name: str) -> Iterator[ModelAdapter]:
"""Hold the lifecycle lock for one inference and release GPU memory after it."""
with self._lock:
adapter = self._prepare(name)
self._active_name = name
try:
yield adapter
finally:
try:
adapter.release_gpu()
finally:
self._active_name = None
cleanup_memory()
def load_flux(self) -> ModelAdapter:
with self._lock:
return self._prepare("flux")
def load_wan(self) -> ModelAdapter:
with self._lock:
return self._prepare("wan")
def load_musicgen(self) -> ModelAdapter:
with self._lock:
return self._prepare("musicgen")
def load_kokoro(self) -> ModelAdapter:
with self._lock:
return self._prepare("kokoro")
def load_whisper(self) -> ModelAdapter:
with self._lock:
return self._prepare("whisper")
def load_sfx(self) -> ModelAdapter:
with self._lock:
return self._prepare("sfx")
def close(self) -> None:
"""Fully unload every constructed adapter."""
with self._lock:
for adapter in self._adapters.clear():
adapter.unload()
self._resident_name = None
self._active_name = None
cleanup_memory()