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,240 @@
1
+ """mflux's FLUX.1 Transformer with the block seam composed in front of its two per-block hooks."""
2
+
3
+ import json
4
+ from collections.abc import Mapping
5
+ from functools import cache
6
+ from pathlib import Path
7
+ from typing import Any
8
+
9
+ import mlx.core as mx
10
+ from mlx.utils import tree_flatten
11
+
12
+ from mlx_dfloat._safetensors import TensorInfo, read_header
13
+ from mlx_dfloat.errors import DFloatFormatError, DFloatIntegrationError
14
+ from mlx_dfloat.format import DF11Checkpoint
15
+ from mlx_dfloat.integrate import seam
16
+ from mlx_dfloat.integrate.coverage import check_extras_cover, extras_plan, read_extra
17
+ from mlx_dfloat.integrate.names import NameMap, Shapes
18
+ from mlx_dfloat.integrate.placeholders import PLACEHOLDER, install_placeholders
19
+ from mlx_dfloat.integrate.providers import WeightProvider
20
+ from mlx_dfloat.mflux import require_mflux
21
+ from mlx_dfloat.mflux.flux1.names import (
22
+ DOUBLE_PREFIX,
23
+ DROPPED_EXTRAS,
24
+ SINGLE_PREFIX,
25
+ check_flux_groups,
26
+ flux_name_map,
27
+ )
28
+
29
+ MAX_BUILD_ACTIVE_BYTES = 2 * 1024**3
30
+
31
+
32
+ class SeamMixin:
33
+ """Overrides mflux's ``_apply_joint_transformer_block`` / ``_apply_single_transformer_block``.
34
+
35
+ Compose it in front of mflux's ``Transformer`` (``seam_transformer_class``) or a fake with the same hooks.
36
+ Call ``attach`` before the first step. A step is ``out = transformer(...)``; with ``verify_in_call`` the
37
+ decode status words are checked before the call returns, otherwise call ``verify_step()`` after the
38
+ step's final eval.
39
+ """
40
+
41
+ _seam: seam.SeamState
42
+
43
+ def attach(
44
+ self,
45
+ provider: WeightProvider,
46
+ shapes: Shapes,
47
+ *,
48
+ eval_policy: str = "per-block",
49
+ tracer: seam.TraceHook | None = None,
50
+ verify_in_call: bool = False,
51
+ ) -> None:
52
+ """Bind the provider, the block shapes (from ``install_placeholders``) and the eval policy."""
53
+ self._seam = seam.attach_state(
54
+ provider, shapes, eval_policy=eval_policy, tracer=tracer, verify_in_call=verify_in_call
55
+ )
56
+
57
+ def _state(self) -> seam.SeamState:
58
+ try:
59
+ return self._seam
60
+ except AttributeError as exc:
61
+ raise DFloatIntegrationError(
62
+ "call attach(provider, shapes) before running a step"
63
+ ) from exc
64
+
65
+ def detach(self) -> None:
66
+ """Forget the provider and the shapes; the next step needs a new ``attach``."""
67
+ if hasattr(self, "_seam"):
68
+ self._seam.provider.reset()
69
+ del self._seam
70
+
71
+ def verify_step(self) -> None:
72
+ """Run the provider's deferred checks; call it after the step's final ``mx.eval``."""
73
+ seam.verify_step(self._state())
74
+
75
+ def __call__(self, *args: Any, **kwargs: Any) -> Any:
76
+ """Run mflux's step; a failed step leaves no stale seam state."""
77
+ state = self._state()
78
+ seam.begin_step(state)
79
+ try:
80
+ out = super().__call__(*args, **kwargs) # type: ignore[misc]
81
+ return seam.end_step(state, out)
82
+ except BaseException:
83
+ seam.abort_step(state)
84
+ raise
85
+
86
+ def _apply_joint_transformer_block(self, idx: int, block: Any, **kwargs: Any) -> Any:
87
+ return seam.run_block(
88
+ self._state(),
89
+ f"{DOUBLE_PREFIX}.{idx}",
90
+ block,
91
+ lambda: super(SeamMixin, self)._apply_joint_transformer_block( # type: ignore[misc]
92
+ idx=idx, block=block, **kwargs
93
+ ),
94
+ )
95
+
96
+ def _apply_single_transformer_block(self, idx: int, block: Any, **kwargs: Any) -> Any:
97
+ return seam.run_block(
98
+ self._state(),
99
+ f"{SINGLE_PREFIX}.{idx}",
100
+ block,
101
+ lambda: super(SeamMixin, self)._apply_single_transformer_block( # type: ignore[misc]
102
+ idx=idx, block=block, **kwargs
103
+ ),
104
+ )
105
+
106
+
107
+ @cache
108
+ def seam_transformer_class() -> type:
109
+ """``SeamMixin`` composed in front of mflux's ``Transformer`` (imports mflux).
110
+
111
+ Raises:
112
+ DFloatDependencyError: The optional ``mflux`` extra is not installed.
113
+ """
114
+ require_mflux()
115
+ from mflux.models.flux.model.flux_transformer.transformer import Transformer
116
+
117
+ return type("SeamTransformer", (SeamMixin, Transformer), {})
118
+
119
+
120
+ def base_transformer_index(
121
+ root: Path, *, index_file: str = "diffusion_pytorch_model.safetensors.index.json"
122
+ ) -> dict[str, tuple[Path, TensorInfo]]:
123
+ """Every tensor of a sharded BF16 transformer directory: name -> (shard path, tensor info), from the index.
124
+
125
+ Raises:
126
+ DFloatFormatError: No usable index (missing, malformed JSON, or a ``weight_map`` that is not
127
+ a non-empty object of shard-name strings), or a shard name that is not a plain file name
128
+ in ``root``.
129
+ """
130
+ try:
131
+ weight_map = json.loads((root / index_file).read_text(encoding="utf-8"))["weight_map"]
132
+ except (OSError, ValueError, KeyError, TypeError) as exc:
133
+ raise DFloatFormatError(f"{root / index_file}: no usable weight index: {exc}") from exc
134
+ if (
135
+ not isinstance(weight_map, dict)
136
+ or not weight_map
137
+ or not all(isinstance(v, str) for v in weight_map.values())
138
+ ):
139
+ raise DFloatFormatError(
140
+ f"{root / index_file}: weight_map must be a non-empty object of shard-name strings"
141
+ )
142
+ index: dict[str, tuple[Path, TensorInfo]] = {}
143
+ for shard_name in sorted(set(weight_map.values())):
144
+ if not isinstance(shard_name, str) or Path(shard_name).name != shard_name:
145
+ raise DFloatFormatError(
146
+ f"{root / index_file}: shard {shard_name!r} escapes the directory"
147
+ )
148
+ shard = root / shard_name
149
+ index.update({name: (shard, info) for name, info in read_header(shard).items()})
150
+ missing = sorted(set(weight_map) - set(index))
151
+ if missing:
152
+ raise DFloatFormatError(
153
+ f"{root / index_file}: tensors named by the index but absent from the shards: {missing[:5]}"
154
+ )
155
+ return index
156
+
157
+
158
+ def base_extras(
159
+ index: Mapping[str, tuple[Path, TensorInfo]], ckpt: DF11Checkpoint
160
+ ) -> dict[str, tuple[Path, TensorInfo]]:
161
+ """The index minus the checkpoint's block matrices: what ``build_transformer`` loads as extras from a BF16 base."""
162
+ matrices = {m for group in ckpt.groups.values() for m in group.matrix_names}
163
+ return {name: entry for name, entry in index.items() if name not in matrices}
164
+
165
+
166
+ def build_transformer(
167
+ model_config: Any,
168
+ ckpt: DF11Checkpoint,
169
+ *,
170
+ name_map: NameMap | None = None,
171
+ n_double: int | None = None,
172
+ n_single: int | None = None,
173
+ extras: Mapping[str, tuple[Path, TensorInfo]] | None = None,
174
+ ) -> tuple[Any, Shapes]:
175
+ """Construct mflux's ``Transformer`` (seamed) with placeholders for the matrices and the extras loaded.
176
+
177
+ Block counts come from the checkpoint's groups, or ``n_double``/``n_single`` for a reduced-depth build
178
+ (each between zero and the checkpoint's own count). Every block matrix is a placeholder and every other
179
+ parameter comes from ``extras`` (a BF16 base's index, via ``base_extras``) or, when ``extras`` is None,
180
+ the checkpoint's own extras; the coverage of both is asserted, and so is the MLX active memory the
181
+ build added (under ``MAX_BUILD_ACTIVE_BYTES``).
182
+
183
+ Raises:
184
+ DFloatFormatError: The groups are not FLUX.1's, an extra is not BF16 or has the wrong shape.
185
+ DFloatDependencyError: The optional ``mflux`` extra is not installed.
186
+ DFloatIntegrationError: A depth override is negative or exceeds the checkpoint's own block
187
+ count, an uncovered parameter, an extra without a target, or too much active memory added
188
+ by the build.
189
+ """
190
+ names = flux_name_map() if name_map is None else name_map
191
+ ckpt_double, ckpt_single = check_flux_groups(ckpt)
192
+ n_double = ckpt_double if n_double is None else n_double
193
+ n_single = ckpt_single if n_single is None else n_single
194
+ if n_double < 0 or n_single < 0:
195
+ raise DFloatIntegrationError(
196
+ f"asked for {n_double} double / {n_single} single blocks; a depth cannot be negative"
197
+ )
198
+ if n_double > ckpt_double or n_single > ckpt_single:
199
+ raise DFloatIntegrationError(
200
+ f"asked for {n_double} double / {n_single} single blocks; the checkpoint has "
201
+ f"{ckpt_double} / {ckpt_single}"
202
+ )
203
+ before = int(mx.get_active_memory())
204
+ transformer = seam_transformer_class()(
205
+ model_config, num_transformer_blocks=n_double, num_single_transformer_blocks=n_single
206
+ )
207
+ shapes = install_placeholders(
208
+ [
209
+ (DOUBLE_PREFIX, transformer.transformer_blocks),
210
+ (SINGLE_PREFIX, transformer.single_transformer_blocks),
211
+ ],
212
+ names,
213
+ )
214
+ matrix_paths = {f"{block}.{attr}.weight" for block, per in shapes.items() for attr in per}
215
+ params = dict(tree_flatten(transformer.parameters()))
216
+ plan = extras_plan(
217
+ ckpt,
218
+ names,
219
+ counts={DOUBLE_PREFIX: n_double, SINGLE_PREFIX: n_single},
220
+ dropped=DROPPED_EXTRAS,
221
+ extras=extras,
222
+ )
223
+ check_extras_cover(params, (name for name, _path, _info in plan), matrix_paths)
224
+ weights: list[tuple[str, mx.array]] = []
225
+ for name, path, info in plan:
226
+ array = read_extra(path, info)
227
+ if tuple(array.shape) != tuple(params[name].shape):
228
+ raise DFloatFormatError(
229
+ f"{name}: extra has shape {tuple(array.shape)}, parameter {tuple(params[name].shape)}"
230
+ )
231
+ weights.append((name, array))
232
+ transformer.load_weights(weights, strict=False)
233
+ mx.eval([array for _name, array in weights], PLACEHOLDER)
234
+ added = int(mx.get_active_memory()) - before
235
+ if added >= MAX_BUILD_ACTIVE_BYTES:
236
+ raise DFloatIntegrationError(
237
+ f"extras added {added / 1024**3:.2f} GiB of active memory "
238
+ f"(limit {MAX_BUILD_ACTIVE_BYTES / 1024**3:g} GiB)"
239
+ )
240
+ return transformer, shapes
mlx_dfloat/py.typed ADDED
File without changes
@@ -0,0 +1,241 @@
1
+ """Pure-NumPy, bit-exact DFloat11 decoder: the oracle every GPU kernel is tested against.
2
+
3
+ It emulates upstream's CUDA algorithm in lockstep: each of 512 threads per 4096-byte block owns one
4
+ 64-bit chunk of the exponent stream and decodes the codes that start in it (from its 5-bit gap); an
5
+ exclusive scan turns per-thread counts into output indices. Structural checks run as it decodes.
6
+ """
7
+
8
+ import os
9
+
10
+ import numpy as np
11
+ import numpy.typing as npt
12
+
13
+ from mlx_dfloat.errors import DFloatFormatError, DFloatResourceError
14
+ from mlx_dfloat.format import (
15
+ BYTES_PER_THREAD,
16
+ LUT_POINTER_MIN,
17
+ THREADS_PER_BLOCK,
18
+ DF11Group,
19
+ GroupArrays,
20
+ validate_group_arrays,
21
+ )
22
+
23
+ _BITS_PER_THREAD = 8 * BYTES_PER_THREAD
24
+
25
+ IntArray = npt.NDArray[np.int64]
26
+ BoolArray = npt.NDArray[np.bool_]
27
+ U16Array = npt.NDArray[np.uint16]
28
+
29
+
30
+ def thread_gaps(gaps: npt.NDArray[np.uint8], n_threads: int) -> IntArray:
31
+ """Unpack the MSB-first 5-bit start offsets, one per thread."""
32
+ need = -(-5 * n_threads // 8)
33
+ bits = np.unpackbits(np.asarray(gaps[:need], dtype=np.uint8))[: 5 * n_threads].reshape(
34
+ n_threads, 5
35
+ )
36
+ packed = (
37
+ (bits[:, 0] << 4) | (bits[:, 1] << 3) | (bits[:, 2] << 2) | (bits[:, 3] << 1) | bits[:, 4]
38
+ )
39
+ return packed.astype(np.int64)
40
+
41
+
42
+ def estimate_decode_bytes(arrays: GroupArrays) -> int:
43
+ """Upper estimate of the decoder's working memory for one group."""
44
+ return 12 * arrays.n_threads * BYTES_PER_THREAD + 3 * arrays.n_elements
45
+
46
+
47
+ def default_memory_budget() -> int:
48
+ """Half of physical RAM."""
49
+ return os.sysconf("SC_PAGE_SIZE") * os.sysconf("SC_PHYS_PAGES") // 2
50
+
51
+
52
+ def max_code_length(luts: npt.NDArray[np.uint8]) -> int:
53
+ """Longest code in a group's codebook (the maximum of the lengths row)."""
54
+ return int(np.asarray(luts)[-1].max())
55
+
56
+
57
+ def _peek32(stream: npt.NDArray[np.uint8], pos: IntArray) -> npt.NDArray[np.uint32]:
58
+ byte = pos >> 3
59
+ shift = (pos & 7).astype(np.uint64)
60
+ word = np.zeros(pos.size, dtype=np.uint64)
61
+ for k in range(5):
62
+ word = (word << np.uint64(8)) | stream[byte + k].astype(np.uint64)
63
+ return ((word >> (np.uint64(8) - shift)) & np.uint64(0xFFFFFFFF)).astype(np.uint32)
64
+
65
+
66
+ def _lookup(
67
+ luts: npt.NDArray[np.uint8], words: npt.NDArray[np.uint32]
68
+ ) -> tuple[IntArray, IntArray, BoolArray]:
69
+ n_luts = luts.shape[0]
70
+ sym = luts[0, words >> 24].astype(np.int64)
71
+ ok = np.ones(words.size, dtype=np.bool_)
72
+ for level in (1, 2, 3):
73
+ ptr = sym >= LUT_POINTER_MIN
74
+ if not ptr.any():
75
+ break
76
+ row = 256 - sym[ptr]
77
+ in_range = (row >= 1) & (row <= n_luts - 2)
78
+ byte = (words[ptr] >> np.uint32(24 - 8 * level)) & np.uint32(0xFF)
79
+ sym[ptr] = np.where(in_range, luts[np.where(in_range, row, 0), byte], LUT_POINTER_MIN)
80
+ ok[ptr] &= in_range
81
+ ok &= sym < LUT_POINTER_MIN
82
+ length = np.where(ok, luts[n_luts - 1, np.minimum(sym, 255)], 0).astype(np.int64)
83
+ ok &= length > 0
84
+ return sym, length, ok
85
+
86
+
87
+ def _count(
88
+ stream: npt.NDArray[np.uint8], luts: npt.NDArray[np.uint8], starts: IntArray, ends: IntArray
89
+ ) -> tuple[IntArray, BoolArray, IntArray]:
90
+ counts = np.zeros(starts.size, dtype=np.int64)
91
+ invalid = np.zeros(starts.size, dtype=np.bool_)
92
+ pos = starts.copy()
93
+ active = np.flatnonzero(pos < ends)
94
+ while active.size:
95
+ _, length, ok = _lookup(luts, _peek32(stream, pos[active]))
96
+ invalid[active[~ok]] = True
97
+ good = active[ok]
98
+ counts[good] += 1
99
+ pos[good] += length[ok]
100
+ active = good[pos[good] < ends[good]]
101
+ return counts, invalid, pos
102
+
103
+
104
+ def _write(
105
+ stream: npt.NDArray[np.uint8],
106
+ luts: npt.NDArray[np.uint8],
107
+ sign_mantissa: npt.NDArray[np.uint8],
108
+ starts: IntArray,
109
+ first: IntArray,
110
+ stop: IntArray,
111
+ ) -> tuple[U16Array, int]:
112
+ n = sign_mantissa.size
113
+ out = np.zeros(n, dtype=np.uint16)
114
+ pos = starts.copy()
115
+ idx = first.copy()
116
+ active = np.flatnonzero(idx < stop)
117
+ while active.size:
118
+ sym, length, _ = _lookup(luts, _peek32(stream, pos[active])) # validated in pass 1
119
+ k = idx[active]
120
+ sm = sign_mantissa[k].astype(np.uint16)
121
+ out[k] = ((sm & 0x80) << 8) | (sym.astype(np.uint16) << 7) | (sm & 0x7F)
122
+ idx[active] += 1
123
+ pos[active] += length
124
+ active = active[idx[active] < stop[active]]
125
+ last_thread = int(np.searchsorted(first, n - 1, side="right") - 1)
126
+ return out, int(pos[last_thread])
127
+
128
+
129
+ def decode_group(
130
+ arrays: GroupArrays,
131
+ *,
132
+ name: str = "<group>",
133
+ max_memory_bytes: int | None = None,
134
+ check_stream_end: bool = True,
135
+ ) -> U16Array:
136
+ """Decode one compressed group to its BF16 bit patterns (flat, concatenation order).
137
+
138
+ ``check_stream_end=False`` skips only the EOF byte-count check; it exists for fixtures that are
139
+ deliberately truncated mid-stream (the committed upstream slice), never for real checkpoints.
140
+
141
+ Raises:
142
+ DFloatFormatError: The arrays are structurally invalid or internally inconsistent.
143
+ DFloatResourceError: The estimated working memory exceeds ``max_memory_bytes`` (default:
144
+ half of physical RAM).
145
+ """
146
+ validate_group_arrays(arrays, name=name)
147
+ budget = default_memory_budget() if max_memory_bytes is None else max_memory_bytes
148
+ need = estimate_decode_bytes(arrays)
149
+ if need > budget:
150
+ raise DFloatResourceError(
151
+ f"{name}: decoding needs about {need} bytes, over the {budget}-byte budget"
152
+ )
153
+ n, n_bytes, n_threads = arrays.n_elements, arrays.n_bytes, arrays.n_threads
154
+ n_bits = 8 * n_bytes
155
+ last_real = (n_bits - 1) // _BITS_PER_THREAD
156
+ g = np.arange(last_real + 1, dtype=np.int64)
157
+ starts = g * _BITS_PER_THREAD + thread_gaps(arrays.gaps, n_threads)[: last_real + 1]
158
+ ends = (g + 1) * _BITS_PER_THREAD
159
+ stream = np.zeros(n_threads * BYTES_PER_THREAD + 8, dtype=np.uint8)
160
+ stream[:n_bytes] = arrays.encoded_exponent
161
+ luts = np.ascontiguousarray(arrays.luts)
162
+
163
+ counts, invalid, ended = _count(stream, luts, starts, ends)
164
+ bad = np.flatnonzero(invalid[:-1])
165
+ if bad.size:
166
+ t = int(bad[0])
167
+ raise DFloatFormatError(
168
+ f"{name}: invalid code in thread {t} (block {t // THREADS_PER_BLOCK})"
169
+ )
170
+ inside = ended[:-1] < n_bits
171
+ broken = np.flatnonzero(inside & (ended[:-1] != starts[1:]))
172
+ if broken.size:
173
+ t = int(broken[0])
174
+ raise DFloatFormatError(
175
+ f"{name}: thread {t} ends at bit {int(ended[t])} but thread {t + 1} starts at bit "
176
+ f"{int(starts[t + 1])} (corrupt gaps?)"
177
+ )
178
+ first = np.zeros_like(counts)
179
+ np.cumsum(counts[:-1], out=first[1:])
180
+ positions = arrays.output_positions.astype(np.int64)
181
+ block_starts = np.arange(positions.size - 1, dtype=np.int64) * THREADS_PER_BLOCK
182
+ mismatch = np.flatnonzero(first[block_starts] != positions[:-1])
183
+ if mismatch.size:
184
+ b = int(mismatch[0])
185
+ raise DFloatFormatError(
186
+ f"{name}: block {b} starts at element {int(first[block_starts[b]])}, output_positions says "
187
+ f"{int(positions[b])}"
188
+ )
189
+ if positions.size - 1 == arrays.n_blocks - 1:
190
+ # The last block holds at least one byte, so its first thread is always a real one
191
+ # (tail <= last_real) and first[tail] is in range.
192
+ tail = (arrays.n_blocks - 1) * THREADS_PER_BLOCK
193
+ if first[tail] < n:
194
+ raise DFloatFormatError(
195
+ f"{name}: block {arrays.n_blocks - 1} starts codes but output_positions has no entry for it"
196
+ )
197
+ total = int(first[-1] + counts[-1])
198
+ if total < n:
199
+ hint = " (invalid code in the last thread)" if invalid[-1] else ""
200
+ raise DFloatFormatError(f"{name}: stream holds {total} codes, expected {n}{hint}")
201
+ stop = np.minimum(first + counts, n)
202
+ out, end_bit = _write(stream, luts, np.asarray(arrays.sign_mantissa), starts, first, stop)
203
+ if check_stream_end and -(-end_bit // 8) != n_bytes:
204
+ raise DFloatFormatError(
205
+ f"{name}: data ends at bit {end_bit}, so the stream should be {-(-end_bit // 8)} bytes, not {n_bytes}"
206
+ )
207
+ return out
208
+
209
+
210
+ def split_matrices(flat: U16Array, split_positions: npt.NDArray[np.int64]) -> list[U16Array]:
211
+ """Split a decoded group into its matrices (flat), at ``split_positions``."""
212
+ return list(np.split(flat, split_positions.astype(np.int64)))
213
+
214
+
215
+ def decode_matrices(
216
+ group: DF11Group, *, max_memory_bytes: int | None = None
217
+ ) -> dict[str, U16Array]:
218
+ """Decode a group and return its matrices keyed by weight name (flat BF16 bit patterns)."""
219
+ arrays = group.load()
220
+ parts = split_matrices(
221
+ decode_group(arrays, name=group.name, max_memory_bytes=max_memory_bytes),
222
+ arrays.split_positions,
223
+ )
224
+ return dict(zip(group.matrix_names, parts, strict=True))
225
+
226
+
227
+ def max_elements_per_block(output_positions: npt.NDArray[np.uint32]) -> int:
228
+ """Largest number of elements whose codes start in one 4096-byte block (threadgroup sizing)."""
229
+ return int(np.diff(output_positions.astype(np.int64)).max())
230
+
231
+
232
+ __all__ = [
233
+ "decode_group",
234
+ "decode_matrices",
235
+ "default_memory_budget",
236
+ "estimate_decode_bytes",
237
+ "max_code_length",
238
+ "max_elements_per_block",
239
+ "split_matrices",
240
+ "thread_gaps",
241
+ ]