Spaces:
Running on Zero
Running on Zero
| """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 | |
| 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 | |
| def active_model(self) -> str | None: | |
| return self._active_name | |
| def resident_model(self) -> str | None: | |
| return self._resident_name | |
| 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 | |
| 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() | |