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,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
|
mlx_dfloat/reference.py
ADDED
|
@@ -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
|
+
]
|