mlx-dfloat 0.1.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 (41) hide show
  1. mlx_dfloat/__init__.py +25 -0
  2. mlx_dfloat/_memory_caps.py +74 -0
  3. mlx_dfloat/_metal_decode.py +420 -0
  4. mlx_dfloat/_safetensors.py +185 -0
  5. mlx_dfloat/_scrub.py +29 -0
  6. mlx_dfloat/_version.py +24 -0
  7. mlx_dfloat/_watchdog.py +253 -0
  8. mlx_dfloat/bench/__init__.py +4 -0
  9. mlx_dfloat/bench/capped.py +165 -0
  10. mlx_dfloat/bench/preflight.py +216 -0
  11. mlx_dfloat/bench/results.py +227 -0
  12. mlx_dfloat/bench/scenario.py +191 -0
  13. mlx_dfloat/bench/table.py +251 -0
  14. mlx_dfloat/cli.py +30 -0
  15. mlx_dfloat/decode.py +120 -0
  16. mlx_dfloat/errors.py +45 -0
  17. mlx_dfloat/format.py +462 -0
  18. mlx_dfloat/integrate/__init__.py +1 -0
  19. mlx_dfloat/integrate/coverage.py +122 -0
  20. mlx_dfloat/integrate/memory.py +55 -0
  21. mlx_dfloat/integrate/names.py +121 -0
  22. mlx_dfloat/integrate/placeholders.py +72 -0
  23. mlx_dfloat/integrate/providers.py +271 -0
  24. mlx_dfloat/integrate/seam.py +196 -0
  25. mlx_dfloat/mflux/__init__.py +31 -0
  26. mlx_dfloat/mflux/flux1/__init__.py +1 -0
  27. mlx_dfloat/mflux/flux1/cli.py +419 -0
  28. mlx_dfloat/mflux/flux1/init.py +245 -0
  29. mlx_dfloat/mflux/flux1/lifecycle.py +159 -0
  30. mlx_dfloat/mflux/flux1/memory.py +201 -0
  31. mlx_dfloat/mflux/flux1/model.py +553 -0
  32. mlx_dfloat/mflux/flux1/names.py +79 -0
  33. mlx_dfloat/mflux/flux1/transformer.py +240 -0
  34. mlx_dfloat/py.typed +0 -0
  35. mlx_dfloat/reference.py +241 -0
  36. mlx_dfloat-0.1.0.dist-info/METADATA +325 -0
  37. mlx_dfloat-0.1.0.dist-info/RECORD +41 -0
  38. mlx_dfloat-0.1.0.dist-info/WHEEL +4 -0
  39. mlx_dfloat-0.1.0.dist-info/entry_points.txt +2 -0
  40. mlx_dfloat-0.1.0.dist-info/licenses/LICENSE +202 -0
  41. mlx_dfloat-0.1.0.dist-info/licenses/NOTICE +40 -0
@@ -0,0 +1,553 @@
1
+ """``DFloatFlux1``: mflux's FLUX.1 pipeline over a DFloat11 transformer, decoded one block at a time.
2
+
3
+ The class subclasses mflux's ``Flux1`` so its loop, scheduler, callbacks, VAE decode and image
4
+ metadata run unchanged; only construction and the prelude of ``generate_image`` differ. The text
5
+ encoders and the compressed transformer are never resident together (see ``lifecycle``). The
6
+ Python API installs no memory caps and no watchdog; the ``mlx-dfloat`` command does both.
7
+ """
8
+
9
+ import logging
10
+ import traceback
11
+ from dataclasses import asdict, dataclass
12
+ from importlib import metadata
13
+ from pathlib import Path
14
+ from typing import Any
15
+
16
+ import mlx.core as mx
17
+ import mlx.nn as nn
18
+ from mlx.utils import tree_flatten
19
+
20
+ from mlx_dfloat._watchdog import phys_footprint
21
+ from mlx_dfloat.errors import DFloatFormatError, DFloatResourceError, DFloatUnsupportedError
22
+ from mlx_dfloat.format import DF11Checkpoint, MxGroup, open_checkpoint
23
+ from mlx_dfloat.integrate.coverage import load_resident_set
24
+ from mlx_dfloat.integrate.memory import FitEstimate, budget_bytes, largest_decoded_bytes
25
+ from mlx_dfloat.integrate.names import NameMap, Shapes
26
+ from mlx_dfloat.integrate.providers import DF11Provider
27
+ from mlx_dfloat.mflux import require_mflux
28
+ from mlx_dfloat.mflux.flux1 import init as base_init
29
+ from mlx_dfloat.mflux.flux1.init import BaseComponents, ResolvedRepo
30
+ from mlx_dfloat.mflux.flux1.lifecycle import Lifecycle
31
+ from mlx_dfloat.mflux.flux1.memory import (
32
+ MAX_MEASURED_PIXELS,
33
+ FluxSizes,
34
+ activation_allowance,
35
+ cache_limit_for,
36
+ fit_for,
37
+ sizes_for,
38
+ text_tokens,
39
+ )
40
+ from mlx_dfloat.mflux.flux1.names import (
41
+ DOUBLE_PREFIX,
42
+ SINGLE_PREFIX,
43
+ check_flux_groups,
44
+ flux_name_map,
45
+ )
46
+ from mlx_dfloat.mflux.flux1.transformer import build_transformer
47
+
48
+ require_mflux()
49
+ from mflux.models.flux.variants.txt2img.flux import Flux1 # noqa: E402
50
+
51
+ log = logging.getLogger("mlx_dfloat.mflux.flux1")
52
+
53
+ MODELS: dict[str, tuple[str, str]] = {
54
+ "schnell": ("DFloat11/FLUX.1-schnell-DF11", "black-forest-labs/FLUX.1-schnell"),
55
+ "dev": ("DFloat11/FLUX.1-dev-DF11", "black-forest-labs/FLUX.1-dev"),
56
+ "krea-dev": ("DFloat11/FLUX.1-Krea-dev-DF11", "black-forest-labs/FLUX.1-Krea-dev"),
57
+ }
58
+ POLICIES: tuple[str, ...] = ("per-block", "depth2")
59
+ RETAINED_SLACK_BYTES = 2 * 1024**3
60
+
61
+
62
+ @dataclass(frozen=True, slots=True, kw_only=True)
63
+ class CallPlan:
64
+ """What one generate call will do about memory: the cache limit, the estimate, and whether the set is dropped before the VAE decode."""
65
+
66
+ cache_limit: int
67
+ estimate: FitEstimate
68
+ drop_set_before_vae: bool
69
+
70
+
71
+ def _refuse(name: str, reason: str) -> None:
72
+ raise DFloatUnsupportedError(f"{name}: {reason}; not on the DFloat11 path in this version")
73
+
74
+
75
+ def _check_eval_policy(eval_policy: str) -> None:
76
+ """Refuse an eval policy outside ``POLICIES``.
77
+
78
+ Checked both in ``__init__`` and in ``_assemble``, so a ``_from_parts`` build gets the same
79
+ guard as the public constructor.
80
+ """
81
+ if eval_policy not in POLICIES:
82
+ raise DFloatUnsupportedError(f"eval_policy {eval_policy!r}: choose from {POLICIES}")
83
+
84
+
85
+ class _VaePoolGuard:
86
+ """mflux after-loop subscriber: close the denoise phase, clear and cap the buffer pool before the VAE decode."""
87
+
88
+ def __init__(self, model: "DFloatFlux1") -> None:
89
+ self._model = model
90
+
91
+ def call_after_loop(self, seed: int, prompt: str, latents: mx.array, config: Any) -> None:
92
+ """Runs once per generation, after the last denoise step and before the VAE decode.
93
+
94
+ Drops the compressed set first when the call was planned that way (the resident set plus the
95
+ float32 decode would exceed the budget; measured 23.29 GiB at 1024² on a 32 GB Mac).
96
+ """
97
+ del seed, prompt, latents, config
98
+ self._model._phase_end("denoise")
99
+ if self._model._plan is not None and self._model._plan.drop_set_before_vae:
100
+ self._model._lifecycle.drop_set()
101
+ mx.clear_cache()
102
+ mx.set_cache_limit(0) # the transformer-shaped buffers cannot serve the decoder
103
+ self._model._phase_begin("vae")
104
+
105
+
106
+ class DFloatFlux1(Flux1): # type: ignore[misc] # mflux ships no type information
107
+ """FLUX.1 (schnell, dev, Krea-dev) generation from a DFloat11 transformer through mflux.
108
+
109
+ Bit-identical to the same transformer run from the BF16 shards, block by block; the
110
+ transformer stays compressed and each block is decoded on the GPU as it runs. Text encoders
111
+ and the compressed set are never resident together: a new prompt after a generation drops the
112
+ set, reloads the encoders, encodes and reloads the set; ``encode(*prompts)`` pays that once for
113
+ several prompts. In-loop callbacks that decode through the VAE run with the set resident and
114
+ under the call's cache limit, and are not supported on this path. Each call resets MLX's
115
+ process-wide peak-memory counter at every phase boundary.
116
+ """
117
+
118
+ def __init__(
119
+ self,
120
+ model: str = "schnell",
121
+ *,
122
+ df11_path: str | None = None,
123
+ base_path: str | None = None,
124
+ eval_policy: str = "per-block",
125
+ cache_limit: int | None = None,
126
+ fit_check: bool = True,
127
+ budget_bytes: int | None = None,
128
+ quantize: int | None = None,
129
+ lora_paths: list[str] | None = None,
130
+ lora_scales: list[float] | None = None,
131
+ bake_lora: bool = True,
132
+ ) -> None:
133
+ """Resolve the checkpoint and the base repository and build the model.
134
+
135
+ The checkpoint headers, the base's tokenizers and the transformer's extras are read; the
136
+ encoder, VAE and compressed weights stay lazy until used. ``budget_bytes`` replaces the
137
+ device's budget (its recommended working set minus 2 GiB) in every call's fit estimate and
138
+ VAE decision, e.g. with a smaller Mac's watchdog ceiling when running under its limits.
139
+
140
+ Raises:
141
+ DFloatUnsupportedError: ``quantize``, ``lora_paths``, ``lora_scales`` or ``bake_lora=False``
142
+ (accepted only to refuse), an eval policy other than ``"per-block"`` / ``"depth2"``,
143
+ or an unknown ``model``.
144
+ DFloatFormatError: The checkpoint is not a FLUX.1 DF11 checkpoint, or the base lacks the
145
+ encoders or is a quantized save.
146
+ DFloatAccessError: A gated repository this account may not read.
147
+ DFloatDependencyError: The ``mlx-dfloat[mflux]`` extra is not installed.
148
+ DFloatIntegrationError: mflux's own weight definition or mapping does not shape the
149
+ way this adapter expects (a name map ambiguity, an uncovered extra).
150
+ """
151
+ if quantize is not None:
152
+ _refuse(
153
+ "quantize",
154
+ "quantisation on top of DFloat11 changes the output the format exists to keep",
155
+ )
156
+ if lora_paths:
157
+ _refuse("lora_paths", "LoRA of any kind")
158
+ if lora_scales:
159
+ _refuse("lora_scales", "LoRA of any kind")
160
+ if not bake_lora:
161
+ _refuse("bake_lora", "LoRA of any kind")
162
+ _check_eval_policy(eval_policy)
163
+ if model not in MODELS:
164
+ raise DFloatUnsupportedError(f"model {model!r}: this path runs {sorted(MODELS)}")
165
+ from mflux.models.common.config.model_config import ModelConfig
166
+
167
+ model_config = ModelConfig.from_name(model_name=model, base_model=None)
168
+ df11_default, base_default = MODELS[model]
169
+ df11 = base_init.resolve(df11_path or df11_default, patterns=base_init.DF11_PATTERNS)
170
+ ckpt = open_checkpoint(df11.root)
171
+ check_flux_groups(ckpt)
172
+ base = base_init.resolve(base_path or base_default, patterns=base_init.BASE_PATTERNS)
173
+ components = base_init.load_base(base.root, model_config)
174
+ transformer, shapes = build_transformer(model_config, ckpt)
175
+ self._assemble(
176
+ model=model,
177
+ model_config=model_config,
178
+ ckpt=ckpt,
179
+ df11=df11,
180
+ base=base,
181
+ components=components,
182
+ transformer=transformer,
183
+ shapes=shapes,
184
+ sizes=sizes_for(ckpt, base.root),
185
+ name_map=None,
186
+ eval_policy=eval_policy,
187
+ cache_limit=cache_limit,
188
+ fit_check=fit_check,
189
+ budget_bytes=budget_bytes,
190
+ )
191
+
192
+ @classmethod
193
+ def _from_parts(cls, **parts: Any) -> "DFloatFlux1":
194
+ """Assemble a model from already loaded parts (what ``__init__`` ends with)."""
195
+ self = cls.__new__(cls)
196
+ self._assemble(**parts)
197
+ return self
198
+
199
+ def _assemble(
200
+ self,
201
+ *,
202
+ model: str = "custom",
203
+ model_config: Any,
204
+ ckpt: DF11Checkpoint,
205
+ df11: ResolvedRepo,
206
+ base: ResolvedRepo,
207
+ components: BaseComponents,
208
+ transformer: Any,
209
+ shapes: Shapes,
210
+ sizes: FluxSizes,
211
+ name_map: NameMap | None = None,
212
+ eval_policy: str = "per-block",
213
+ cache_limit: int | None = None,
214
+ fit_check: bool = True,
215
+ budget_bytes: int | None = None,
216
+ ) -> None:
217
+ from mflux.models.flux.flux_initializer import FluxInitializer
218
+
219
+ _check_eval_policy(eval_policy)
220
+ nn.Module.__init__(self) # type: ignore[attr-defined] # not Flux1.__init__: mflux's own initializer runs over the whole base
221
+ FluxInitializer._init_config(
222
+ self, model_config
223
+ ) # prompt_cache, model_config, callbacks, tiling_config
224
+ self.vae = components.vae
225
+ self.t5_text_encoder = components.t5
226
+ self.clip_text_encoder = components.clip
227
+ self.tokenizers = components.tokenizers
228
+ self.transformer = transformer
229
+ self.bits = None
230
+ self.lora_paths: list[str] = []
231
+ self.lora_scales: list[float] = []
232
+ self._model = model
233
+ self._ckpt, self._df11, self._base, self._shapes, self._sizes = (
234
+ ckpt,
235
+ df11,
236
+ base,
237
+ shapes,
238
+ sizes,
239
+ )
240
+ self._names = flux_name_map() if name_map is None else name_map
241
+ self._largest = largest_decoded_bytes(ckpt, self._names)
242
+ self._policy, self._cache_limit_override, self._fit_check = (
243
+ eval_policy,
244
+ cache_limit,
245
+ fit_check,
246
+ )
247
+ self._budget_override = budget_bytes
248
+ self._provider: DF11Provider | None = None
249
+ self._plan: CallPlan | None = None
250
+ self._peaks: dict[str, dict[str, int]] = {}
251
+ self._open_phase: str | None = None
252
+ self._launches = 0
253
+ self._baseline_active = int(mx.get_active_memory())
254
+ self._lifecycle = Lifecycle(
255
+ load_encoders=self._load_encoders,
256
+ unload_encoders=self._unload_encoders,
257
+ encode=self._encode,
258
+ load_set=self._load_set,
259
+ unload_set=self._unload_set,
260
+ prompt_cache=self.prompt_cache,
261
+ retained_bound=self._retained_bound,
262
+ encoders_loaded=True,
263
+ )
264
+ self.callbacks.register(_VaePoolGuard(self))
265
+
266
+ # --- lifecycle callbacks ------------------------------------------------------------------------
267
+
268
+ def _retained_bound(self) -> int:
269
+ """What may stay active after a drop: the assembly baseline, the cached embeddings, the VAE, slack."""
270
+ embeddings = sum(int(a.nbytes) for pair in self.prompt_cache.values() for a in pair)
271
+ flat: list[tuple[str, mx.array]] = list(tree_flatten(self.vae.parameters())) # type: ignore[arg-type]
272
+ vae = sum(int(v.nbytes) for _name, v in flat)
273
+ return self._baseline_active + embeddings + vae + RETAINED_SLACK_BYTES
274
+
275
+ def _load_encoders(self) -> None:
276
+ self.t5_text_encoder, self.clip_text_encoder = base_init.load_encoders(self._base.root)
277
+
278
+ def _unload_encoders(self) -> None:
279
+ self.t5_text_encoder = None
280
+ self.clip_text_encoder = None
281
+
282
+ def _encode(self, prompt: str) -> tuple[mx.array, mx.array]:
283
+ from mflux.models.flux.model.flux_text_encoder.prompt_encoder import PromptEncoder
284
+
285
+ pair: tuple[mx.array, mx.array] = PromptEncoder.encode_prompt(
286
+ prompt,
287
+ prompt_cache={}, # the lifecycle caches the evaluated pair itself
288
+ t5_tokenizer=self.tokenizers["t5"],
289
+ clip_tokenizer=self.tokenizers["clip"],
290
+ t5_text_encoder=self.t5_text_encoder,
291
+ clip_text_encoder=self.clip_text_encoder,
292
+ )
293
+ return pair
294
+
295
+ def _load_set(self) -> None:
296
+ resident: dict[str, MxGroup] = load_resident_set(self._ckpt)
297
+ self._provider = DF11Provider(
298
+ resident, {n: self._ckpt.groups[n].matrix_names for n in resident}, self._names
299
+ )
300
+ self.transformer.attach(
301
+ self._provider, self._shapes, eval_policy=self._policy, verify_in_call=True
302
+ )
303
+
304
+ def _unload_set(self) -> None:
305
+ if self._provider is not None:
306
+ self._launches += self._provider.launches
307
+ self.transformer.detach()
308
+ self._provider = None
309
+
310
+ # --- phases -----------------------------------------------------------------------------------
311
+
312
+ def _phase_begin(self, name: str) -> None:
313
+ mx.reset_peak_memory() # resets to zero, not to what is active
314
+ self._open_phase = name
315
+ self._peaks[name] = {
316
+ "footprint_start": phys_footprint(),
317
+ "mlx_peak": 0,
318
+ "footprint_end": 0,
319
+ "active_at_start": int(mx.get_active_memory()),
320
+ }
321
+
322
+ def _phase_end(self, name: str) -> None:
323
+ if self._open_phase != name:
324
+ return
325
+ self._peaks[name]["mlx_peak"] = max(
326
+ int(mx.get_peak_memory()), self._peaks[name]["active_at_start"]
327
+ )
328
+ self._peaks[name]["footprint_end"] = phys_footprint()
329
+ self._open_phase = None
330
+
331
+ # --- the public surface -----------------------------------------------------------------------
332
+
333
+ def plan_call(self, *, height: int, width: int) -> CallPlan:
334
+ """The cache limit, the fit estimate and the VAE strategy of a call at this size (rounded down to multiples of 16, as mflux does).
335
+
336
+ The set stays resident through the VAE decode when the estimate allows it; otherwise the call drops it
337
+ before decoding and the next call reloads it. On a 32 GB Mac (a 22.96 GiB budget) every FLUX.1 call
338
+ drops the set before the decode, whatever the size: the measured decode transient is a floor, not
339
+ smaller for smaller images. A larger budget keeps the set. Sizes above 1024² were not measured on
340
+ this path; the estimate there is an extrapolation.
341
+
342
+ Raises:
343
+ DFloatResourceError: The size is above 1024² or the predicted peak exceeds the budget even with
344
+ the set dropped, and ``fit_check`` is on.
345
+ """
346
+ height, width = 16 * (height // 16), 16 * (width // 16)
347
+ if height * width > MAX_MEASURED_PIXELS:
348
+ if self._fit_check:
349
+ raise DFloatResourceError(
350
+ f"{height}x{width}: above the measured ceiling of 1024x1024 ({MAX_MEASURED_PIXELS} "
351
+ "pixels) on this path; pass fit_check=False to run on an extrapolated estimate"
352
+ )
353
+ log.warning(
354
+ "%dx%d: no measurement above 1024² on this path; the estimate is an extrapolation",
355
+ height,
356
+ width,
357
+ )
358
+ tokens = text_tokens(self.model_config)
359
+ derived_minimum = self._largest[DOUBLE_PREFIX] + self._largest[SINGLE_PREFIX]
360
+ limit = cache_limit_for(
361
+ self._largest,
362
+ policy=self._policy,
363
+ height=height,
364
+ width=width,
365
+ text_tokens=tokens,
366
+ override=self._cache_limit_override,
367
+ )
368
+ if limit < derived_minimum:
369
+ log.warning(
370
+ "cache_limit %d is below the derived minimum %d (the two largest decoded groups): "
371
+ "every block's decode output will be allocated fresh",
372
+ limit,
373
+ derived_minimum,
374
+ )
375
+ allowance = activation_allowance(height=height, width=width, text_tokens=tokens)
376
+ budget = self._budget_override if self._budget_override is not None else budget_bytes()
377
+ estimate = fit_for(
378
+ sizes=self._sizes,
379
+ largest=self._largest,
380
+ policy=self._policy,
381
+ cache_limit=limit,
382
+ allowance=allowance,
383
+ budget=budget,
384
+ height=height,
385
+ width=width,
386
+ text_tokens=tokens,
387
+ )
388
+ drop_set_before_vae = estimate.phases["vae"] > budget
389
+ if drop_set_before_vae:
390
+ estimate = fit_for(
391
+ sizes=self._sizes,
392
+ largest=self._largest,
393
+ policy=self._policy,
394
+ cache_limit=limit,
395
+ allowance=allowance,
396
+ budget=budget,
397
+ height=height,
398
+ width=width,
399
+ text_tokens=tokens,
400
+ vae_with_set=False,
401
+ )
402
+ if not estimate.fits:
403
+ phases = ", ".join(f"{k} {v / 1024**3:.1f} GiB" for k, v in estimate.phases.items())
404
+ message = (
405
+ f"predicted peak {estimate.peak_bytes / 1024**3:.1f} GiB in the {estimate.peak_phase} phase "
406
+ f"exceeds the budget {estimate.budget_bytes / 1024**3:.1f} GiB ({phases})"
407
+ )
408
+ if self._fit_check:
409
+ raise DFloatResourceError(f"{message}; pass fit_check=False to run anyway")
410
+ log.warning("%s; running anyway (fit_check=False)", message)
411
+ plan = CallPlan(
412
+ cache_limit=limit, estimate=estimate, drop_set_before_vae=drop_set_before_vae
413
+ )
414
+ self._plan = plan
415
+ return plan
416
+
417
+ def encode(self, *prompts: str) -> None:
418
+ """Encode prompts now, so several generations pay the encoder reload once (drops a resident set first)."""
419
+ self._lifecycle.ensure_embeddings(*prompts)
420
+
421
+ def generate_image(
422
+ self,
423
+ seed: int,
424
+ prompt: str,
425
+ num_inference_steps: int = 4,
426
+ height: int = 1024,
427
+ width: int = 1024,
428
+ guidance: float = 4.0,
429
+ image_path: Path | str | None = None,
430
+ image_strength: float | None = None,
431
+ scheduler: str = "linear",
432
+ negative_prompt: str | None = None,
433
+ pid_decode: bool = False,
434
+ pid_degrade_sigma: float = 0.0,
435
+ ) -> Any:
436
+ """Mflux's ``generate_image`` behind a prelude: refusals, cache limit, fit check, prompt encoding, set load.
437
+
438
+ ``negative_prompt`` is accepted and ignored, as upstream does for FLUX.1.
439
+
440
+ Raises:
441
+ DFloatUnsupportedError: ``image_path`` / ``image_strength`` (img2img) or ``pid_decode``.
442
+ DFloatResourceError: The fit check refuses the call (the estimate, or a size above 1024²), or
443
+ memory stayed active after a drop.
444
+ DFloatFormatError: A block's decode reported an error; the set is dropped for a clean retry.
445
+ """
446
+ if image_path is not None or image_strength is not None:
447
+ _refuse("image_path/image_strength", "img2img")
448
+ if pid_decode:
449
+ _refuse(
450
+ "pid_decode",
451
+ "the PiD decoder loads an 8 GB caption encoder next to the compressed set",
452
+ )
453
+ plan = self.plan_call(height=height, width=width)
454
+ self._phase_begin("encode")
455
+ self._lifecycle.ensure_embeddings(prompt)
456
+ self._phase_end("encode")
457
+ self._phase_begin("set_load")
458
+ self._lifecycle.ensure_set()
459
+ self._phase_end("set_load")
460
+ self._phase_begin("denoise") # before set_cache_limit: a raise here must change nothing
461
+ previous = mx.set_cache_limit(plan.cache_limit)
462
+ try:
463
+ return super().generate_image(
464
+ seed=seed,
465
+ prompt=prompt,
466
+ num_inference_steps=num_inference_steps,
467
+ height=height,
468
+ width=width,
469
+ guidance=guidance,
470
+ image_path=None,
471
+ image_strength=None,
472
+ scheduler=scheduler,
473
+ negative_prompt=negative_prompt,
474
+ pid_decode=False,
475
+ pid_degrade_sigma=pid_degrade_sigma,
476
+ )
477
+ except DFloatFormatError as exc:
478
+ # The exception's own traceback keeps every frame between the raise site (in
479
+ # production, inside SeamMixin.__call__ / DF11Provider.verify()) and here alive, and
480
+ # those frames can reference the whole resident set. Clear them before drop_set's own
481
+ # gc.collect() runs, or _reclaim's active-memory check turns this clean format error
482
+ # into a DFloatResourceError with the real error demoted to __context__.
483
+ traceback.clear_frames(exc.__traceback__)
484
+ self._lifecycle.drop_set() # a corrupt block: the retry starts from a clean load
485
+ raise
486
+ finally:
487
+ mx.set_cache_limit(previous)
488
+ self._phase_end("denoise")
489
+ self._phase_end("vae")
490
+
491
+ def report(self) -> dict[str, Any]:
492
+ """What this model is and what its calls cost: repositories, policy, limits, estimate, peaks, counts, versions.
493
+
494
+ Peaks are sampled at phase boundaries (labelled ``"sampled"``); the fit estimate is a prediction.
495
+ """
496
+ launches = self._launches + (self._provider.launches if self._provider is not None else 0)
497
+ fit = None if self._plan is None else self._plan.estimate
498
+ try:
499
+ mflux_version: str | None = metadata.version("mflux")
500
+ except metadata.PackageNotFoundError:
501
+ mflux_version = None
502
+ return {
503
+ "model": self._model,
504
+ "df11": {
505
+ "root": str(self._df11.root),
506
+ "repo_id": self._df11.repo_id,
507
+ "revision": self._df11.revision,
508
+ },
509
+ "base": {
510
+ "root": str(self._base.root),
511
+ "repo_id": self._base.repo_id,
512
+ "revision": self._base.revision,
513
+ },
514
+ "eval_policy": self._policy,
515
+ "cache_limit_in_force": None if self._plan is None else self._plan.cache_limit,
516
+ "drop_set_before_vae": None if self._plan is None else self._plan.drop_set_before_vae,
517
+ "fit": None
518
+ if fit is None
519
+ else {
520
+ "label": "predicted",
521
+ "phases": dict(fit.phases),
522
+ "peak_phase": fit.peak_phase,
523
+ "peak_bytes": fit.peak_bytes,
524
+ "budget_bytes": fit.budget_bytes,
525
+ "fits": fit.fits,
526
+ },
527
+ "sizes": asdict(self._sizes),
528
+ "peaks": {"label": "sampled at phase boundaries", **self._peaks},
529
+ "decode_launches": launches,
530
+ "lifecycle": self._lifecycle.counters.as_dict(),
531
+ "versions": {"mlx": mx.__version__, "mflux": mflux_version}, # type: ignore[attr-defined]
532
+ }
533
+
534
+ @staticmethod
535
+ def from_name(model_name: str, quantize: int | None = None) -> "DFloatFlux1":
536
+ """A model by its mflux name; ``quantize`` is accepted only to refuse it."""
537
+ return DFloatFlux1(model_name, quantize=quantize)
538
+
539
+ def save_model(self, base_path: str) -> None:
540
+ """Refused: the DFloat11 checkpoint is the saved form; there is nothing of mflux's to write.
541
+
542
+ Raises:
543
+ DFloatUnsupportedError: Always.
544
+ """
545
+ del base_path
546
+ _refuse("save_model", "the DFloat11 repository is the checkpoint")
547
+
548
+ def freeze(self, **kwargs: Any) -> None:
549
+ """Freeze the components that are loaded (a dropped encoder is ``None``)."""
550
+ del kwargs
551
+ for module in (self.vae, self.transformer, self.t5_text_encoder, self.clip_text_encoder):
552
+ if module is not None:
553
+ module.freeze()
@@ -0,0 +1,79 @@
1
+ """FLUX.1 naming: the map derived from mflux's own weight mapping, and the group-kind check."""
2
+
3
+ from typing import Any
4
+
5
+ from mlx_dfloat.errors import DFloatFormatError, DFloatIntegrationError
6
+ from mlx_dfloat.integrate.names import StaticNameMap
7
+ from mlx_dfloat.mflux import require_mflux
8
+
9
+ DOUBLE_PREFIX = "transformer_blocks"
10
+ SINGLE_PREFIX = "single_transformer_blocks"
11
+ MATRICES_PER_KIND = {DOUBLE_PREFIX: 14, SINGLE_PREFIX: 6}
12
+ # mflux 0.20 builds `norm_out.linear` with bias=False and drops this bias through update(strict=False);
13
+ # the seam follows mflux so its output equals mflux's.
14
+ DROPPED_EXTRAS: frozenset[str] = frozenset({"norm_out.linear.bias"})
15
+
16
+
17
+ def flux_name_map() -> StaticNameMap:
18
+ """The block-matrix map read from mflux's ``FluxWeightMapping`` (a pure rename; imports mflux).
19
+
20
+ Raises:
21
+ DFloatDependencyError: The optional ``mflux`` extra is not installed.
22
+ DFloatIntegrationError: A block matrix target has several source patterns or a transform, which this
23
+ adapter cannot express.
24
+ """
25
+ require_mflux()
26
+ from mflux.models.flux.weights.flux_weight_mapping import FluxWeightMapping
27
+
28
+ tables: dict[str, dict[str, str]] = {DOUBLE_PREFIX: {}, SINGLE_PREFIX: {}}
29
+ for target in FluxWeightMapping.get_transformer_mapping():
30
+ to = target.to_pattern
31
+ kind = next((k for k in tables if to.startswith(f"{k}.{{block}}.")), None)
32
+ if kind is None or not to.endswith(".weight"):
33
+ continue
34
+ sources = (
35
+ target.from_pattern if isinstance(target.from_pattern, list) else [target.from_pattern]
36
+ )
37
+ if len(sources) != 1 or getattr(target, "transform", None) is not None:
38
+ raise DFloatIntegrationError(
39
+ f"{to}: mflux maps it from {sources} with a transform; not a pure rename"
40
+ )
41
+ prefix = f"{kind}.{{block}}."
42
+ sub = sources[0].removeprefix(prefix).removesuffix(".weight")
43
+ attr = to.removeprefix(prefix).removesuffix(".weight")
44
+ tables[kind][sub] = attr
45
+ keep = {k: {s: a for s, a in v.items() if _is_matrix(k, s)} for k, v in tables.items()}
46
+ return StaticNameMap(keep)
47
+
48
+
49
+ def _is_matrix(kind: str, sub: str) -> bool:
50
+ """Whether ``sub`` (a checkpoint sub-path of ``kind``) is a DF11 matrix rather than a norm scale."""
51
+ del kind # kept for a readable call site; every FLUX kind uses the same norm-scale rule
52
+ # Norm weights (`attn.norm_q`, `norm_k`, `norm_added_q`, ...) are RMSNorm scales, not DF11 matrices.
53
+ return not any(part.startswith("norm_") for part in sub.split("."))
54
+
55
+
56
+ def check_flux_groups(ckpt: Any) -> tuple[int, int]:
57
+ """Count the double and single blocks; the groups must be exactly those two contiguous families.
58
+
59
+ Raises:
60
+ DFloatFormatError: A group outside the two kinds, a hole in an index sequence, a missing kind, or a
61
+ group with the wrong number of matrices.
62
+ """
63
+ seen: dict[str, set[int]] = {DOUBLE_PREFIX: set(), SINGLE_PREFIX: set()}
64
+ for name, group in ckpt.groups.items():
65
+ kind, _dot, idx = name.partition(".")
66
+ if kind not in seen or not (idx.isascii() and idx.isdigit()):
67
+ raise DFloatFormatError(f"{name}: not a FLUX block group")
68
+ n = len(group.matrix_names)
69
+ if n != MATRICES_PER_KIND[kind]:
70
+ raise DFloatFormatError(
71
+ f"{name}: {n} matrices, a {kind} group has {MATRICES_PER_KIND[kind]}"
72
+ )
73
+ seen[kind].add(int(idx))
74
+ for kind, indices in seen.items():
75
+ if not indices:
76
+ raise DFloatFormatError(f"no {kind} groups in the checkpoint")
77
+ if indices != set(range(len(indices))):
78
+ raise DFloatFormatError(f"{kind} groups are not contiguous: {sorted(indices)}")
79
+ return len(seen[DOUBLE_PREFIX]), len(seen[SINGLE_PREFIX])