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