Spaces:
Running on Zero
Running on Zero
File size: 5,386 Bytes
36333c5 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 | """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()
|