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,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])
|