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,159 @@
|
|
|
1
|
+
"""The encoder / compressed-set lifecycle: the text encoders and the compressed transformer are never resident together.
|
|
2
|
+
|
|
3
|
+
On a 32 GB Mac the T5 encoder (8.9 GiB) and the compressed FLUX.1 set (15.2 GiB) do not fit under the
|
|
4
|
+
fit rule, so a prompt is encoded first, its embeddings evaluated, the encoders dropped, and only then
|
|
5
|
+
the set loaded. A new prompt after the set is resident drops the set, reloads the encoders, encodes,
|
|
6
|
+
drops them and reloads the set. The object here holds no arrays and no modules: the callbacks own
|
|
7
|
+
them, so a drop releases them.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
import gc
|
|
11
|
+
import logging
|
|
12
|
+
import time
|
|
13
|
+
import traceback
|
|
14
|
+
from collections.abc import Callable
|
|
15
|
+
from dataclasses import asdict, dataclass
|
|
16
|
+
|
|
17
|
+
import mlx.core as mx
|
|
18
|
+
|
|
19
|
+
from mlx_dfloat.errors import DFloatResourceError
|
|
20
|
+
|
|
21
|
+
log = logging.getLogger("mlx_dfloat.mflux.flux1")
|
|
22
|
+
_eval = mx.eval # looked up on the module, so a test can record the evaluation
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
@dataclass(slots=True)
|
|
26
|
+
class LifecycleCounters:
|
|
27
|
+
"""Reloads and their seconds, for the report."""
|
|
28
|
+
|
|
29
|
+
encoder_loads: int = 0
|
|
30
|
+
encoder_seconds: float = 0.0
|
|
31
|
+
set_loads: int = 0
|
|
32
|
+
set_seconds: float = 0.0
|
|
33
|
+
forced_set_drops: int = 0
|
|
34
|
+
|
|
35
|
+
def as_dict(self) -> dict[str, float | int]:
|
|
36
|
+
"""The counters as a JSON-ready dict."""
|
|
37
|
+
return asdict(self)
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
class Lifecycle:
|
|
41
|
+
"""Drives the encoder and set callbacks so that the two are never resident at once."""
|
|
42
|
+
|
|
43
|
+
def __init__(
|
|
44
|
+
self,
|
|
45
|
+
*,
|
|
46
|
+
load_encoders: Callable[[], None],
|
|
47
|
+
unload_encoders: Callable[[], None],
|
|
48
|
+
encode: Callable[[str], tuple[mx.array, mx.array]],
|
|
49
|
+
load_set: Callable[[], None],
|
|
50
|
+
unload_set: Callable[[], None],
|
|
51
|
+
prompt_cache: dict[str, tuple[mx.array, mx.array]],
|
|
52
|
+
retained_bound: Callable[[], int],
|
|
53
|
+
encoders_loaded: bool = False,
|
|
54
|
+
) -> None:
|
|
55
|
+
"""Bind the callbacks, the prompt cache mflux reads, and the active-memory bound a drop must reach.
|
|
56
|
+
|
|
57
|
+
``retained_bound`` is called at each check, so what legitimately stays resident (cached
|
|
58
|
+
embeddings, a decoded VAE) can grow between drops.
|
|
59
|
+
"""
|
|
60
|
+
self._load_encoders = load_encoders
|
|
61
|
+
self._unload_encoders = unload_encoders
|
|
62
|
+
self._encode = encode
|
|
63
|
+
self._load_set = load_set
|
|
64
|
+
self._unload_set = unload_set
|
|
65
|
+
self._cache = prompt_cache
|
|
66
|
+
self.retained_bound = retained_bound
|
|
67
|
+
self.encoders_loaded = encoders_loaded
|
|
68
|
+
self.set_resident = False
|
|
69
|
+
self.counters = LifecycleCounters()
|
|
70
|
+
|
|
71
|
+
def ensure_embeddings(self, *prompts: str) -> None:
|
|
72
|
+
"""Encode the prompts not yet cached: the set is dropped first, the encoders after.
|
|
73
|
+
|
|
74
|
+
Each pair is evaluated before it is cached: a lazy pair would keep every encoder weight
|
|
75
|
+
alive past the drop. A failure in the encode or the evaluation still drops the encoders
|
|
76
|
+
(instead of leaving them resident next to a set loaded later) and re-raises the original
|
|
77
|
+
error; a memory check that fails on that path is logged, never raised over it.
|
|
78
|
+
"""
|
|
79
|
+
missing = [p for p in prompts if p not in self._cache]
|
|
80
|
+
if not missing:
|
|
81
|
+
return
|
|
82
|
+
if self.set_resident:
|
|
83
|
+
log.info(
|
|
84
|
+
"new prompt with the compressed set resident: dropping the set to load the text encoders"
|
|
85
|
+
)
|
|
86
|
+
self.counters.forced_set_drops += 1
|
|
87
|
+
self.drop_set()
|
|
88
|
+
if not self.encoders_loaded:
|
|
89
|
+
start = time.perf_counter()
|
|
90
|
+
self._load_encoders()
|
|
91
|
+
self.encoders_loaded = True
|
|
92
|
+
self.counters.encoder_loads += 1
|
|
93
|
+
self.counters.encoder_seconds += time.perf_counter() - start
|
|
94
|
+
pair: tuple[mx.array, mx.array] | None = None
|
|
95
|
+
try:
|
|
96
|
+
for prompt in missing:
|
|
97
|
+
pair = self._encode(prompt)
|
|
98
|
+
_eval(*pair)
|
|
99
|
+
self._cache[prompt] = pair
|
|
100
|
+
except BaseException as exc:
|
|
101
|
+
# The traceback keeps the failing frames alive, and with them a lazy pair whose graph
|
|
102
|
+
# holds the encoder weights; clear them before the memory check.
|
|
103
|
+
pair = None
|
|
104
|
+
traceback.clear_frames(exc.__traceback__)
|
|
105
|
+
self._unload_encoders()
|
|
106
|
+
self.encoders_loaded = False
|
|
107
|
+
try:
|
|
108
|
+
self._reclaim("text encoders")
|
|
109
|
+
except DFloatResourceError as reclaim_error:
|
|
110
|
+
log.warning("%s (while handling %s: %s)", reclaim_error, type(exc).__name__, exc)
|
|
111
|
+
raise
|
|
112
|
+
del pair
|
|
113
|
+
self.drop_encoders()
|
|
114
|
+
|
|
115
|
+
def ensure_set(self) -> None:
|
|
116
|
+
"""Load and attach the compressed set unless it is resident.
|
|
117
|
+
|
|
118
|
+
Raises:
|
|
119
|
+
DFloatResourceError: MLX memory above the retained bound is still active before the
|
|
120
|
+
load (a previous set not released), so a second copy is never loaded next to it.
|
|
121
|
+
"""
|
|
122
|
+
if self.set_resident:
|
|
123
|
+
return
|
|
124
|
+
gc.collect()
|
|
125
|
+
active = int(mx.get_active_memory())
|
|
126
|
+
bound = self.retained_bound()
|
|
127
|
+
if active > bound:
|
|
128
|
+
raise DFloatResourceError(
|
|
129
|
+
f"a previous set is still resident: {active / 1024**3:.2f} GiB of MLX memory is active "
|
|
130
|
+
f"before the load (bound {bound / 1024**3:.2f} GiB)"
|
|
131
|
+
)
|
|
132
|
+
start = time.perf_counter()
|
|
133
|
+
self._load_set()
|
|
134
|
+
self.set_resident = True
|
|
135
|
+
self.counters.set_loads += 1
|
|
136
|
+
self.counters.set_seconds += time.perf_counter() - start
|
|
137
|
+
|
|
138
|
+
def drop_set(self) -> None:
|
|
139
|
+
"""Detach and release the compressed set; assert the memory came back."""
|
|
140
|
+
self._unload_set()
|
|
141
|
+
self.set_resident = False
|
|
142
|
+
self._reclaim("compressed set")
|
|
143
|
+
|
|
144
|
+
def drop_encoders(self) -> None:
|
|
145
|
+
"""Release the text encoders; assert the memory came back."""
|
|
146
|
+
self._unload_encoders()
|
|
147
|
+
self.encoders_loaded = False
|
|
148
|
+
self._reclaim("text encoders")
|
|
149
|
+
|
|
150
|
+
def _reclaim(self, what: str) -> None:
|
|
151
|
+
gc.collect()
|
|
152
|
+
mx.clear_cache()
|
|
153
|
+
active = int(mx.get_active_memory())
|
|
154
|
+
bound = self.retained_bound()
|
|
155
|
+
if active > bound:
|
|
156
|
+
raise DFloatResourceError(
|
|
157
|
+
f"after dropping the {what}, {active / 1024**3:.2f} GiB of MLX memory is still active "
|
|
158
|
+
f"(bound {bound / 1024**3:.2f} GiB): something else holds it"
|
|
159
|
+
)
|
|
@@ -0,0 +1,201 @@
|
|
|
1
|
+
"""FLUX.1 memory rules: the cache limit and the fit phases. Measured at schnell 1024² only; elsewhere predicted."""
|
|
2
|
+
|
|
3
|
+
from collections.abc import Mapping
|
|
4
|
+
from dataclasses import dataclass
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
from typing import Any
|
|
7
|
+
|
|
8
|
+
from mlx_dfloat.format import DF11Checkpoint
|
|
9
|
+
from mlx_dfloat.integrate.memory import FitEstimate, fit_estimate
|
|
10
|
+
from mlx_dfloat.mflux.flux1.names import DOUBLE_PREFIX, SINGLE_PREFIX
|
|
11
|
+
|
|
12
|
+
ALLOWANCE_AT_REFERENCE = (
|
|
13
|
+
1_500_000_000 # recycled activation volume measured at 1024², 256 text tokens
|
|
14
|
+
)
|
|
15
|
+
REFERENCE_TOKENS = 4096 + 256
|
|
16
|
+
ALLOWANCE_FLOOR = 500_000_000
|
|
17
|
+
MEASURED_LIMIT_AT_1024 = 2_500_000_000
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def text_tokens(model_config: Any) -> int:
|
|
21
|
+
"""The model's text sequence length (256 for schnell, 512 for dev and Krea-dev)."""
|
|
22
|
+
return int(model_config.max_sequence_length)
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def activation_allowance(*, height: int, width: int, text_tokens: int) -> int:
|
|
26
|
+
"""Cache room for the activation buffers a block frees, scaled linearly with the token count."""
|
|
27
|
+
tokens = height * width // 256 + text_tokens
|
|
28
|
+
return max(ALLOWANCE_FLOOR, int(ALLOWANCE_AT_REFERENCE * tokens / REFERENCE_TOKENS))
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def cache_limit_for(
|
|
32
|
+
largest: Mapping[str, int],
|
|
33
|
+
*,
|
|
34
|
+
policy: str,
|
|
35
|
+
height: int,
|
|
36
|
+
width: int,
|
|
37
|
+
text_tokens: int,
|
|
38
|
+
override: int | None = None,
|
|
39
|
+
) -> int:
|
|
40
|
+
"""The MLX buffer-cache limit for one generate call: room for a decoded buffer of each kind next to the activations."""
|
|
41
|
+
if override is not None:
|
|
42
|
+
return override
|
|
43
|
+
limit = (
|
|
44
|
+
largest[DOUBLE_PREFIX]
|
|
45
|
+
+ largest[SINGLE_PREFIX]
|
|
46
|
+
+ activation_allowance(height=height, width=width, text_tokens=text_tokens)
|
|
47
|
+
)
|
|
48
|
+
if policy == "depth2":
|
|
49
|
+
limit += max(largest.values())
|
|
50
|
+
if height * width >= 1024 * 1024:
|
|
51
|
+
limit = max(limit, MEASURED_LIMIT_AT_1024)
|
|
52
|
+
return limit
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def flux_phases(
|
|
56
|
+
*,
|
|
57
|
+
compressed_bytes: int,
|
|
58
|
+
extras_bytes: int,
|
|
59
|
+
largest: Mapping[str, int],
|
|
60
|
+
policy: str,
|
|
61
|
+
cache_limit: int,
|
|
62
|
+
allowance: int,
|
|
63
|
+
encoders_bytes: int,
|
|
64
|
+
vae_bytes: int,
|
|
65
|
+
vae_transient_bytes: int,
|
|
66
|
+
overhead_bytes: int,
|
|
67
|
+
denoise_activation_bytes: int = 0,
|
|
68
|
+
) -> dict[str, dict[str, int]]:
|
|
69
|
+
"""The three phases of a generation, each term once."""
|
|
70
|
+
in_flight = max(largest.values()) * (2 if policy == "depth2" else 1)
|
|
71
|
+
return {
|
|
72
|
+
"encode": {
|
|
73
|
+
"encoders": encoders_bytes,
|
|
74
|
+
"activations": allowance,
|
|
75
|
+
"overhead": overhead_bytes,
|
|
76
|
+
},
|
|
77
|
+
"denoise": {
|
|
78
|
+
"compressed": compressed_bytes,
|
|
79
|
+
"extras": extras_bytes,
|
|
80
|
+
"decoded": in_flight,
|
|
81
|
+
"cache": cache_limit,
|
|
82
|
+
"activations": denoise_activation_bytes,
|
|
83
|
+
"overhead": overhead_bytes,
|
|
84
|
+
},
|
|
85
|
+
"vae": {
|
|
86
|
+
"compressed": compressed_bytes,
|
|
87
|
+
"extras": extras_bytes,
|
|
88
|
+
"vae": vae_bytes,
|
|
89
|
+
"transient": vae_transient_bytes,
|
|
90
|
+
"overhead": overhead_bytes,
|
|
91
|
+
},
|
|
92
|
+
}
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
# Measured on 2026-09-28 (schnell 1024², one process: encode, drop, set, one step, VAE decode; M1 Max 32 GB,
|
|
96
|
+
# macOS 27.0, mlx 0.32.2, mflux 0.20.0, git 4601a8b). OVERHEAD is the footprint that MLX's counters do not see
|
|
97
|
+
# after construction (Python, torch and transformers imports, the runtime); VAE_TRANSIENT is what the float32
|
|
98
|
+
# decode adds at 1024² on top of everything resident after the set load (the VAE phase peaked at 23.29 GiB with
|
|
99
|
+
# the set resident, over the 23.0 GiB fit rule, which is why a call may drop the set before decoding).
|
|
100
|
+
# Re-measure before changing either.
|
|
101
|
+
OVERHEAD_BYTES = 247_712_510 # 0.23 GiB: build-phase footprint minus MLX active and cache
|
|
102
|
+
VAE_TRANSIENT_BYTES = (
|
|
103
|
+
8_392_982_528 # 7.82 GiB: VAE-phase footprint peak minus the footprint after the set load
|
|
104
|
+
)
|
|
105
|
+
# Measured on 2026-09-28 (the `mlx-dfloat generate` runs, schnell 1024², 4 steps, per-block, 2.5 GB cache limit;
|
|
106
|
+
# M1 Max 32 GB, mlx 0.32.2, mflux 0.20.0, git 2fe3b2c): the process footprint peaked at 19.94 GiB against a
|
|
107
|
+
# denoise estimate of 18.38 GiB without this term. The 1.56 GiB gap is the activation volume the step holds
|
|
108
|
+
# beyond the cache limit, at 4096 image + 256 text tokens; dev (512 text tokens) measured 20.10 GiB.
|
|
109
|
+
DENOISE_ACTIVATION_AT_REFERENCE = 1_675_000_000
|
|
110
|
+
MAX_MEASURED_PIXELS = 1024 * 1024 # no run above 1024² on this path
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
def vae_transient_bytes(*, height: int, width: int) -> int:
|
|
114
|
+
"""What the float32 VAE decode adds on top of what is resident: the measured 1024² value as a floor.
|
|
115
|
+
|
|
116
|
+
At or below 1024² this is the measured value (smaller images were not measured lower, so the
|
|
117
|
+
floor stays). Above 1024² it grows with the pixel count: a prediction, not a measurement.
|
|
118
|
+
"""
|
|
119
|
+
return int(VAE_TRANSIENT_BYTES * max(1.0, (height * width) / (1024 * 1024)))
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
def denoise_activation_bytes(*, height: int, width: int, text_tokens: int) -> int:
|
|
123
|
+
"""The activation volume a denoise step holds beyond the cache limit, linear in the token count.
|
|
124
|
+
|
|
125
|
+
Calibrated on the measured schnell 1024² peak (4096 image + 256 text tokens); every other size
|
|
126
|
+
is a prediction.
|
|
127
|
+
"""
|
|
128
|
+
tokens = height * width // 256 + text_tokens
|
|
129
|
+
return int(DENOISE_ACTIVATION_AT_REFERENCE * tokens / REFERENCE_TOKENS)
|
|
130
|
+
|
|
131
|
+
|
|
132
|
+
@dataclass(frozen=True, slots=True, kw_only=True)
|
|
133
|
+
class FluxSizes:
|
|
134
|
+
"""Resident bytes of the components, from file sizes (exact) and the extras' tensor sizes."""
|
|
135
|
+
|
|
136
|
+
compressed: int
|
|
137
|
+
extras: int
|
|
138
|
+
encoders: int
|
|
139
|
+
vae: int
|
|
140
|
+
|
|
141
|
+
|
|
142
|
+
def safetensors_bytes(root: Path, *subdirs: str) -> int:
|
|
143
|
+
"""The size of every ``*.safetensors`` file directly under each ``root/subdir`` (missing subdirs count zero)."""
|
|
144
|
+
return sum(
|
|
145
|
+
p.stat().st_size
|
|
146
|
+
for sub in subdirs
|
|
147
|
+
for p in (root / sub).glob("*.safetensors")
|
|
148
|
+
if p.is_file()
|
|
149
|
+
)
|
|
150
|
+
|
|
151
|
+
|
|
152
|
+
def sizes_for(ckpt: DF11Checkpoint, base_root: Path) -> FluxSizes:
|
|
153
|
+
"""Sizes of a DF11 checkpoint (compressed set apart from its extras) and of a base's encoders and VAE."""
|
|
154
|
+
total = sum(p.stat().st_size for p in ckpt.root.glob("*.safetensors") if p.is_file())
|
|
155
|
+
extras = sum(info.nbytes for _path, info in ckpt.extras.values())
|
|
156
|
+
return FluxSizes(
|
|
157
|
+
compressed=total - extras,
|
|
158
|
+
extras=extras,
|
|
159
|
+
encoders=safetensors_bytes(base_root, "text_encoder", "text_encoder_2"),
|
|
160
|
+
vae=safetensors_bytes(base_root, "vae"),
|
|
161
|
+
)
|
|
162
|
+
|
|
163
|
+
|
|
164
|
+
def fit_for(
|
|
165
|
+
*,
|
|
166
|
+
sizes: FluxSizes,
|
|
167
|
+
largest: Mapping[str, int],
|
|
168
|
+
policy: str,
|
|
169
|
+
cache_limit: int,
|
|
170
|
+
allowance: int,
|
|
171
|
+
budget: int,
|
|
172
|
+
height: int,
|
|
173
|
+
width: int,
|
|
174
|
+
text_tokens: int,
|
|
175
|
+
vae_with_set: bool = True,
|
|
176
|
+
) -> FitEstimate:
|
|
177
|
+
"""The phase estimate for one generate call against ``budget`` (a prediction, labelled as such by the caller).
|
|
178
|
+
|
|
179
|
+
``height`` and ``width`` scale the VAE transient (above 1024²) and, with ``text_tokens``, the
|
|
180
|
+
denoise activation term. ``vae_with_set=False`` plans the VAE phase after the compressed set has
|
|
181
|
+
been dropped (what a call does when the resident variant would not fit).
|
|
182
|
+
"""
|
|
183
|
+
phases = flux_phases(
|
|
184
|
+
compressed_bytes=sizes.compressed,
|
|
185
|
+
extras_bytes=sizes.extras,
|
|
186
|
+
largest=largest,
|
|
187
|
+
policy=policy,
|
|
188
|
+
cache_limit=cache_limit,
|
|
189
|
+
allowance=allowance,
|
|
190
|
+
encoders_bytes=sizes.encoders,
|
|
191
|
+
vae_bytes=sizes.vae,
|
|
192
|
+
vae_transient_bytes=vae_transient_bytes(height=height, width=width),
|
|
193
|
+
overhead_bytes=OVERHEAD_BYTES,
|
|
194
|
+
denoise_activation_bytes=denoise_activation_bytes(
|
|
195
|
+
height=height, width=width, text_tokens=text_tokens
|
|
196
|
+
),
|
|
197
|
+
)
|
|
198
|
+
if not vae_with_set:
|
|
199
|
+
phases["vae"].pop("compressed")
|
|
200
|
+
phases["vae"].pop("extras")
|
|
201
|
+
return fit_estimate(phases, budget_bytes=budget)
|