nodetool-mlx 0.7.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- nodetool/mlx/__init__.py +5 -0
- nodetool/mlx/flux_model_loader.py +445 -0
- nodetool/mlx/mlx_provider.py +3015 -0
- nodetool/mlx/stable_audio_3/LICENSE +21 -0
- nodetool/mlx/stable_audio_3/NOTICE.md +23 -0
- nodetool/mlx/stable_audio_3/__init__.py +24 -0
- nodetool/mlx/stable_audio_3/defs/__init__.py +0 -0
- nodetool/mlx/stable_audio_3/defs/dit_mlx.py +344 -0
- nodetool/mlx/stable_audio_3/defs/dit_mlx_medium.py +458 -0
- nodetool/mlx/stable_audio_3/defs/sa3_pipeline.py +199 -0
- nodetool/mlx/stable_audio_3/defs/same_l_decoder.py +345 -0
- nodetool/mlx/stable_audio_3/defs/same_l_encoder.py +146 -0
- nodetool/mlx/stable_audio_3/defs/same_s_decoder.py +294 -0
- nodetool/mlx/stable_audio_3/defs/same_s_encoder.py +161 -0
- nodetool/mlx/stable_audio_3/defs/t5gemma_mlx.py +313 -0
- nodetool/mlx/stable_audio_3/pipeline.py +386 -0
- nodetool/mlx/stable_audio_3/weights.py +64 -0
- nodetool/nodes/mlx/_hf_cache.py +55 -0
- nodetool/nodes/mlx/automatic_speech_recognition.py +184 -0
- nodetool/nodes/mlx/image_to_image.py +3199 -0
- nodetool/nodes/mlx/image_to_text.py +219 -0
- nodetool/nodes/mlx/speech_enhancement.py +252 -0
- nodetool/nodes/mlx/speech_to_text.py +411 -0
- nodetool/nodes/mlx/text_generation.py +245 -0
- nodetool/nodes/mlx/text_to_audio.py +372 -0
- nodetool/nodes/mlx/text_to_image.py +336 -0
- nodetool/nodes/mlx/text_to_music.py +729 -0
- nodetool/nodes/mlx/text_to_speech.py +1464 -0
- nodetool/package_metadata/nodetool-mlx.json +10105 -0
- nodetool_mlx-0.7.0.dist-info/METADATA +219 -0
- nodetool_mlx-0.7.0.dist-info/RECORD +32 -0
- nodetool_mlx-0.7.0.dist-info/WHEEL +4 -0
nodetool/mlx/__init__.py
ADDED
|
@@ -0,0 +1,445 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Centralized Flux model loading utilities for MLX.
|
|
3
|
+
|
|
4
|
+
This module provides a unified way to load Flux models across both nodes and image providers,
|
|
5
|
+
with proper caching and availability checks.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import asyncio
|
|
11
|
+
from typing import TYPE_CHECKING
|
|
12
|
+
|
|
13
|
+
from nodetool.config.logging_config import get_logger
|
|
14
|
+
from nodetool.integrations.huggingface.hf_cache import has_cached_files
|
|
15
|
+
from nodetool.ml.core.model_manager import ModelManager
|
|
16
|
+
|
|
17
|
+
if TYPE_CHECKING: # pragma: no cover - import only for type checking
|
|
18
|
+
from mflux.models.flux.variants.controlnet.flux_controlnet import Flux1Controlnet
|
|
19
|
+
from mflux.models.flux.variants.depth.flux_depth import Flux1Depth
|
|
20
|
+
from mflux.models.flux.variants.fill.flux_fill import Flux1Fill
|
|
21
|
+
from mflux.models.flux.variants.kontext.flux_kontext import Flux1Kontext
|
|
22
|
+
from mflux.models.flux.variants.redux.flux_redux import Flux1Redux
|
|
23
|
+
from mflux.models.flux.variants.txt2img.flux import Flux1
|
|
24
|
+
|
|
25
|
+
log = get_logger(__name__)
|
|
26
|
+
|
|
27
|
+
# Estimated memory requirements in GB for different quantization levels
|
|
28
|
+
# These are rough estimates for Flux.1 models (approx 12B params)
|
|
29
|
+
# 4-bit: ~12GB (model) + overhead
|
|
30
|
+
# 8-bit: ~20GB (model) + overhead
|
|
31
|
+
# 16-bit: ~34GB (model) + overhead
|
|
32
|
+
FLUX_MEMORY_ESTIMATES = {
|
|
33
|
+
4: 12.0,
|
|
34
|
+
8: 20.0,
|
|
35
|
+
None: 34.0, # Default/16-bit
|
|
36
|
+
}
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def estimate_required_memory(
|
|
40
|
+
quantize: int | None, is_controlnet: bool = False
|
|
41
|
+
) -> float:
|
|
42
|
+
"""Estimate the required memory in GB for loading the model."""
|
|
43
|
+
base_mem = FLUX_MEMORY_ESTIMATES.get(quantize, 34.0)
|
|
44
|
+
|
|
45
|
+
# If quantization is not one of the standard values, estimate linearly
|
|
46
|
+
if quantize not in FLUX_MEMORY_ESTIMATES and quantize is not None:
|
|
47
|
+
# Rough linear interpolation: ~1.5GB per bit for 12B params + overhead
|
|
48
|
+
base_mem = quantize * 2.5
|
|
49
|
+
|
|
50
|
+
if is_controlnet:
|
|
51
|
+
# ControlNet adds extra parameters (approx 2-3GB for Flux ControlNet)
|
|
52
|
+
base_mem += 3.0
|
|
53
|
+
|
|
54
|
+
return base_mem
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def check_memory_availability(required_gb: float):
|
|
58
|
+
"""
|
|
59
|
+
Check if there is enough available RAM to load the model.
|
|
60
|
+
Raises MemoryError if insufficient memory.
|
|
61
|
+
"""
|
|
62
|
+
import psutil
|
|
63
|
+
|
|
64
|
+
vm = psutil.virtual_memory()
|
|
65
|
+
|
|
66
|
+
# Calculate available memory in GB
|
|
67
|
+
# available includes memory that can be reclaimed (cache/buffers)
|
|
68
|
+
available_gb = vm.available / (1024**3)
|
|
69
|
+
total_gb = vm.total / (1024**3)
|
|
70
|
+
|
|
71
|
+
# 5% safety buffer
|
|
72
|
+
buffer_gb = total_gb * 0.05
|
|
73
|
+
|
|
74
|
+
if available_gb < (required_gb + buffer_gb):
|
|
75
|
+
raise MemoryError(
|
|
76
|
+
f"Insufficient memory to load Flux model.\n"
|
|
77
|
+
f"Required: {required_gb:.1f} GB\n"
|
|
78
|
+
f"Available: {available_gb:.1f} GB\n"
|
|
79
|
+
f"Safety Buffer: {buffer_gb:.1f} GB (5% of total)\n"
|
|
80
|
+
f"Total System RAM: {total_gb:.1f} GB\n\n"
|
|
81
|
+
f"Please close other applications or try a higher quantization level (e.g., 4-bit)."
|
|
82
|
+
)
|
|
83
|
+
|
|
84
|
+
log.info(
|
|
85
|
+
f"Memory check passed: {available_gb:.1f} GB available, {required_gb:.1f} GB required"
|
|
86
|
+
)
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
class FluxModelNotAvailableError(Exception):
|
|
90
|
+
"""Raised when a Flux model is not available in the local cache."""
|
|
91
|
+
|
|
92
|
+
def __init__(self, model_id: str):
|
|
93
|
+
self.model_id = model_id
|
|
94
|
+
super().__init__(
|
|
95
|
+
f"Model '{model_id}' is not available in your local cache.\n\n"
|
|
96
|
+
f"To use this model:\n"
|
|
97
|
+
f"1. Open the Model Manager (in the NodeTool UI)\n"
|
|
98
|
+
f"2. Search for '{model_id}'\n"
|
|
99
|
+
f"3. Download the model\n"
|
|
100
|
+
f"4. Try again once the download is complete\n\n"
|
|
101
|
+
f"Note: MLX models are only loaded from local cache and will not be downloaded automatically "
|
|
102
|
+
f"to avoid unexpected network usage and storage consumption."
|
|
103
|
+
)
|
|
104
|
+
|
|
105
|
+
|
|
106
|
+
def is_flux_model_available(model_id: str) -> bool:
|
|
107
|
+
"""
|
|
108
|
+
Check if a Flux model is available in the local HuggingFace cache.
|
|
109
|
+
|
|
110
|
+
Args:
|
|
111
|
+
model_id: The HuggingFace repo ID (e.g., "black-forest-labs/FLUX.1-schnell")
|
|
112
|
+
|
|
113
|
+
Returns:
|
|
114
|
+
True if the model is cached locally, False otherwise
|
|
115
|
+
"""
|
|
116
|
+
return has_cached_files(model_id)
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
async def load_flux_model(
|
|
120
|
+
model_id: str,
|
|
121
|
+
quantize: int | None = 4,
|
|
122
|
+
node_id: str | None = None,
|
|
123
|
+
task: str = "flux",
|
|
124
|
+
force_reload: bool = False,
|
|
125
|
+
) -> "Flux1":
|
|
126
|
+
"""
|
|
127
|
+
Load a Flux model with proper caching and availability checks.
|
|
128
|
+
|
|
129
|
+
This function:
|
|
130
|
+
1. Checks if the model is available in the HuggingFace cache
|
|
131
|
+
2. Checks the ModelManager cache
|
|
132
|
+
3. Only loads from local files (does not download)
|
|
133
|
+
4. Caches the loaded model for reuse
|
|
134
|
+
|
|
135
|
+
Args:
|
|
136
|
+
model_id: The HuggingFace repo ID (e.g., "black-forest-labs/FLUX.1-schnell")
|
|
137
|
+
quantize: Quantization level (3, 4, 5, 6, 8, or None)
|
|
138
|
+
node_id: Optional node ID for ModelManager association
|
|
139
|
+
task: Task identifier for ModelManager (default: "flux")
|
|
140
|
+
force_reload: If True, bypass ModelManager cache and reload from disk
|
|
141
|
+
|
|
142
|
+
Returns:
|
|
143
|
+
The loaded Flux1 model instance
|
|
144
|
+
|
|
145
|
+
Raises:
|
|
146
|
+
RuntimeError: If mflux is not available or not on macOS
|
|
147
|
+
FluxModelNotAvailableError: If the model is not in the local cache
|
|
148
|
+
"""
|
|
149
|
+
# Check if model is available in HF cache
|
|
150
|
+
if not is_flux_model_available(model_id):
|
|
151
|
+
raise FluxModelNotAvailableError(model_id)
|
|
152
|
+
|
|
153
|
+
# Construct cache key
|
|
154
|
+
cache_key = f"{model_id}_{task}_q{quantize}"
|
|
155
|
+
|
|
156
|
+
# Check ModelManager cache (unless forcing reload)
|
|
157
|
+
if not force_reload:
|
|
158
|
+
cached_model = ModelManager.get_model(cache_key)
|
|
159
|
+
if cached_model is not None:
|
|
160
|
+
log.info(f"Using cached Flux model: {model_id}")
|
|
161
|
+
return cached_model
|
|
162
|
+
|
|
163
|
+
# Load model from local cache
|
|
164
|
+
required_mem = estimate_required_memory(quantize)
|
|
165
|
+
check_memory_availability(required_mem)
|
|
166
|
+
|
|
167
|
+
loop = asyncio.get_running_loop()
|
|
168
|
+
|
|
169
|
+
def _load() -> "Flux1":
|
|
170
|
+
log.info(
|
|
171
|
+
f"Loading Flux model {model_id} from local cache "
|
|
172
|
+
f"(quantize={quantize if quantize is not None else 'none'})"
|
|
173
|
+
)
|
|
174
|
+
from mflux.models.flux.variants.txt2img.flux import Flux1
|
|
175
|
+
|
|
176
|
+
model = Flux1.from_name(
|
|
177
|
+
model_name=model_id,
|
|
178
|
+
quantize=quantize,
|
|
179
|
+
)
|
|
180
|
+
return model
|
|
181
|
+
|
|
182
|
+
model = await loop.run_in_executor(None, _load)
|
|
183
|
+
|
|
184
|
+
# Cache in ModelManager if node_id provided
|
|
185
|
+
if node_id:
|
|
186
|
+
ModelManager.set_model(node_id, cache_key, model)
|
|
187
|
+
|
|
188
|
+
return model
|
|
189
|
+
|
|
190
|
+
|
|
191
|
+
async def load_flux_controlnet_model(
|
|
192
|
+
base_model_id: str,
|
|
193
|
+
controlnet_model_id: str,
|
|
194
|
+
quantize: int | None = 4,
|
|
195
|
+
node_id: str | None = None,
|
|
196
|
+
force_reload: bool = False,
|
|
197
|
+
) -> "Flux1Controlnet":
|
|
198
|
+
"""
|
|
199
|
+
Load a Flux ControlNet model with proper caching and availability checks.
|
|
200
|
+
|
|
201
|
+
Args:
|
|
202
|
+
base_model_id: Base Flux model repo ID
|
|
203
|
+
controlnet_model_id: ControlNet model repo ID
|
|
204
|
+
quantize: Quantization level
|
|
205
|
+
node_id: Optional node ID for ModelManager association
|
|
206
|
+
force_reload: If True, bypass ModelManager cache
|
|
207
|
+
|
|
208
|
+
Returns:
|
|
209
|
+
The loaded Flux1Controlnet model instance
|
|
210
|
+
|
|
211
|
+
Raises:
|
|
212
|
+
RuntimeError: If mflux is not available or not on macOS
|
|
213
|
+
FluxModelNotAvailableError: If either model is not in the local cache
|
|
214
|
+
"""
|
|
215
|
+
# Check if both models are available
|
|
216
|
+
if not is_flux_model_available(base_model_id):
|
|
217
|
+
raise FluxModelNotAvailableError(base_model_id)
|
|
218
|
+
if not is_flux_model_available(controlnet_model_id):
|
|
219
|
+
raise FluxModelNotAvailableError(controlnet_model_id)
|
|
220
|
+
|
|
221
|
+
# Construct cache key
|
|
222
|
+
cache_key = f"{base_model_id}:{controlnet_model_id}_flux-controlnet_q{quantize}"
|
|
223
|
+
|
|
224
|
+
# Check ModelManager cache
|
|
225
|
+
if not force_reload:
|
|
226
|
+
cached_model = ModelManager.get_model(cache_key)
|
|
227
|
+
if cached_model is not None:
|
|
228
|
+
log.info(f"Using cached Flux ControlNet model: {cache_key}")
|
|
229
|
+
return cached_model
|
|
230
|
+
|
|
231
|
+
required_mem = estimate_required_memory(quantize, is_controlnet=True)
|
|
232
|
+
check_memory_availability(required_mem)
|
|
233
|
+
|
|
234
|
+
loop = asyncio.get_running_loop()
|
|
235
|
+
|
|
236
|
+
def _load() -> "Flux1Controlnet":
|
|
237
|
+
log.info(
|
|
238
|
+
f"Loading Flux ControlNet model {base_model_id} with controlnet {controlnet_model_id} "
|
|
239
|
+
f"(quantize={quantize if quantize is not None else 'none'})"
|
|
240
|
+
)
|
|
241
|
+
from mflux.models.common.config import ModelConfig
|
|
242
|
+
from mflux.models.flux.variants.controlnet.flux_controlnet import (
|
|
243
|
+
Flux1Controlnet,
|
|
244
|
+
)
|
|
245
|
+
|
|
246
|
+
model_config = ModelConfig.from_name(base_model_id)
|
|
247
|
+
model_config.controlnet_model = controlnet_model_id
|
|
248
|
+
|
|
249
|
+
model = Flux1Controlnet(
|
|
250
|
+
model_config=model_config,
|
|
251
|
+
quantize=quantize,
|
|
252
|
+
)
|
|
253
|
+
return model
|
|
254
|
+
|
|
255
|
+
model = await loop.run_in_executor(None, _load)
|
|
256
|
+
|
|
257
|
+
# Cache in ModelManager
|
|
258
|
+
if node_id:
|
|
259
|
+
ModelManager.set_model(node_id, cache_key, model)
|
|
260
|
+
|
|
261
|
+
return model
|
|
262
|
+
|
|
263
|
+
|
|
264
|
+
async def load_flux_fill_model(
|
|
265
|
+
model_id: str = "black-forest-labs/FLUX.1-Fill-dev",
|
|
266
|
+
quantize: int | None = 4,
|
|
267
|
+
node_id: str | None = None,
|
|
268
|
+
force_reload: bool = False,
|
|
269
|
+
) -> "Flux1Fill":
|
|
270
|
+
"""
|
|
271
|
+
Load a Flux Fill (inpainting/outpainting) model.
|
|
272
|
+
|
|
273
|
+
Args:
|
|
274
|
+
model_id: Fill model repo ID
|
|
275
|
+
quantize: Quantization level
|
|
276
|
+
node_id: Optional node ID for ModelManager association
|
|
277
|
+
force_reload: If True, bypass ModelManager cache
|
|
278
|
+
|
|
279
|
+
Returns:
|
|
280
|
+
The loaded Flux1Fill model instance
|
|
281
|
+
"""
|
|
282
|
+
if not is_flux_model_available(model_id):
|
|
283
|
+
raise FluxModelNotAvailableError(model_id)
|
|
284
|
+
|
|
285
|
+
# Construct cache key
|
|
286
|
+
cache_key = f"{model_id}_flux-fill_q{quantize}"
|
|
287
|
+
|
|
288
|
+
if not force_reload:
|
|
289
|
+
cached_model = ModelManager.get_model(cache_key)
|
|
290
|
+
if cached_model is not None:
|
|
291
|
+
log.info(f"Using cached Flux Fill model: {model_id}")
|
|
292
|
+
return cached_model
|
|
293
|
+
|
|
294
|
+
required_mem = estimate_required_memory(quantize)
|
|
295
|
+
check_memory_availability(required_mem)
|
|
296
|
+
|
|
297
|
+
loop = asyncio.get_running_loop()
|
|
298
|
+
|
|
299
|
+
def _load() -> "Flux1Fill":
|
|
300
|
+
log.info(
|
|
301
|
+
f"Loading Flux Fill model {model_id} "
|
|
302
|
+
f"(quantize={quantize if quantize is not None else 'none'})"
|
|
303
|
+
)
|
|
304
|
+
from mflux.models.flux.variants.fill.flux_fill import Flux1Fill
|
|
305
|
+
|
|
306
|
+
model = Flux1Fill(quantize=quantize)
|
|
307
|
+
return model
|
|
308
|
+
|
|
309
|
+
model = await loop.run_in_executor(None, _load)
|
|
310
|
+
|
|
311
|
+
if node_id:
|
|
312
|
+
ModelManager.set_model(node_id, cache_key, model)
|
|
313
|
+
|
|
314
|
+
return model
|
|
315
|
+
|
|
316
|
+
|
|
317
|
+
async def load_flux_depth_model(
|
|
318
|
+
model_id: str = "black-forest-labs/FLUX.1-Depth-dev",
|
|
319
|
+
quantize: int | None = 4,
|
|
320
|
+
node_id: str | None = None,
|
|
321
|
+
force_reload: bool = False,
|
|
322
|
+
) -> "Flux1Depth":
|
|
323
|
+
"""Load a Flux Depth model."""
|
|
324
|
+
if not is_flux_model_available(model_id):
|
|
325
|
+
raise FluxModelNotAvailableError(model_id)
|
|
326
|
+
|
|
327
|
+
# Construct cache key
|
|
328
|
+
cache_key = f"{model_id}_flux-depth_q{quantize}"
|
|
329
|
+
|
|
330
|
+
if not force_reload:
|
|
331
|
+
cached_model = ModelManager.get_model(cache_key)
|
|
332
|
+
if cached_model is not None:
|
|
333
|
+
log.info(f"Using cached Flux Depth model: {model_id}")
|
|
334
|
+
return cached_model
|
|
335
|
+
|
|
336
|
+
required_mem = estimate_required_memory(quantize)
|
|
337
|
+
check_memory_availability(required_mem)
|
|
338
|
+
|
|
339
|
+
loop = asyncio.get_running_loop()
|
|
340
|
+
|
|
341
|
+
def _load() -> "Flux1Depth":
|
|
342
|
+
log.info(
|
|
343
|
+
f"Loading Flux Depth model {model_id} "
|
|
344
|
+
f"(quantize={quantize if quantize is not None else 'none'})"
|
|
345
|
+
)
|
|
346
|
+
from mflux.models.flux.variants.depth.flux_depth import Flux1Depth
|
|
347
|
+
|
|
348
|
+
model = Flux1Depth(quantize=quantize)
|
|
349
|
+
return model
|
|
350
|
+
|
|
351
|
+
model = await loop.run_in_executor(None, _load)
|
|
352
|
+
|
|
353
|
+
if node_id:
|
|
354
|
+
ModelManager.set_model(node_id, cache_key, model)
|
|
355
|
+
|
|
356
|
+
return model
|
|
357
|
+
|
|
358
|
+
|
|
359
|
+
async def load_flux_redux_model(
|
|
360
|
+
model_id: str = "black-forest-labs/FLUX.1-Redux-dev",
|
|
361
|
+
quantize: int | None = 4,
|
|
362
|
+
node_id: str | None = None,
|
|
363
|
+
force_reload: bool = False,
|
|
364
|
+
) -> "Flux1Redux":
|
|
365
|
+
"""Load a Flux Redux model."""
|
|
366
|
+
if not is_flux_model_available(model_id):
|
|
367
|
+
raise FluxModelNotAvailableError(model_id)
|
|
368
|
+
|
|
369
|
+
# Construct cache key
|
|
370
|
+
cache_key = f"{model_id}_flux-redux_q{quantize}"
|
|
371
|
+
|
|
372
|
+
if not force_reload:
|
|
373
|
+
cached_model = ModelManager.get_model(cache_key)
|
|
374
|
+
if cached_model is not None:
|
|
375
|
+
log.info(f"Using cached Flux Redux model: {model_id}")
|
|
376
|
+
return cached_model
|
|
377
|
+
|
|
378
|
+
required_mem = estimate_required_memory(quantize)
|
|
379
|
+
check_memory_availability(required_mem)
|
|
380
|
+
|
|
381
|
+
loop = asyncio.get_running_loop()
|
|
382
|
+
|
|
383
|
+
def _load() -> "Flux1Redux":
|
|
384
|
+
log.info(
|
|
385
|
+
f"Loading Flux Redux model {model_id} "
|
|
386
|
+
f"(quantize={quantize if quantize is not None else 'none'})"
|
|
387
|
+
)
|
|
388
|
+
from mflux.models.common.config import ModelConfig
|
|
389
|
+
from mflux.models.flux.variants.redux.flux_redux import Flux1Redux
|
|
390
|
+
|
|
391
|
+
model_config = ModelConfig.dev_redux()
|
|
392
|
+
model = Flux1Redux(
|
|
393
|
+
model_config=model_config,
|
|
394
|
+
quantize=quantize,
|
|
395
|
+
)
|
|
396
|
+
return model
|
|
397
|
+
|
|
398
|
+
model = await loop.run_in_executor(None, _load)
|
|
399
|
+
|
|
400
|
+
if node_id:
|
|
401
|
+
ModelManager.set_model(node_id, cache_key, model)
|
|
402
|
+
|
|
403
|
+
return model
|
|
404
|
+
|
|
405
|
+
|
|
406
|
+
async def load_flux_kontext_model(
|
|
407
|
+
model_id: str = "black-forest-labs/FLUX.1-Kontext-dev",
|
|
408
|
+
quantize: int | None = 4,
|
|
409
|
+
node_id: str | None = None,
|
|
410
|
+
force_reload: bool = False,
|
|
411
|
+
) -> "Flux1Kontext":
|
|
412
|
+
"""Load a Flux Kontext model."""
|
|
413
|
+
if not is_flux_model_available(model_id):
|
|
414
|
+
raise FluxModelNotAvailableError(model_id)
|
|
415
|
+
|
|
416
|
+
# Construct cache key
|
|
417
|
+
cache_key = f"{model_id}_flux-kontext_q{quantize}"
|
|
418
|
+
|
|
419
|
+
if not force_reload:
|
|
420
|
+
cached_model = ModelManager.get_model(cache_key)
|
|
421
|
+
if cached_model is not None:
|
|
422
|
+
log.info(f"Using cached Flux Kontext model: {model_id}")
|
|
423
|
+
return cached_model
|
|
424
|
+
|
|
425
|
+
required_mem = estimate_required_memory(quantize)
|
|
426
|
+
check_memory_availability(required_mem)
|
|
427
|
+
|
|
428
|
+
loop = asyncio.get_running_loop()
|
|
429
|
+
|
|
430
|
+
def _load() -> "Flux1Kontext":
|
|
431
|
+
log.info(
|
|
432
|
+
f"Loading Flux Kontext model {model_id} "
|
|
433
|
+
f"(quantize={quantize if quantize is not None else 'none'})"
|
|
434
|
+
)
|
|
435
|
+
from mflux.models.flux.variants.kontext.flux_kontext import Flux1Kontext
|
|
436
|
+
|
|
437
|
+
model = Flux1Kontext(quantize=quantize)
|
|
438
|
+
return model
|
|
439
|
+
|
|
440
|
+
model = await loop.run_in_executor(None, _load)
|
|
441
|
+
|
|
442
|
+
if node_id:
|
|
443
|
+
ModelManager.set_model(node_id, cache_key, model)
|
|
444
|
+
|
|
445
|
+
return model
|