"""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()