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 @@
1
+ """FLUX.1 (dev, schnell, Krea-dev) on mflux from a DFloat11 transformer."""
@@ -0,0 +1,419 @@
1
+ """``mlx-dfloat generate``: one FLUX.1 image from a DFloat11 transformer, with memory caps and a watchdog.
2
+
3
+ Flag names follow ``mflux-generate`` so a pasted command works; the options that path cannot
4
+ honour are parsed only to refuse them with a reason (exit 2). ``--tier GB`` runs under a smaller
5
+ Mac's MLX limits, with that tier's watchdog ceiling and fit budget (the host's own tier keeps the
6
+ host caps); ``--memory-ceiling BYTES`` sets the watchdog ceiling alone, under the host caps.
7
+ """
8
+
9
+ import argparse
10
+ import json
11
+ import sys
12
+ import traceback
13
+ from collections.abc import Callable
14
+ from pathlib import Path
15
+ from typing import Any
16
+
17
+ import mlx.core as mx
18
+
19
+ from mlx_dfloat._memory_caps import install_memory_caps
20
+ from mlx_dfloat._scrub import scrub_home
21
+ from mlx_dfloat._watchdog import Watchdog, default_ceiling, phys_footprint
22
+ from mlx_dfloat.bench import capped
23
+ from mlx_dfloat.bench.capped import TierLimits, host_tier_gb, limits_record, tier_limits
24
+ from mlx_dfloat.errors import DFloatError
25
+
26
+ EXIT_OK, EXIT_ERROR = 0, 2
27
+ DEFAULT_STEPS = {"schnell": 4, "dev": 25, "krea-dev": 25} # mflux 0.20's per-model defaults
28
+ DEFAULT_GUIDANCE = 3.5 # mflux-generate's default; schnell ignores it
29
+ FOOTPRINT_PEAK_LABEL = "OS phys_footprint, sampled every 0.05 s by the watchdog"
30
+ REFUSED: dict[str, str] = {
31
+ "--quantize": "quantisation on top of DFloat11 changes the output the format exists to keep",
32
+ "--lora-paths": "LoRA is not on the DFloat11 path",
33
+ "--lora-scales": "LoRA is not on the DFloat11 path",
34
+ "--lora": "LoRA is not on the DFloat11 path",
35
+ "--lora-style": "LoRA is not on the DFloat11 path",
36
+ "--image-path": "img2img is not on the DFloat11 path",
37
+ "--image-strength": "img2img is not on the DFloat11 path",
38
+ "--image": "img2img is not on the DFloat11 path",
39
+ "--pid-decode": "the PiD decoder loads an 8 GB caption encoder next to the compressed set",
40
+ "--controlnet-image-path": "ControlNet is another model class, not provided on the DFloat11 path",
41
+ "--controlnet-strength": "ControlNet is another model class, not provided on the DFloat11 path",
42
+ }
43
+ _REFUSED_ALIASES = {"--quantize": ["-q"]}
44
+
45
+
46
+ def _positive_bytes(text: str, what: str) -> int:
47
+ try:
48
+ value = int(float(text))
49
+ except OverflowError as exc: # "inf"
50
+ raise ValueError(f"{what} must be a finite byte count, got {text!r}") from exc
51
+ if value <= 0:
52
+ raise ValueError(f"{what} must be positive, got {text!r}")
53
+ return value
54
+
55
+
56
+ def cache_limit_bytes(text: str) -> int:
57
+ """A positive byte count, ``2.5e9`` accepted.
58
+
59
+ Raises:
60
+ ValueError: Not a finite number, or not positive.
61
+ """
62
+ return _positive_bytes(text, "cache limit")
63
+
64
+
65
+ def memory_ceiling_bytes(text: str) -> int:
66
+ """A positive byte count for the watchdog ceiling, ``1.8e10`` accepted.
67
+
68
+ Raises:
69
+ ValueError: Not a finite number, or not positive.
70
+ """
71
+ return _positive_bytes(text, "memory ceiling")
72
+
73
+
74
+ def tier_gb(text: str) -> int:
75
+ """A tier in whole GB, at least 1.
76
+
77
+ Raises:
78
+ ValueError: Not an integer, or below 1.
79
+ """
80
+ value = int(text)
81
+ if value < 1:
82
+ raise ValueError(f"tier must be at least 1 GB, got {text!r}")
83
+ return value
84
+
85
+
86
+ def add_generate_parser(sub: Any) -> argparse.ArgumentParser:
87
+ """Register ``generate`` on a subparsers object."""
88
+ p: argparse.ArgumentParser = sub.add_parser(
89
+ "generate",
90
+ help="generate one FLUX.1 image from a DFloat11 transformer",
91
+ description=__doc__,
92
+ )
93
+ _add_arguments(p)
94
+ p.set_defaults(run=run)
95
+ return p
96
+
97
+
98
+ def build_parser() -> argparse.ArgumentParser:
99
+ """A standalone parser for ``generate`` (tests parse with it)."""
100
+ p = argparse.ArgumentParser(prog="mlx-dfloat generate", description=__doc__)
101
+ _add_arguments(p)
102
+ return p
103
+
104
+
105
+ def _add_arguments(p: argparse.ArgumentParser) -> None:
106
+ p.add_argument("--model", "-m", choices=tuple(DEFAULT_STEPS), default="schnell")
107
+ p.add_argument("--prompt", required=True, help="the text to generate")
108
+ p.add_argument("--seed", type=int, default=42)
109
+ p.add_argument(
110
+ "--steps",
111
+ type=int,
112
+ default=None,
113
+ help="denoise steps (default per model: schnell 4, dev 25)",
114
+ )
115
+ p.add_argument("--height", type=int, default=1024)
116
+ p.add_argument("--width", type=int, default=1024)
117
+ p.add_argument(
118
+ "--guidance",
119
+ type=float,
120
+ default=None,
121
+ help=f"guidance (default {DEFAULT_GUIDANCE}; schnell ignores it)",
122
+ )
123
+ p.add_argument("--scheduler", default="linear")
124
+ p.add_argument(
125
+ "--negative-prompt", default=None, help="accepted and ignored, as mflux does for FLUX.1"
126
+ )
127
+ p.add_argument(
128
+ "--output",
129
+ default="image.png",
130
+ help="the image file; an existing file is kept and the new one gets a numbered name",
131
+ )
132
+ p.add_argument(
133
+ "--metadata", action="store_true", help="also write mflux's JSON metadata sidecar"
134
+ )
135
+ p.add_argument(
136
+ "--df11",
137
+ default=None,
138
+ help="DFloat11 checkpoint: a directory or a Hub id (default per model)",
139
+ )
140
+ p.add_argument(
141
+ "--base", default=None, help="base repository for the encoders and VAE (default per model)"
142
+ )
143
+ p.add_argument("--eval-policy", choices=("per-block", "depth2"), default="per-block")
144
+ p.add_argument(
145
+ "--cache-limit",
146
+ type=cache_limit_bytes,
147
+ default=None,
148
+ help="MLX buffer-cache limit in bytes (derived per call by default)",
149
+ )
150
+ p.add_argument(
151
+ "--no-fit-check",
152
+ action="store_true",
153
+ help="skip the memory fit estimate (a warning instead of a refusal)",
154
+ )
155
+ p.add_argument(
156
+ "--report",
157
+ type=Path,
158
+ default=None,
159
+ help="write the run report as JSON (an existing file is replaced)",
160
+ )
161
+ p.add_argument(
162
+ "--wall-budget",
163
+ type=float,
164
+ default=3600.0,
165
+ help="seconds before the watchdog aborts (exit 71)",
166
+ )
167
+ p.add_argument(
168
+ "--tier",
169
+ type=tier_gb,
170
+ default=None,
171
+ metavar="GB",
172
+ help="run under a smaller Mac's MLX limits, watchdog ceiling and fit budget "
173
+ "(default: this host's own tier and caps)",
174
+ )
175
+ p.add_argument(
176
+ "--memory-ceiling",
177
+ type=memory_ceiling_bytes,
178
+ default=None,
179
+ metavar="BYTES",
180
+ help="a lower watchdog ceiling alone, under the host caps (not with --tier)",
181
+ )
182
+ for flag, reason in REFUSED.items():
183
+ p.add_argument(
184
+ flag,
185
+ *_REFUSED_ALIASES.get(flag, []),
186
+ nargs="*",
187
+ default=None,
188
+ help=argparse.SUPPRESS,
189
+ metavar=reason,
190
+ )
191
+
192
+
193
+ def refused_option(args: argparse.Namespace) -> str | None:
194
+ """The first refused flag present on the command line, with its reason; None when there is none."""
195
+ for flag, reason in REFUSED.items():
196
+ if getattr(args, flag.lstrip("-").replace("-", "_")) is not None:
197
+ return f"{flag}: {reason}"
198
+ return None
199
+
200
+
201
+ def ceiling_for(
202
+ args: argparse.Namespace,
203
+ *,
204
+ host_ram_bytes: int,
205
+ host_recommended_bytes: int,
206
+ default_ceiling_bytes: int,
207
+ ) -> tuple[int, TierLimits, str]:
208
+ """The watchdog ceiling, the limits to install and the report label for these flags (pure).
209
+
210
+ ``--memory-ceiling`` gives that ceiling under the host's limits (``PROOF``); ``--tier`` gives
211
+ the tier's ceiling and limits, except the host's own tier, which keeps ``default_ceiling_bytes``
212
+ (``MEASURED``); neither gives the host's limits and ``default_ceiling_bytes`` (``MEASURED``).
213
+
214
+ Raises:
215
+ ValueError: Both ``--tier`` and ``--memory-ceiling`` were given, or ``--memory-ceiling`` is
216
+ above ``default_ceiling_bytes`` (it may only lower the ceiling).
217
+ DFloatUnsupportedError: The tier is larger than the host.
218
+ """
219
+ tier, ceiling = args.tier, args.memory_ceiling
220
+ if tier is not None and ceiling is not None:
221
+ raise ValueError(
222
+ "--tier and --memory-ceiling cannot be combined: --tier sets a tier's limits and "
223
+ "ceiling, --memory-ceiling the watchdog ceiling alone under the host caps"
224
+ )
225
+ limits = tier_limits(
226
+ host_tier_gb(host_ram_bytes) if tier is None else tier,
227
+ host_ram_bytes=host_ram_bytes,
228
+ host_recommended_bytes=host_recommended_bytes,
229
+ )
230
+ if ceiling is not None:
231
+ if ceiling > default_ceiling_bytes:
232
+ raise ValueError(
233
+ f"--memory-ceiling {ceiling} is above this host's ceiling "
234
+ f"{default_ceiling_bytes / 1024**3:.1f} GiB: it may only lower the watchdog ceiling"
235
+ )
236
+ return ceiling, limits, "PROOF"
237
+ return (default_ceiling_bytes if limits.is_host else limits.ceiling_bytes), limits, limits.label
238
+
239
+
240
+ def _steps(args: argparse.Namespace) -> int:
241
+ """The denoise steps this call runs: ``--steps``, else the model's mflux default."""
242
+ return int(args.steps) if args.steps is not None else DEFAULT_STEPS[args.model]
243
+
244
+
245
+ def _run_context(args: argparse.Namespace) -> dict[str, Any]:
246
+ """What the watchdog's abort artifact records about the run it may stop."""
247
+ return {
248
+ "model": args.model,
249
+ "height": args.height,
250
+ "width": args.width,
251
+ "seed": args.seed,
252
+ "steps": _steps(args),
253
+ }
254
+
255
+
256
+ def _host_facts() -> tuple[int, int]:
257
+ """This Mac's RAM and recommended working set, from MLX's device info (0 when not reported)."""
258
+ info = mx.device_info()
259
+ return int(info.get("memory_size", 0)), int(info.get("max_recommended_working_set_size", 0))
260
+
261
+
262
+ def _model_class() -> Callable[..., Any]:
263
+ from mlx_dfloat.mflux.flux1.model import (
264
+ DFloatFlux1, # raises DFloatDependencyError without mflux
265
+ )
266
+
267
+ return DFloatFlux1
268
+
269
+
270
+ def _output_resolver() -> Callable[[Path], Path]:
271
+ """The saved name by mflux's own rule (an existing file gets a numbered name next to it)."""
272
+ from mflux.utils.image_util import ImageUtil
273
+
274
+ resolve: Callable[[Path], Path] = ImageUtil.resolve_output_path
275
+ return resolve
276
+
277
+
278
+ def _finish(args: argparse.Namespace, report: dict[str, Any]) -> int:
279
+ """Write the report (when ``--report`` was given) and announce success. The one path every exit uses."""
280
+ if args.report is not None:
281
+ try:
282
+ args.report.parent.mkdir(parents=True, exist_ok=True)
283
+ args.report.write_text(json.dumps(scrub_home(report), indent=1, default=str))
284
+ except OSError as exc:
285
+ print(f"error: the report was not written: {exc}", file=sys.stderr)
286
+ return EXIT_ERROR
287
+ if report["exit_code"] == EXIT_OK:
288
+ peak = report.get("footprint_peak_bytes")
289
+ suffix = f" (footprint peak {peak / 1024**3:.2f} GiB)" if peak is not None else ""
290
+ print(f"ok: {report['output']}{suffix}")
291
+ return int(report["exit_code"])
292
+
293
+
294
+ def run(
295
+ args: argparse.Namespace,
296
+ *,
297
+ model_factory: Callable[..., Any] | None = None,
298
+ install_caps: Callable[[], tuple[int, int]] = install_memory_caps,
299
+ watchdog_factory: Callable[..., Any] = Watchdog,
300
+ resolve_output: Callable[[Path], Path] | None = None,
301
+ host_facts: Callable[[], tuple[int, int]] = _host_facts,
302
+ ceiling_default: Callable[[], int] = default_ceiling,
303
+ apply_limits: Callable[[TierLimits], object] = capped.apply,
304
+ read_limits: Callable[[], dict[str, int]] = capped.current_limits,
305
+ ) -> int:
306
+ """Refuse, cap, watch, build, generate, save, report. Returns the exit code.
307
+
308
+ ``resolve_output`` maps the requested file to the one written (default: mflux's rule, looked
309
+ up lazily next to the model class). ``host_facts`` returns this host's RAM and recommended
310
+ working set; ``ceiling_default`` the host's watchdog ceiling; ``apply_limits`` installs a
311
+ smaller tier's limits (in place of ``install_caps``); ``read_limits`` reads the limits in force.
312
+ """
313
+ refused = refused_option(args)
314
+ if refused is not None:
315
+ print(f"error: {refused}", file=sys.stderr)
316
+ return EXIT_ERROR
317
+ if args.negative_prompt:
318
+ print(
319
+ "warning: --negative-prompt is ignored: FLUX.1 has no negative branch", file=sys.stderr
320
+ )
321
+ output = Path(args.output)
322
+ model_kwargs = {
323
+ "model": args.model,
324
+ "df11_path": args.df11,
325
+ "base_path": args.base,
326
+ "eval_policy": args.eval_policy,
327
+ "cache_limit": args.cache_limit,
328
+ "fit_check": not args.no_fit_check,
329
+ }
330
+ report: dict[str, Any] = {
331
+ "exit_code": EXIT_ERROR,
332
+ "output": str(output),
333
+ "memory_caps_gb": None,
334
+ "model_kwargs": model_kwargs,
335
+ "height": args.height,
336
+ "width": args.width,
337
+ "memory_ceiling_bytes": args.memory_ceiling,
338
+ "tier_gb": None,
339
+ "label": None,
340
+ "watchdog_ceiling_bytes": None,
341
+ "limits": None,
342
+ }
343
+ try:
344
+ host_ram, host_recommended = host_facts()
345
+ try:
346
+ ceiling, limits, label = ceiling_for(
347
+ args,
348
+ host_ram_bytes=host_ram,
349
+ host_recommended_bytes=host_recommended,
350
+ default_ceiling_bytes=ceiling_default(),
351
+ )
352
+ except ValueError as exc:
353
+ if isinstance(exc, DFloatError): # DFloatUnsupportedError: the tier above the host
354
+ raise
355
+ print(f"error: {exc}", file=sys.stderr)
356
+ report["error"] = str(exc)
357
+ return _finish(args, report)
358
+ report.update(tier_gb=limits.tier_gb, label=label, watchdog_ceiling_bytes=ceiling)
359
+ output.parent.mkdir(parents=True, exist_ok=True)
360
+ if limits.is_host:
361
+ report["memory_caps_gb"] = list(install_caps())
362
+ applied = "host-caps"
363
+ else:
364
+ apply_limits(limits)
365
+ applied = "tier-defaults"
366
+ model_kwargs["budget_bytes"] = limits.ceiling_bytes
367
+ effective = read_limits()
368
+ report["limits"] = limits_record(
369
+ limits,
370
+ effective_memory_limit=effective["memory"],
371
+ effective_cache_limit=effective["cache"],
372
+ effective_wired_limit=effective["wired"],
373
+ applied=applied,
374
+ )
375
+ watchdog = watchdog_factory(
376
+ output.parent, ceiling=ceiling, budget=args.wall_budget, context=_run_context(args)
377
+ ).start()
378
+ except (DFloatError, OSError) as exc:
379
+ print(f"error: {type(exc).__name__}: {exc}", file=sys.stderr)
380
+ report["error"] = f"{type(exc).__name__}: {exc}"
381
+ return _finish(args, report)
382
+ except Exception as exc: # symmetric with the block below: never a silent exit 1
383
+ traceback.print_exc()
384
+ report["error"] = f"{type(exc).__name__}: {exc}"
385
+ return _finish(args, report)
386
+ try:
387
+ factory = model_factory if model_factory is not None else _model_class()
388
+ resolve = resolve_output if resolve_output is not None else _output_resolver()
389
+ model = factory(**model_kwargs)
390
+ image = model.generate_image(
391
+ seed=args.seed,
392
+ prompt=args.prompt,
393
+ num_inference_steps=_steps(args),
394
+ height=args.height,
395
+ width=args.width,
396
+ guidance=args.guidance if args.guidance is not None else DEFAULT_GUIDANCE,
397
+ scheduler=args.scheduler,
398
+ )
399
+ final = Path(resolve(output))
400
+ image.save(str(final), export_json_metadata=args.metadata, overwrite=True)
401
+ if not final.is_file() or final.stat().st_size == 0:
402
+ # mflux's save logs a write failure and returns normally
403
+ raise DFloatError(f"{final}: the image was not written")
404
+ report["output"] = str(final)
405
+ report.update(exit_code=EXIT_OK, **model.report())
406
+ except DFloatError as exc:
407
+ print(f"error: {type(exc).__name__}: {exc}", file=sys.stderr)
408
+ report["error"] = f"{type(exc).__name__}: {exc}"
409
+ except Exception as exc: # anything else is still a tool error (2), never a silent success
410
+ traceback.print_exc()
411
+ report["error"] = f"{type(exc).__name__}: {exc}"
412
+ finally:
413
+ watchdog.stop()
414
+ final_sample = phys_footprint()
415
+ report["footprint_peak_bytes"] = max(watchdog.peak_footprint, final_sample)
416
+ report["footprint_peak_label"] = FOOTPRINT_PEAK_LABEL
417
+ report["watched_peak_bytes"] = max(watchdog.peak_watched, final_sample)
418
+ report["mlx_peak_bytes"] = watchdog.peak_mlx
419
+ return _finish(args, report)
@@ -0,0 +1,245 @@
1
+ """The FLUX.1 base repository without its transformer: VAE, T5, CLIP and tokenizers, resolved and loaded.
2
+
3
+ mflux's ``FluxInitializer.init`` wants every component under one path, so it would download and
4
+ size the 23.8 GB BF16 transformer; this module loads the three other components through mflux's
5
+ own loader and applier with weight definitions that name only them.
6
+ """
7
+
8
+ from dataclasses import dataclass
9
+ from pathlib import Path
10
+ from typing import Any
11
+
12
+ from mlx_dfloat.errors import DFloatAccessError, DFloatFormatError, DFloatIntegrationError
13
+ from mlx_dfloat.mflux import require_mflux
14
+
15
+ DF11_PATTERNS: tuple[str, ...] = ("*.safetensors", "config.json")
16
+ ENCODER_PATTERNS: tuple[str, ...] = (
17
+ "text_encoder/*.safetensors",
18
+ "text_encoder/*.json",
19
+ "text_encoder_2/*.safetensors",
20
+ "text_encoder_2/*.json",
21
+ )
22
+ VAE_PATTERNS: tuple[str, ...] = ("vae/*.safetensors", "vae/*.json")
23
+ TOKENIZER_PATTERNS: tuple[str, ...] = ("tokenizer/**", "tokenizer_2/**")
24
+ BASE_PATTERNS: tuple[str, ...] = ENCODER_PATTERNS + VAE_PATTERNS + TOKENIZER_PATTERNS
25
+ _SHA_LENGTH = 40
26
+
27
+
28
+ @dataclass(frozen=True, slots=True, kw_only=True)
29
+ class ResolvedRepo:
30
+ """Where a repository's files are, and which Hub repository and commit they came from (if any)."""
31
+
32
+ root: Path
33
+ repo_id: str | None
34
+ revision: str | None
35
+
36
+
37
+ @dataclass(slots=True, kw_only=True)
38
+ class BaseComponents:
39
+ """The loaded base components: mflux modules (weights still lazy) and the two tokenizers."""
40
+
41
+ vae: Any
42
+ t5: Any
43
+ clip: Any
44
+ tokenizers: dict[str, Any]
45
+
46
+
47
+ def is_hub_id(spec: str) -> bool:
48
+ """Whether ``spec`` is an ``org/name`` Hub id rather than a path (an existing path always wins)."""
49
+ if Path(spec).expanduser().exists():
50
+ return False
51
+ return "/" in spec and spec.count("/") == 1 and not spec.startswith(("./", "../", "~/", "/"))
52
+
53
+
54
+ def hub_revision(root: Path) -> str | None:
55
+ """The commit SHA of a Hub cache snapshot directory (``…/snapshots/<sha>``); None for any other directory."""
56
+ return root.name if root.parent.name == "snapshots" and len(root.name) == _SHA_LENGTH else None
57
+
58
+
59
+ def hub_error(exc: Exception, repo_id: str) -> Exception:
60
+ """Translate a Hub failure: gated or unauthorised → access error, unknown repo → format error; else ``exc``.
61
+
62
+ A cache miss whose cause is a Hub HTTP error is translated through that cause.
63
+ """
64
+ from huggingface_hub.errors import (
65
+ GatedRepoError,
66
+ HfHubHTTPError,
67
+ LocalEntryNotFoundError,
68
+ RepositoryNotFoundError,
69
+ RevisionNotFoundError,
70
+ )
71
+
72
+ if isinstance(exc, LocalEntryNotFoundError) and isinstance(exc.__cause__, HfHubHTTPError):
73
+ # huggingface_hub reports a refused download as a cache miss, the refusal as its cause
74
+ mapped = hub_error(exc.__cause__, repo_id)
75
+ return exc if mapped is exc.__cause__ else mapped
76
+ hint = "accept the licence on the Hub and run `hf auth login`"
77
+ if isinstance(exc, GatedRepoError):
78
+ return DFloatAccessError(f"{repo_id}: this repository is gated; {hint}")
79
+ if isinstance(exc, RepositoryNotFoundError | RevisionNotFoundError):
80
+ return DFloatFormatError(
81
+ f"{repo_id}: not a model repository on the Hub (or private without access)"
82
+ )
83
+ if isinstance(exc, HfHubHTTPError):
84
+ response = getattr(exc, "response", None)
85
+ status = getattr(response, "status_code", None)
86
+ if status in (401, 403):
87
+ return DFloatAccessError(f"{repo_id}: the Hub refused access ({status}); {hint}")
88
+ return exc
89
+
90
+
91
+ def _snapshot_download(*, repo_id: str, allow_patterns: list[str]) -> str:
92
+ """``huggingface_hub.snapshot_download`` behind one name a test can replace."""
93
+ from huggingface_hub import snapshot_download
94
+
95
+ return str(snapshot_download(repo_id=repo_id, allow_patterns=allow_patterns))
96
+
97
+
98
+ def resolve(spec: str, *, patterns: tuple[str, ...]) -> ResolvedRepo:
99
+ """A local directory as is, or a Hub id fetched (or found in the cache) with ``patterns`` only.
100
+
101
+ Raises:
102
+ DFloatFormatError: Neither a directory nor a Hub id, or the Hub has no such repository.
103
+ DFloatAccessError: The repository is gated or this account may not read it.
104
+ """
105
+ if not is_hub_id(spec):
106
+ root = Path(spec).expanduser()
107
+ if not root.is_dir():
108
+ raise DFloatFormatError(f"{spec}: not a directory or a Hub repository id")
109
+ return ResolvedRepo(root=root, repo_id=None, revision=hub_revision(root))
110
+ try:
111
+ root = Path(_snapshot_download(repo_id=spec, allow_patterns=list(patterns)))
112
+ except Exception as exc:
113
+ translated = hub_error(exc, spec)
114
+ if translated is exc:
115
+ raise
116
+ raise translated from exc
117
+ return ResolvedRepo(root=root, repo_id=spec, revision=hub_revision(root))
118
+
119
+
120
+ def refuse_quantized(weights: Any, root: Path) -> None:
121
+ """Refuse an mflux-saved quantized base: the applier would honour its stored bits even with ``quantize_arg=None``.
122
+
123
+ Raises:
124
+ DFloatFormatError: The weights carry a quantization level.
125
+ """
126
+ level = weights.meta_data.quantization_level
127
+ if level is not None:
128
+ raise DFloatFormatError(
129
+ f"{root}: an mflux-saved {level}-bit model; the base must hold the BF16 encoders and VAE"
130
+ )
131
+
132
+
133
+ def _definition(names: tuple[str, ...], patterns: tuple[str, ...]) -> type:
134
+ """An mflux weight definition holding only the named FLUX.1 components (mflux's own ``ComponentDefinition``s)."""
135
+ require_mflux()
136
+ from mflux.models.flux.weights.flux_weight_definition import FluxWeightDefinition
137
+
138
+ components = [c for c in FluxWeightDefinition.get_components() if c.name in names]
139
+ if [c.name for c in components] != list(names):
140
+ raise DFloatIntegrationError(f"mflux's FLUX.1 weight definition does not define {names}")
141
+ return type(
142
+ "DFloatFluxBaseDefinition",
143
+ (),
144
+ {
145
+ "get_components": staticmethod(lambda: list(components)),
146
+ "get_tokenizers": staticmethod(FluxWeightDefinition.get_tokenizers),
147
+ "get_download_patterns": staticmethod(lambda: list(patterns)),
148
+ "quantization_predicate": staticmethod(FluxWeightDefinition.quantization_predicate),
149
+ },
150
+ )
151
+
152
+
153
+ def encoders_definition() -> type:
154
+ """The T5 and CLIP components of mflux's FLUX.1 definition, with their download patterns."""
155
+ return _definition(("t5_encoder", "clip_encoder"), ENCODER_PATTERNS)
156
+
157
+
158
+ def vae_definition() -> type:
159
+ """The VAE component of mflux's FLUX.1 definition, with its download patterns."""
160
+ return _definition(("vae",), VAE_PATTERNS)
161
+
162
+
163
+ def _load_into(root: Path, definition: Any, models: dict[str, Any]) -> None:
164
+ """Mflux's loader and applier over ``definition`` into ``models``; the weights stay lazy until evaluated.
165
+
166
+ The ``LoadedWeights`` object is not kept: it references every lazily loaded array (9.5 GB of
167
+ T5), and the encoders must be droppable.
168
+
169
+ Raises:
170
+ DFloatFormatError: The directory lacks the components (a checkpoint or model repository
171
+ without them), or holds an mflux-saved quantized model.
172
+ """
173
+ from mflux.models.common.weights.loading.weight_applier import WeightApplier
174
+ from mflux.models.common.weights.loading.weight_loader import WeightLoader
175
+
176
+ names = [c.name for c in definition.get_components()]
177
+ try:
178
+ weights = WeightLoader.load(
179
+ weight_definition=definition,
180
+ model_path=str(root),
181
+ download_patterns=definition.get_download_patterns(),
182
+ )
183
+ except (FileNotFoundError, ValueError) as exc:
184
+ raise DFloatFormatError(
185
+ f"{root}: not a FLUX.1 base with {names} (a DFloat11 checkpoint or model repository?): {exc}"
186
+ ) from exc
187
+ refuse_quantized(weights, root)
188
+ bits = WeightApplier.apply_and_quantize(
189
+ weights=weights, models=models, quantize_arg=None, weight_definition=definition
190
+ )
191
+ if bits is not None:
192
+ raise DFloatIntegrationError(
193
+ f"{root}: the applier quantized {names} to {bits} bits unasked"
194
+ )
195
+
196
+
197
+ def load_encoders(root: Path) -> tuple[Any, Any]:
198
+ """Fresh ``T5Encoder`` and ``CLIPEncoder`` modules with the base's weights (lazy)."""
199
+ require_mflux()
200
+ from mflux.models.flux.model.flux_text_encoder.clip_encoder.clip_encoder import CLIPEncoder
201
+ from mflux.models.flux.model.flux_text_encoder.t5_encoder.t5_encoder import T5Encoder
202
+
203
+ t5, clip = T5Encoder(), CLIPEncoder()
204
+ _load_into(root, encoders_definition(), {"t5_encoder": t5, "clip_encoder": clip})
205
+ return t5, clip
206
+
207
+
208
+ def load_vae(root: Path) -> Any:
209
+ """A fresh ``VAE`` module with the base's weights (lazy)."""
210
+ require_mflux()
211
+ from mflux.models.flux.model.flux_vae.vae import VAE
212
+
213
+ vae = VAE()
214
+ _load_into(root, vae_definition(), {"vae": vae})
215
+ return vae
216
+
217
+
218
+ def load_tokenizers(root: Path, model_config: Any) -> dict[str, Any]:
219
+ """The CLIP and T5 tokenizers; T5's length is the model's (256 schnell, 512 dev and Krea-dev).
220
+
221
+ Raises:
222
+ DFloatFormatError: The tokenizer files are missing.
223
+ """
224
+ require_mflux()
225
+ from mflux.models.common.tokenizer.tokenizer_loader import TokenizerLoader
226
+ from mflux.models.flux.weights.flux_weight_definition import FluxWeightDefinition
227
+
228
+ try:
229
+ return dict(
230
+ TokenizerLoader.load_all(
231
+ definitions=FluxWeightDefinition.get_tokenizers(),
232
+ model_path=str(root),
233
+ max_length_overrides={"t5": int(model_config.max_sequence_length)},
234
+ )
235
+ )
236
+ except (FileNotFoundError, RuntimeError) as exc:
237
+ raise DFloatFormatError(f"{root}: no usable FLUX.1 tokenizers: {exc}") from exc
238
+
239
+
240
+ def load_base(root: Path, model_config: Any) -> BaseComponents:
241
+ """VAE, encoders and tokenizers of a base repository (everything lazy; nothing read from disk yet)."""
242
+ t5, clip = load_encoders(root)
243
+ return BaseComponents(
244
+ vae=load_vae(root), t5=t5, clip=clip, tokenizers=load_tokenizers(root, model_config)
245
+ )