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
mlx_dfloat/__init__.py ADDED
@@ -0,0 +1,25 @@
1
+ """Run DFloat11 losslessly compressed BF16 models on Apple Silicon with MLX."""
2
+
3
+ from mlx_dfloat._version import __version__
4
+ from mlx_dfloat.errors import (
5
+ DFloatAccessError,
6
+ DFloatBackendError,
7
+ DFloatDependencyError,
8
+ DFloatError,
9
+ DFloatFormatError,
10
+ DFloatIntegrationError,
11
+ DFloatResourceError,
12
+ DFloatUnsupportedError,
13
+ )
14
+
15
+ __all__ = [
16
+ "DFloatAccessError",
17
+ "DFloatBackendError",
18
+ "DFloatDependencyError",
19
+ "DFloatError",
20
+ "DFloatFormatError",
21
+ "DFloatIntegrationError",
22
+ "DFloatResourceError",
23
+ "DFloatUnsupportedError",
24
+ "__version__",
25
+ ]
@@ -0,0 +1,74 @@
1
+ """Hardware-aware MLX memory caps (kernel-watchdog panic guard).
2
+
3
+ Derives wired + memory caps from the device's reported working-set size and
4
+ clamps strictly below it. Returns (0, 0) as a no-op signal on devices/CI images
5
+ that report no working-set size. Mirrors the mlx-taef and mlx-quant-fidelity pattern.
6
+ """
7
+
8
+ import mlx.core as mx
9
+
10
+ DESIRED_WIRED_GB = 20
11
+ DESIRED_MEMORY_GB = 22
12
+ HEADROOM_GB = 2
13
+
14
+
15
+ def _clamp_caps_gb(max_recommended_gb: int) -> tuple[int, int]:
16
+ """Clamp the desired caps to fit a device with `max_recommended_gb` working set."""
17
+ if max_recommended_gb <= 0:
18
+ return (0, 0)
19
+ wired_gb = min(DESIRED_WIRED_GB, max(1, max_recommended_gb - HEADROOM_GB))
20
+ memory_gb = min(DESIRED_MEMORY_GB, max(wired_gb + 1, max_recommended_gb))
21
+ return (wired_gb, memory_gb)
22
+
23
+
24
+ def compute_safe_caps_gb() -> tuple[int, int]:
25
+ """Return (wired_gb, memory_gb) that fit the current device, or (0, 0)."""
26
+ try:
27
+ info = mx.device_info()
28
+ max_gb = int(info.get("max_recommended_working_set_size", 0)) // (1024**3)
29
+ except Exception:
30
+ return (0, 0)
31
+ return _clamp_caps_gb(max_gb)
32
+
33
+
34
+ def device_string() -> str | None:
35
+ """Human-readable chip + unified-memory size (e.g. 'Apple M1 Max, 32 GB'), or None.
36
+
37
+ Report provenance only — never used for gating. Returns None when MLX does not
38
+ report a device name (CI containers, future backends).
39
+ """
40
+ try:
41
+ info = mx.device_info()
42
+ name = info.get("device_name")
43
+ if not isinstance(name, str) or not name:
44
+ return None
45
+ mem = info.get("memory_size")
46
+ if isinstance(mem, int) and mem > 0:
47
+ return f"{name}, {round(mem / 1024**3)} GB"
48
+ return name
49
+ except Exception:
50
+ return None
51
+
52
+
53
+ def install_memory_caps() -> tuple[int, int]:
54
+ """Apply wired + memory caps for the current device. Idempotent; never raises.
55
+
56
+ Returns the (wired_gb, memory_gb) actually installed; each is 0 when that cap could not be
57
+ applied, and both are 0 on a device with no reported working-set size. The two caps are
58
+ applied independently, so a failed wired cap still leaves the memory cap in place.
59
+ """
60
+ wired_gb, memory_gb = compute_safe_caps_gb()
61
+ if wired_gb == 0:
62
+ return (0, 0)
63
+ try:
64
+ mx.set_wired_limit(wired_gb * 1024**3)
65
+ except Exception:
66
+ wired_gb = 0
67
+ try:
68
+ mx.set_memory_limit(memory_gb * 1024**3)
69
+ except Exception:
70
+ memory_gb = 0
71
+ return (wired_gb, memory_gb)
72
+
73
+
74
+ __all__ = ["compute_safe_caps_gb", "device_string", "install_memory_caps"]
@@ -0,0 +1,420 @@
1
+ """The Metal DF11 decode kernel: a per-thread, bounded, status-reporting mirror of the reference algorithm."""
2
+
3
+ from dataclasses import dataclass
4
+ from typing import Any
5
+
6
+ import mlx.core as mx
7
+ import numpy as np
8
+ import numpy.typing as npt
9
+
10
+ from mlx_dfloat.decode import DecodeResult
11
+ from mlx_dfloat.errors import DFloatBackendError
12
+ from mlx_dfloat.format import GroupArrays, MxGroup
13
+
14
+ THREADS = 512
15
+ CAP = 16192
16
+ _SCRATCH_BYTES = 4 * (16 + 16 + 1 + 16)
17
+
18
+
19
+ def _metal_static_bytes(declared: int) -> int:
20
+ """The static threadgroup memory Metal reserves for ``declared`` bytes of arrays.
21
+
22
+ Metal reports ``staticThreadgroupMemoryLength`` at 16-byte granularity (measured with
23
+ standalone compiles: 196 -> 208, 192 -> 192; mlx 0.32.2, macOS 27.0, M1 Max).
24
+ """
25
+ return -(-declared // 16) * 16
26
+
27
+
28
+ THREADGROUP_BYTES_STAGED = _metal_static_bytes(
29
+ 2 * CAP + _SCRATCH_BYTES
30
+ ) # 32,580 -> 32,592 < 32,768
31
+ THREADGROUP_BYTES_DIRECT = _metal_static_bytes(2 + _SCRATCH_BYTES) # 198 -> 208
32
+
33
+ _KERNEL: Any | None = None
34
+ # (force_direct, poison_buf, unguarded_gap_read) -> warmed
35
+ _PIPELINES: dict[tuple[bool, bool, bool], bool] = {}
36
+
37
+ _SOURCE = r"""
38
+ const uint t = thread_position_in_threadgroup.x;
39
+ const uint b = threadgroup_position_in_grid.x;
40
+ const uint g = b * 512u + t;
41
+ const uint n_bytes = (uint)encoded_shape[0];
42
+ const uint n = (uint)sm_shape[0];
43
+ const uint n_luts = (uint)luts_shape[0];
44
+ const uint n_launch = (uint)positions_shape[0] - 1u;
45
+ const uint blk_start = positions[b];
46
+ const uint blk_end = positions[b + 1u];
47
+ const uint interval = blk_end - blk_start;
48
+ const bool staged = (!FORCE_DIRECT) && (interval <= (uint)CAP);
49
+ const uint last_real = (n_bytes - 1u) / 8u;
50
+ const bool real = (8u * g) < n_bytes;
51
+
52
+ threadgroup ushort buf[FORCE_DIRECT ? 1 : CAP];
53
+ threadgroup uint simd_totals[16];
54
+ threadgroup uint simd_pre[16];
55
+ threadgroup uint block_total[1];
56
+ threadgroup uint simd_bad[16];
57
+ if (POISON_BUF && staged) { for (uint i = t; i < (uint)CAP; i += 512u) buf[i] = (ushort)0xDEAD; }
58
+ threadgroup_barrier(mem_flags::mem_threadgroup);
59
+
60
+ const uint gbit = 5u * g;
61
+ const uint gbyte = gbit >> 3;
62
+ const uint gsh = gbit & 7u;
63
+ const uint g0 = (uint)gaps[gbyte];
64
+ const uint g1 = (UNGUARDED_GAP_READ || gsh > 3u) ? (uint)gaps[gbyte + 1u] : 0u;
65
+ const uint gap = (((g0 << 8) | g1) >> (11u - gsh)) & 31u;
66
+
67
+ ulong w = 0;
68
+ for (uint k = 0; k < 8u; ++k) {
69
+ const uint i = 8u * g + k;
70
+ w = (w << 8) | (ulong)((i < n_bytes) ? encoded[i] : (uchar)0);
71
+ }
72
+ uint la = 0;
73
+ for (uint k = 8u; k < 12u; ++k) {
74
+ const uint i = 8u * g + k;
75
+ la = (la << 8) | (uint)((i < n_bytes) ? encoded[i] : (uchar)0);
76
+ }
77
+
78
+ uint count = 0;
79
+ uint p = gap;
80
+ bool bad = false;
81
+ if (real) {
82
+ for (uint it = 0; it < 64u && p < 64u; ++it) {
83
+ const uint word = (p <= 32u) ? (uint)(w >> (32u - p)) : ((uint)(w << (p - 32u)) | (la >> (64u - p)));
84
+ uint sym = (uint)luts[word >> 24];
85
+ for (uint level = 1u; level <= 3u && sym >= 240u; ++level) {
86
+ const uint row = 256u - sym;
87
+ if (row < 1u || row > n_luts - 2u) { sym = 240u; break; }
88
+ sym = (uint)luts[row * 256u + ((word >> (24u - 8u * level)) & 0xFFu)];
89
+ }
90
+ const uint len = (sym < 240u) ? (uint)luts[(n_luts - 1u) * 256u + sym] : 0u;
91
+ if (len == 0u) { bad = true; break; }
92
+ count += 1u;
93
+ p += len;
94
+ }
95
+ }
96
+ const uint sg = t / 32u;
97
+ const uint lane = t % 32u;
98
+ const uint local = simd_prefix_exclusive_sum(count);
99
+
100
+ // Chain check (reference `broken`): a clean thread before the last real one must end exactly where the
101
+ // next thread's gap says its first code starts, unless its last code runs past the stream end.
102
+ bool broken = false;
103
+ if (real && !bad && g < last_real) {
104
+ const uint nbit = 5u * (g + 1u);
105
+ const uint nb0 = (uint)gaps[nbit >> 3];
106
+ const uint nsh = nbit & 7u;
107
+ const uint nb1 = (nsh > 3u) ? (uint)gaps[(nbit >> 3) + 1u] : 0u;
108
+ const uint next_gap = (((nb0 << 8) | nb1) >> (11u - nsh)) & 31u;
109
+ const ulong end_bit = 64ul * (ulong)g + (ulong)p;
110
+ broken = (end_bit < 8ul * (ulong)n_bytes) && (p != 64u + next_gap);
111
+ }
112
+ const bool relevant_bad = bad && (g < last_real);
113
+ const uint vote = (simd_any(relevant_bad) ? 1u : 0u) | (simd_any(broken) ? 4u : 0u); // uniform flow
114
+
115
+ if (lane == 31u) { simd_totals[sg] = local + count; simd_bad[sg] = vote; }
116
+ threadgroup_barrier(mem_flags::mem_threadgroup);
117
+ if (sg == 0u) {
118
+ const uint v = (lane < 16u) ? simd_totals[lane] : 0u;
119
+ const uint pre = simd_prefix_exclusive_sum(v);
120
+ if (lane < 16u) simd_pre[lane] = pre;
121
+ if (lane == 15u) block_total[0] = pre + v;
122
+ }
123
+ threadgroup_barrier(mem_flags::mem_threadgroup);
124
+ const uint first = blk_start + simd_pre[sg] + local;
125
+ const uint total = block_total[0];
126
+
127
+ if (real && count > 0u) {
128
+ uint idx = first;
129
+ const uint stop = min(min(first + count, n), blk_end);
130
+ uint q = gap;
131
+ for (uint it = 0; it < 64u && idx < stop; ++it) {
132
+ const uint word = (q <= 32u) ? (uint)(w >> (32u - q)) : ((uint)(w << (q - 32u)) | (la >> (64u - q)));
133
+ uint sym = (uint)luts[word >> 24];
134
+ for (uint level = 1u; level <= 3u && sym >= 240u; ++level) {
135
+ const uint row = 256u - sym;
136
+ if (row < 1u || row > n_luts - 2u) { sym = 240u; break; }
137
+ sym = (uint)luts[row * 256u + ((word >> (24u - 8u * level)) & 0xFFu)];
138
+ }
139
+ const uint len = (sym < 240u) ? (uint)luts[(n_luts - 1u) * 256u + sym] : 0u;
140
+ if (len == 0u) break;
141
+ const uint smv = (uint)sm[idx];
142
+ const ushort v = (ushort)(((smv & 0x80u) << 8) | (sym << 7) | (smv & 0x7Fu));
143
+ if (staged) buf[idx - blk_start] = v; else out[idx] = v;
144
+ idx += 1u;
145
+ q += len;
146
+ }
147
+ }
148
+ threadgroup_barrier(mem_flags::mem_threadgroup);
149
+ if (staged) {
150
+ const uint limit = min(interval, n - blk_start);
151
+ for (uint i = t; i < limit; i += 512u) out[blk_start + i] = buf[i];
152
+ }
153
+ if (t == 511u) {
154
+ uint flags = 0u;
155
+ for (uint s = 0; s < 16u; ++s) flags |= simd_bad[s];
156
+ uint word = flags; // 1 invalid code, 4 broken chain
157
+ const bool last_block = (b + 1u == n_launch);
158
+ if ((!last_block && total != interval) || (last_block && total < interval)) word |= 2u;
159
+ if (!staged) word |= 8u; // informational: direct path taken
160
+ status[b] = word;
161
+ }
162
+ """
163
+
164
+
165
+ def _build_kernel() -> Any:
166
+ global _KERNEL
167
+ if _KERNEL is None:
168
+ _KERNEL = mx.fast.metal_kernel(
169
+ name="df11_decode",
170
+ input_names=["encoded", "sm", "luts", "gaps", "positions"],
171
+ output_names=["out", "status"],
172
+ source=_SOURCE,
173
+ )
174
+ return _KERNEL
175
+
176
+
177
+ def _dispatch(
178
+ group: MxGroup,
179
+ *,
180
+ force_direct: bool,
181
+ poison_buf: bool,
182
+ init_value: int | None,
183
+ unguarded_gap_read: bool = False,
184
+ ) -> tuple[mx.array, mx.array]:
185
+ out, status = _build_kernel()(
186
+ inputs=[
187
+ group.encoded_exponent,
188
+ group.sign_mantissa,
189
+ group.luts,
190
+ group.gaps,
191
+ group.positions,
192
+ ],
193
+ template=[
194
+ ("CAP", CAP),
195
+ ("FORCE_DIRECT", force_direct),
196
+ ("POISON_BUF", poison_buf),
197
+ ("UNGUARDED_GAP_READ", unguarded_gap_read),
198
+ ],
199
+ grid=(THREADS * group.n_launch, 1, 1),
200
+ threadgroup=(THREADS, 1, 1),
201
+ output_shapes=[(group.n_elements,), (group.n_launch,)],
202
+ output_dtypes=[mx.uint16, mx.uint32],
203
+ init_value=init_value,
204
+ )
205
+ return out, status
206
+
207
+
208
+ # The warm-up groups. MLX binds a read-only input of fewer than 8 entries in the `constant` address
209
+ # space and names (and compiles) one pipeline per input-binding signature, so readiness must warm every
210
+ # signature a real group can present, not one of them:
211
+ #
212
+ # 1. `warmup-7-block`: seven full 4096-byte blocks, `positions` has eight entries and every input binds
213
+ # `device`, the signature of any group with at least seven launched blocks. Each thread's 8 bytes
214
+ # FF FF FF FF FF FF FF FE decode under the H1 tables (0xxxxxxx -> exponent 127, length 1;
215
+ # 10xxxxxx -> 126, length 2; 11xxxxxx -> 128, length 3) to 21 codes `111` (exponent 128) then one
216
+ # `0` (exponent 127): 22 elements per thread, 11,264 per block, under CAP, so the staged
217
+ # instantiation stages every block. An all-zero stream would decode to 64 elements per thread,
218
+ # 32,768 per block, over CAP, and never exercise staging.
219
+ # 2. `warmup-1-block`: one thread of that same stream (8 bytes, 22 elements), so `positions` has two
220
+ # entries and binds `constant` while every other input still binds `device`: the signature of any
221
+ # group with fewer than seven launched blocks (positions.size <= 7).
222
+ # 3. `warmup-tiny`: the H1 hand fixture (one byte 0x9B, three elements), so `encoded`, `sm` and
223
+ # `positions` all bind `constant`: the signature of a very small group.
224
+ _WARMUP_BLOCKS = 7
225
+ _WARMUP_CODES_PER_THREAD = 22
226
+ _WARMUP_ELEMENTS_PER_BLOCK = THREADS * _WARMUP_CODES_PER_THREAD
227
+ _WARMUP_THREAD_BYTES = (0xFF,) * 7 + (0xFE,)
228
+ _GAPS_BYTES_PER_BLOCK = 320 # 512 threads x 5 bits
229
+
230
+
231
+ @dataclass(frozen=True, slots=True, kw_only=True)
232
+ class _Warmup:
233
+ """One warm-up group: the signature it exercises, its host arrays and its hand-derived BF16 bits."""
234
+
235
+ name: str
236
+ arrays: GroupArrays
237
+ expected: npt.NDArray[np.uint16]
238
+
239
+
240
+ def _h1_luts() -> npt.NDArray[np.uint8]:
241
+ """H1: 127 `0`, 126 `10`, 128 `110` (EOF `111` inherits 128, as upstream); row 1 holds the lengths."""
242
+ row0 = np.zeros(256, np.uint8)
243
+ row0[0:128], row0[128:192], row0[192:256] = 127, 126, 128
244
+ lens = np.zeros(256, np.uint8)
245
+ lens[126], lens[127], lens[128] = 2, 1, 3
246
+ return np.stack([row0, lens])
247
+
248
+
249
+ def _thread_stream_arrays(n_threads: int, n_blocks: int, name: str) -> _Warmup:
250
+ """`n_threads` copies of the FF..FE thread stream over `n_blocks` blocks, zero gaps, sm[i] = i & 0xFF."""
251
+ n_elements = n_threads * _WARMUP_CODES_PER_THREAD
252
+ i = np.arange(n_elements, dtype=np.uint32)
253
+ exponent = np.where(i % _WARMUP_CODES_PER_THREAD == _WARMUP_CODES_PER_THREAD - 1, 127, 128)
254
+ sm = i & 0xFF
255
+ expected = (((sm & 0x80) << 8) | (exponent.astype(np.uint32) << 7) | (sm & 0x7F)).astype(
256
+ np.uint16
257
+ )
258
+ per_block = _WARMUP_ELEMENTS_PER_BLOCK if n_blocks > 1 else n_elements
259
+ arrays = GroupArrays(
260
+ encoded_exponent=np.tile(np.array(_WARMUP_THREAD_BYTES, np.uint8), n_threads),
261
+ sign_mantissa=sm.astype(np.uint8),
262
+ luts=_h1_luts(),
263
+ gaps=np.zeros(_GAPS_BYTES_PER_BLOCK * n_blocks, np.uint8),
264
+ output_positions=(np.arange(n_blocks + 1) * per_block).astype(np.uint32),
265
+ split_positions=np.zeros(0, np.int64),
266
+ )
267
+ return _Warmup(name=name, arrays=arrays, expected=expected)
268
+
269
+
270
+ def _warmups() -> tuple[_Warmup, ...]:
271
+ """The three warm-up groups, one per input-binding signature (see the note above)."""
272
+ tiny = GroupArrays(
273
+ encoded_exponent=np.array([0x9B], np.uint8), # 10 | 0 | 110 -> exponents 126, 127, 128
274
+ sign_mantissa=np.array([0x00, 0x80, 0x7F], np.uint8),
275
+ luts=_h1_luts(),
276
+ gaps=np.zeros(_GAPS_BYTES_PER_BLOCK, np.uint8),
277
+ output_positions=np.array([0, 3], np.uint32),
278
+ split_positions=np.zeros(0, np.int64),
279
+ )
280
+ return (
281
+ _thread_stream_arrays(_WARMUP_BLOCKS * THREADS, _WARMUP_BLOCKS, "warmup-7-block"),
282
+ _thread_stream_arrays(1, 1, "warmup-1-block"),
283
+ _Warmup(
284
+ name="warmup-tiny",
285
+ arrays=tiny,
286
+ expected=np.array([0x3F00, 0xBF80, 0x407F], np.uint16),
287
+ ),
288
+ )
289
+
290
+
291
+ def _warmup_arrays() -> GroupArrays:
292
+ """The seven-block warm-up group as host arrays."""
293
+ return _warmups()[0].arrays
294
+
295
+
296
+ def _warmup_expected() -> npt.NDArray[np.uint16]:
297
+ """The seven-block warm-up group's BF16 bits, derived by hand (none equals 0)."""
298
+ return _warmups()[0].expected
299
+
300
+
301
+ def _warmup_group() -> MxGroup:
302
+ """The seven-block warm-up group as evaluated MLX arrays (also the register-pressure probe's launch group)."""
303
+ return _warmup_arrays().to_mx(name="warmup-7-block")
304
+
305
+
306
+ def _warmup_groups() -> list[tuple[MxGroup, npt.NDArray[np.uint16]]]:
307
+ """Every warm-up group as evaluated MLX arrays with its expected bits, in signature order."""
308
+ return [(w.arrays.to_mx(name=w.name), w.expected) for w in _warmups()]
309
+
310
+
311
+ def ensure_pipeline(
312
+ *, force_direct: bool, poison_buf: bool = False, unguarded_gap_read: bool = False
313
+ ) -> None:
314
+ """Compile this instantiation's pipelines for every input-binding signature and prove each decodes bit-exactly.
315
+
316
+ Each distinct template tuple is its own JIT compile, so each is warmed once and the result
317
+ is cached in `_PIPELINES`. MLX compiles one pipeline per input-binding signature (a read-only
318
+ input of fewer than 8 entries binds `constant`, the rest `device`), so the warm-up dispatches
319
+ three groups: seven blocks with every input in `device` memory, one block whose `positions`
320
+ alone binds `constant`, and an all-tiny group whose `encoded`, `sign_mantissa` and `positions`
321
+ bind `constant`. Those are the signatures real groups present, so the pipelines compiled here
322
+ are the ones the caller's decodes reuse, whatever the group's size. Readiness means every one
323
+ of them compiled, dispatched 512-thread groups, and wrote every element of its group
324
+ bit-exactly, with no error bit on any block and every block on the expected write path
325
+ (staged unless `force_direct`; every warm-up block fits under `CAP`). Each output is prefilled
326
+ with 0, which no expected word equals, so an element the kernel fails to write cannot pass on
327
+ recycled memory. `unguarded_gap_read` is a test-only mutant (see `decode`).
328
+
329
+ Raises:
330
+ DFloatBackendError: Metal is unavailable, a pipeline fails to compile or dispatch here
331
+ (a pipeline ceiling below 512 threads, a driver refusal), a warm-up group decodes to
332
+ the wrong bits or reports an error bit on any block, or any block runs on the wrong
333
+ write path.
334
+ """
335
+ key = (force_direct, poison_buf, unguarded_gap_read)
336
+ if _PIPELINES.get(key):
337
+ return
338
+ if not mx.metal.is_available():
339
+ raise DFloatBackendError("Metal is not available on this machine")
340
+ for group, expected in _warmup_groups():
341
+ try:
342
+ out, status = _dispatch(
343
+ group,
344
+ force_direct=force_direct,
345
+ poison_buf=poison_buf,
346
+ init_value=0,
347
+ unguarded_gap_read=unguarded_gap_read,
348
+ )
349
+ mx.eval(out, status)
350
+ except Exception as exc: # compile error, pipeline ceiling below 512, driver refusal
351
+ raise DFloatBackendError(
352
+ f"the Metal decode kernel cannot run here ({group.name}): {exc}"
353
+ ) from exc
354
+ words = np.array(status)
355
+ if not np.array_equal(np.array(out), expected) or np.any(words & 7):
356
+ raise DFloatBackendError(
357
+ f"the Metal decode kernel produced wrong bits on the warm-up group {group.name}"
358
+ )
359
+ if np.any((words & 8) != (8 if force_direct else 0)):
360
+ want = "direct" if force_direct else "staged"
361
+ raise DFloatBackendError(
362
+ f"the Metal decode kernel took the wrong write path on the warm-up group "
363
+ f"{group.name} (want {want})"
364
+ )
365
+ _PIPELINES[key] = True
366
+
367
+
368
+ def metal_ready() -> bool:
369
+ """Whether both production instantiations (direct and staged) compile here and decode the warm-up groups.
370
+
371
+ True means each instantiation compiled a pipeline for every input-binding signature a real
372
+ group can present and decoded each warm-up group bit-exactly on its expected write path (see
373
+ `ensure_pipeline`), so no later decode meets a first-time compile.
374
+ """
375
+ try:
376
+ ensure_pipeline(force_direct=True)
377
+ ensure_pipeline(force_direct=False)
378
+ except DFloatBackendError:
379
+ return False
380
+ return True
381
+
382
+
383
+ def decode(
384
+ group: MxGroup,
385
+ *,
386
+ force_direct: bool = False,
387
+ _init_value: int | None = None,
388
+ _poison_buf: bool = False,
389
+ _unguarded_gap_read: bool = False,
390
+ ) -> DecodeResult:
391
+ """Lazy Metal decode of one group.
392
+
393
+ `_init_value` and `_poison_buf` are test-only poisons for the write-once tests: they prefill
394
+ the output and the staging buffer so an unwritten element cannot pass as a correct one.
395
+ `_unguarded_gap_read` is a test-only mutant for the shader-validation test: it drops the
396
+ guard on the second `gaps` byte, so the last thread of a full group reads one byte past
397
+ `gaps` (the value is shifted out, so the bits stay correct).
398
+
399
+ Raises:
400
+ DFloatBackendError: The kernel instantiation cannot be warmed up here (see
401
+ `ensure_pipeline`).
402
+ """
403
+ ensure_pipeline(
404
+ force_direct=force_direct, poison_buf=_poison_buf, unguarded_gap_read=_unguarded_gap_read
405
+ )
406
+ out, status = _dispatch(
407
+ group,
408
+ force_direct=force_direct,
409
+ poison_buf=_poison_buf,
410
+ init_value=_init_value,
411
+ unguarded_gap_read=_unguarded_gap_read,
412
+ )
413
+ direct = group.n_launch if force_direct else int(np.count_nonzero(group.intervals > CAP))
414
+ return DecodeResult(
415
+ bits=out,
416
+ status=status,
417
+ backend="metal",
418
+ direct_blocks=direct,
419
+ threadgroup_bytes=THREADGROUP_BYTES_DIRECT if force_direct else THREADGROUP_BYTES_STAGED,
420
+ )