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.
- mlx_dfloat/__init__.py +25 -0
- mlx_dfloat/_memory_caps.py +74 -0
- mlx_dfloat/_metal_decode.py +420 -0
- mlx_dfloat/_safetensors.py +185 -0
- mlx_dfloat/_scrub.py +29 -0
- mlx_dfloat/_version.py +24 -0
- mlx_dfloat/_watchdog.py +253 -0
- mlx_dfloat/bench/__init__.py +4 -0
- mlx_dfloat/bench/capped.py +165 -0
- mlx_dfloat/bench/preflight.py +216 -0
- mlx_dfloat/bench/results.py +227 -0
- mlx_dfloat/bench/scenario.py +191 -0
- mlx_dfloat/bench/table.py +251 -0
- mlx_dfloat/cli.py +30 -0
- mlx_dfloat/decode.py +120 -0
- mlx_dfloat/errors.py +45 -0
- mlx_dfloat/format.py +462 -0
- mlx_dfloat/integrate/__init__.py +1 -0
- mlx_dfloat/integrate/coverage.py +122 -0
- mlx_dfloat/integrate/memory.py +55 -0
- mlx_dfloat/integrate/names.py +121 -0
- mlx_dfloat/integrate/placeholders.py +72 -0
- mlx_dfloat/integrate/providers.py +271 -0
- mlx_dfloat/integrate/seam.py +196 -0
- mlx_dfloat/mflux/__init__.py +31 -0
- mlx_dfloat/mflux/flux1/__init__.py +1 -0
- mlx_dfloat/mflux/flux1/cli.py +419 -0
- mlx_dfloat/mflux/flux1/init.py +245 -0
- mlx_dfloat/mflux/flux1/lifecycle.py +159 -0
- mlx_dfloat/mflux/flux1/memory.py +201 -0
- mlx_dfloat/mflux/flux1/model.py +553 -0
- mlx_dfloat/mflux/flux1/names.py +79 -0
- mlx_dfloat/mflux/flux1/transformer.py +240 -0
- mlx_dfloat/py.typed +0 -0
- mlx_dfloat/reference.py +241 -0
- mlx_dfloat-0.1.0.dist-info/METADATA +325 -0
- mlx_dfloat-0.1.0.dist-info/RECORD +41 -0
- mlx_dfloat-0.1.0.dist-info/WHEEL +4 -0
- mlx_dfloat-0.1.0.dist-info/entry_points.txt +2 -0
- mlx_dfloat-0.1.0.dist-info/licenses/LICENSE +202 -0
- 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
|
+
)
|