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