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.
Files changed (32) hide show
  1. nodetool/mlx/__init__.py +5 -0
  2. nodetool/mlx/flux_model_loader.py +445 -0
  3. nodetool/mlx/mlx_provider.py +3015 -0
  4. nodetool/mlx/stable_audio_3/LICENSE +21 -0
  5. nodetool/mlx/stable_audio_3/NOTICE.md +23 -0
  6. nodetool/mlx/stable_audio_3/__init__.py +24 -0
  7. nodetool/mlx/stable_audio_3/defs/__init__.py +0 -0
  8. nodetool/mlx/stable_audio_3/defs/dit_mlx.py +344 -0
  9. nodetool/mlx/stable_audio_3/defs/dit_mlx_medium.py +458 -0
  10. nodetool/mlx/stable_audio_3/defs/sa3_pipeline.py +199 -0
  11. nodetool/mlx/stable_audio_3/defs/same_l_decoder.py +345 -0
  12. nodetool/mlx/stable_audio_3/defs/same_l_encoder.py +146 -0
  13. nodetool/mlx/stable_audio_3/defs/same_s_decoder.py +294 -0
  14. nodetool/mlx/stable_audio_3/defs/same_s_encoder.py +161 -0
  15. nodetool/mlx/stable_audio_3/defs/t5gemma_mlx.py +313 -0
  16. nodetool/mlx/stable_audio_3/pipeline.py +386 -0
  17. nodetool/mlx/stable_audio_3/weights.py +64 -0
  18. nodetool/nodes/mlx/_hf_cache.py +55 -0
  19. nodetool/nodes/mlx/automatic_speech_recognition.py +184 -0
  20. nodetool/nodes/mlx/image_to_image.py +3199 -0
  21. nodetool/nodes/mlx/image_to_text.py +219 -0
  22. nodetool/nodes/mlx/speech_enhancement.py +252 -0
  23. nodetool/nodes/mlx/speech_to_text.py +411 -0
  24. nodetool/nodes/mlx/text_generation.py +245 -0
  25. nodetool/nodes/mlx/text_to_audio.py +372 -0
  26. nodetool/nodes/mlx/text_to_image.py +336 -0
  27. nodetool/nodes/mlx/text_to_music.py +729 -0
  28. nodetool/nodes/mlx/text_to_speech.py +1464 -0
  29. nodetool/package_metadata/nodetool-mlx.json +10105 -0
  30. nodetool_mlx-0.7.0.dist-info/METADATA +219 -0
  31. nodetool_mlx-0.7.0.dist-info/RECORD +32 -0
  32. nodetool_mlx-0.7.0.dist-info/WHEEL +4 -0
@@ -0,0 +1,5 @@
1
+ """MLX integration for NodeTool."""
2
+
3
+ from nodetool.mlx.mlx_provider import MLXProvider
4
+
5
+ __all__ = ["MLXProvider"]
@@ -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