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
mlx_dfloat/format.py
ADDED
|
@@ -0,0 +1,462 @@
|
|
|
1
|
+
"""DFloat11 checkpoint format: config, per-group arrays, validation, discovery."""
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import os
|
|
5
|
+
import re
|
|
6
|
+
import stat
|
|
7
|
+
from collections.abc import Mapping
|
|
8
|
+
from dataclasses import dataclass
|
|
9
|
+
from pathlib import Path
|
|
10
|
+
|
|
11
|
+
import mlx.core as mx
|
|
12
|
+
import numpy as np
|
|
13
|
+
import numpy.typing as npt
|
|
14
|
+
|
|
15
|
+
from mlx_dfloat._safetensors import TensorInfo, read_array, read_header, short_repr
|
|
16
|
+
from mlx_dfloat.errors import DFloatBackendError, DFloatFormatError
|
|
17
|
+
|
|
18
|
+
SUPPORTED_VERSIONS: frozenset[str] = frozenset({"0.2.0", "0.3.1", "0.3.2", "0.5.0"})
|
|
19
|
+
THREADS_PER_BLOCK = 512
|
|
20
|
+
BYTES_PER_THREAD = 8
|
|
21
|
+
BLOCK_BYTES = THREADS_PER_BLOCK * BYTES_PER_THREAD
|
|
22
|
+
LUT_POINTER_MIN = 240
|
|
23
|
+
MAX_LUT_ROWS = (
|
|
24
|
+
18 # row 0, up to 16 pointer-target rows (pointers 255..240 -> rows 1..16), lengths row
|
|
25
|
+
)
|
|
26
|
+
MAX_PATTERNS = 64
|
|
27
|
+
MAX_PATTERN_CHARS = 256
|
|
28
|
+
# pattern_dict keys come from a downloaded config.json and are matched with re.fullmatch, so they
|
|
29
|
+
# are held to the small grammar upstream DF11 configs use (literals, `\.` and a bare `.`, which
|
|
30
|
+
# the 0.2.0 configs of Qwen3-4B and FLUX.1-dev/schnell leave unescaped, `\d`, `\w`, classes such
|
|
31
|
+
# as `[0-9]`, `+`/`*`/`?`, `|` and plain groups), with a cap on every kind of backtracking choice
|
|
32
|
+
# point. Group names are at most 256 characters, so two unbounded quantifiers cost at most
|
|
33
|
+
# ~256^2 steps per match; optional parts and alternation branches each double the worst case.
|
|
34
|
+
_QUANTIFIED_GROUP = re.compile(r"\)[*+{]")
|
|
35
|
+
_PATTERN_TOKEN = re.compile(r"\\[.dw]|[A-Za-z0-9_.\-\[\]()|*+?]")
|
|
36
|
+
MAX_UNBOUNDED_QUANTIFIERS = 2
|
|
37
|
+
MAX_OPTIONAL_QUANTIFIERS = 2
|
|
38
|
+
MAX_ALTERNATIONS = 4
|
|
39
|
+
MAX_ARRAY_ELEMENTS = 2**31 - 1 # Metal shape buffers and grid sizes are int32
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
@dataclass(frozen=True, slots=True, kw_only=True)
|
|
43
|
+
class DF11Config:
|
|
44
|
+
"""The ``dfloat11_config`` block of a checkpoint's ``config.json``."""
|
|
45
|
+
|
|
46
|
+
version: str
|
|
47
|
+
threads_per_block: int
|
|
48
|
+
bytes_per_thread: int
|
|
49
|
+
pattern_dict: Mapping[str, tuple[str, ...]]
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def _check_pattern(pattern: str, *, source: str) -> None:
|
|
53
|
+
if len(pattern) > MAX_PATTERN_CHARS:
|
|
54
|
+
raise DFloatFormatError(
|
|
55
|
+
f"{source}: pattern_dict pattern is too long ({len(pattern)} chars)"
|
|
56
|
+
)
|
|
57
|
+
if _QUANTIFIED_GROUP.search(pattern):
|
|
58
|
+
raise DFloatFormatError(
|
|
59
|
+
f"{source}: pattern {pattern!r} has a quantified group; refused to avoid catastrophic backtracking"
|
|
60
|
+
)
|
|
61
|
+
pos = 0
|
|
62
|
+
while pos < len(pattern):
|
|
63
|
+
token = _PATTERN_TOKEN.match(pattern, pos)
|
|
64
|
+
if token is None:
|
|
65
|
+
raise DFloatFormatError(
|
|
66
|
+
f"{source}: pattern {pattern!r}: {pattern[pos : pos + 2]!r} at offset {pos} is not "
|
|
67
|
+
"allowed in a DF11 pattern"
|
|
68
|
+
)
|
|
69
|
+
pos = token.end()
|
|
70
|
+
if "(?" in pattern:
|
|
71
|
+
raise DFloatFormatError(
|
|
72
|
+
f"{source}: pattern {pattern!r}: inline flags and extension groups '(?' are not allowed"
|
|
73
|
+
)
|
|
74
|
+
unbounded = pattern.count("*") + pattern.count("+")
|
|
75
|
+
if unbounded > MAX_UNBOUNDED_QUANTIFIERS:
|
|
76
|
+
raise DFloatFormatError(
|
|
77
|
+
f"{source}: pattern {pattern!r} has {unbounded} unbounded quantifiers "
|
|
78
|
+
f"(at most {MAX_UNBOUNDED_QUANTIFIERS})"
|
|
79
|
+
)
|
|
80
|
+
for symbol, limit, what in (
|
|
81
|
+
("?", MAX_OPTIONAL_QUANTIFIERS, "optional parts"),
|
|
82
|
+
("|", MAX_ALTERNATIONS, "alternations"),
|
|
83
|
+
):
|
|
84
|
+
if pattern.count(symbol) > limit:
|
|
85
|
+
raise DFloatFormatError(
|
|
86
|
+
f"{source}: pattern {pattern!r} has too many {what} (at most {limit})"
|
|
87
|
+
)
|
|
88
|
+
try:
|
|
89
|
+
re.compile(pattern)
|
|
90
|
+
except re.error as exc:
|
|
91
|
+
raise DFloatFormatError(
|
|
92
|
+
f"{source}: pattern {pattern!r} is not a valid regular expression"
|
|
93
|
+
) from exc
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
def parse_df11_config(raw: object, *, source: str) -> DF11Config:
|
|
97
|
+
"""Validate a ``dfloat11_config`` mapping and return it typed.
|
|
98
|
+
|
|
99
|
+
Raises:
|
|
100
|
+
DFloatFormatError: The version is unsupported, a field is missing or malformed, or a
|
|
101
|
+
pattern is unsafe.
|
|
102
|
+
"""
|
|
103
|
+
if not isinstance(raw, dict):
|
|
104
|
+
raise DFloatFormatError(f"{source}: dfloat11_config is not an object")
|
|
105
|
+
version = raw.get("version")
|
|
106
|
+
if not isinstance(version, str):
|
|
107
|
+
raise DFloatFormatError(f"{source}: dfloat11_config has no version string")
|
|
108
|
+
if version not in SUPPORTED_VERSIONS:
|
|
109
|
+
raise DFloatFormatError(
|
|
110
|
+
f"{source}: unsupported DF11 format version {short_repr(version)} "
|
|
111
|
+
f"(supported: {', '.join(sorted(SUPPORTED_VERSIONS))})"
|
|
112
|
+
)
|
|
113
|
+
if raw.get("threads_per_block") != [THREADS_PER_BLOCK]:
|
|
114
|
+
raise DFloatFormatError(
|
|
115
|
+
f"{source}: threads_per_block must be [512], got {short_repr(raw.get('threads_per_block'))}"
|
|
116
|
+
)
|
|
117
|
+
if raw.get("bytes_per_thread") != BYTES_PER_THREAD:
|
|
118
|
+
raise DFloatFormatError(
|
|
119
|
+
f"{source}: bytes_per_thread must be 8, got {short_repr(raw.get('bytes_per_thread'))}"
|
|
120
|
+
)
|
|
121
|
+
patterns = raw.get("pattern_dict")
|
|
122
|
+
if not isinstance(patterns, dict) or not patterns:
|
|
123
|
+
raise DFloatFormatError(f"{source}: pattern_dict is missing or empty")
|
|
124
|
+
if len(patterns) > MAX_PATTERNS:
|
|
125
|
+
raise DFloatFormatError(f"{source}: pattern_dict has too many patterns ({len(patterns)})")
|
|
126
|
+
parsed: dict[str, tuple[str, ...]] = {}
|
|
127
|
+
for pattern, subpaths in patterns.items():
|
|
128
|
+
if not isinstance(subpaths, list) or not all(isinstance(s, str) for s in subpaths):
|
|
129
|
+
raise DFloatFormatError(
|
|
130
|
+
f"{source}: pattern_dict entry {short_repr(pattern)} is not a list of names"
|
|
131
|
+
)
|
|
132
|
+
_check_pattern(pattern, source=source)
|
|
133
|
+
parsed[pattern] = tuple(subpaths)
|
|
134
|
+
return DF11Config(
|
|
135
|
+
version=version,
|
|
136
|
+
threads_per_block=THREADS_PER_BLOCK,
|
|
137
|
+
bytes_per_thread=BYTES_PER_THREAD,
|
|
138
|
+
pattern_dict=parsed,
|
|
139
|
+
)
|
|
140
|
+
|
|
141
|
+
|
|
142
|
+
def read_df11_config(model_dir: Path) -> DF11Config:
|
|
143
|
+
"""Read ``config.json`` from a DF11 model directory.
|
|
144
|
+
|
|
145
|
+
Raises:
|
|
146
|
+
DFloatFormatError: The directory is a legacy pickle-format DF11 repo, or has no usable config.
|
|
147
|
+
"""
|
|
148
|
+
config_path = model_dir / "config.json"
|
|
149
|
+
has_legacy = any(model_dir.glob("*.pkl")) or any(model_dir.glob("*.ptx"))
|
|
150
|
+
config: object = None
|
|
151
|
+
if config_path.is_file():
|
|
152
|
+
try:
|
|
153
|
+
config = json.loads(config_path.read_text(encoding="utf-8"))
|
|
154
|
+
except (ValueError, RecursionError) as exc:
|
|
155
|
+
raise DFloatFormatError(f"{config_path}: not valid JSON") from exc
|
|
156
|
+
has_df11 = isinstance(config, dict) and "dfloat11_config" in config
|
|
157
|
+
if has_legacy and not has_df11:
|
|
158
|
+
raise DFloatFormatError(
|
|
159
|
+
f"{model_dir}: legacy pickle-format DF11 checkpoint (.pkl/.ptx) is not supported; "
|
|
160
|
+
"mlx-dfloat never unpickles files"
|
|
161
|
+
)
|
|
162
|
+
if config is None:
|
|
163
|
+
raise DFloatFormatError(f"{model_dir}: no config.json")
|
|
164
|
+
if not has_df11:
|
|
165
|
+
raise DFloatFormatError(f"{config_path}: no dfloat11_config block")
|
|
166
|
+
assert isinstance(config, dict) # narrowed by has_df11
|
|
167
|
+
return parse_df11_config(config["dfloat11_config"], source=str(config_path))
|
|
168
|
+
|
|
169
|
+
|
|
170
|
+
def n_blocks_for(n_bytes: int) -> int:
|
|
171
|
+
"""Number of 4096-byte thread-blocks that cover ``n_bytes`` of encoded exponents."""
|
|
172
|
+
return -(-n_bytes // BLOCK_BYTES)
|
|
173
|
+
|
|
174
|
+
|
|
175
|
+
@dataclass(frozen=True, slots=True, kw_only=True)
|
|
176
|
+
class MxGroup:
|
|
177
|
+
"""One compressed group as evaluated MLX arrays, plus its host-side metadata.
|
|
178
|
+
|
|
179
|
+
A decode backend (the Metal kernel, the reference-decode capability) reads this metadata
|
|
180
|
+
without a device sync.
|
|
181
|
+
"""
|
|
182
|
+
|
|
183
|
+
name: str
|
|
184
|
+
encoded_exponent: mx.array # uint8[n_bytes]
|
|
185
|
+
sign_mantissa: mx.array # uint8[n_elements]
|
|
186
|
+
luts: mx.array # uint8[n_luts, 256]
|
|
187
|
+
gaps: mx.array # uint8, as stored (320 * n_blocks bytes)
|
|
188
|
+
positions: mx.array # uint32[n_launch + 1], as stored
|
|
189
|
+
intervals: npt.NDArray[np.int64] # host: np.diff(positions), for path counting without a sync
|
|
190
|
+
n_elements: int
|
|
191
|
+
n_bytes: int
|
|
192
|
+
n_luts: int
|
|
193
|
+
n_launch: int # len(positions) - 1
|
|
194
|
+
max_elements_per_block: int
|
|
195
|
+
split_positions: tuple[int, ...]
|
|
196
|
+
|
|
197
|
+
|
|
198
|
+
@dataclass(frozen=True, slots=True, kw_only=True)
|
|
199
|
+
class GroupArrays:
|
|
200
|
+
"""The six arrays stored for one compressed group, in decoder-ready dtypes."""
|
|
201
|
+
|
|
202
|
+
encoded_exponent: npt.NDArray[np.uint8]
|
|
203
|
+
sign_mantissa: npt.NDArray[np.uint8]
|
|
204
|
+
luts: npt.NDArray[np.uint8]
|
|
205
|
+
gaps: npt.NDArray[np.uint8]
|
|
206
|
+
output_positions: npt.NDArray[np.uint32]
|
|
207
|
+
split_positions: npt.NDArray[np.int64]
|
|
208
|
+
|
|
209
|
+
@property
|
|
210
|
+
def n_elements(self) -> int:
|
|
211
|
+
"""Number of BF16 values in the group."""
|
|
212
|
+
return int(self.sign_mantissa.size)
|
|
213
|
+
|
|
214
|
+
@property
|
|
215
|
+
def n_bytes(self) -> int:
|
|
216
|
+
"""Length of the encoded exponent stream in bytes."""
|
|
217
|
+
return int(self.encoded_exponent.size)
|
|
218
|
+
|
|
219
|
+
@property
|
|
220
|
+
def n_blocks(self) -> int:
|
|
221
|
+
"""Number of 512-thread blocks the upstream kernel launches for this group."""
|
|
222
|
+
return n_blocks_for(self.n_bytes)
|
|
223
|
+
|
|
224
|
+
@property
|
|
225
|
+
def n_threads(self) -> int:
|
|
226
|
+
"""Total decode threads (512 per block)."""
|
|
227
|
+
return self.n_blocks * THREADS_PER_BLOCK
|
|
228
|
+
|
|
229
|
+
def to_mx(self, *, name: str = "<group>") -> MxGroup:
|
|
230
|
+
"""Validate, copy into MLX arrays, evaluate, and compute host metadata once.
|
|
231
|
+
|
|
232
|
+
Raises:
|
|
233
|
+
DFloatFormatError: The arrays fail the structural checks in ``validate_group_arrays``.
|
|
234
|
+
DFloatBackendError: The byte or element count exceeds the int32 bound the Metal
|
|
235
|
+
backend's shape buffers and grid sizes use.
|
|
236
|
+
"""
|
|
237
|
+
validate_group_arrays(self, name=name)
|
|
238
|
+
for label, size in (("bytes", self.n_bytes), ("elements", self.n_elements)):
|
|
239
|
+
if size > MAX_ARRAY_ELEMENTS:
|
|
240
|
+
raise DFloatBackendError(
|
|
241
|
+
f"{name}: {size} {label} exceeds 2^31-1; the Metal backend uses int32 sizes"
|
|
242
|
+
)
|
|
243
|
+
positions = self.output_positions.astype(np.uint32)
|
|
244
|
+
arrays = {
|
|
245
|
+
"encoded_exponent": mx.array(np.ascontiguousarray(self.encoded_exponent)),
|
|
246
|
+
"sign_mantissa": mx.array(np.ascontiguousarray(self.sign_mantissa)),
|
|
247
|
+
"luts": mx.array(np.ascontiguousarray(self.luts)),
|
|
248
|
+
"gaps": mx.array(np.ascontiguousarray(self.gaps)),
|
|
249
|
+
"positions": mx.array(positions),
|
|
250
|
+
}
|
|
251
|
+
mx.eval(*arrays.values())
|
|
252
|
+
positions_i64 = positions.astype(np.int64)
|
|
253
|
+
intervals = np.diff(positions_i64)
|
|
254
|
+
return MxGroup(
|
|
255
|
+
name=name,
|
|
256
|
+
**arrays,
|
|
257
|
+
intervals=intervals,
|
|
258
|
+
n_elements=self.n_elements,
|
|
259
|
+
n_bytes=self.n_bytes,
|
|
260
|
+
n_luts=int(self.luts.shape[0]),
|
|
261
|
+
n_launch=int(positions.size - 1),
|
|
262
|
+
max_elements_per_block=int(intervals.max()),
|
|
263
|
+
split_positions=tuple(int(s) for s in self.split_positions),
|
|
264
|
+
)
|
|
265
|
+
|
|
266
|
+
|
|
267
|
+
def validate_group_arrays(arrays: GroupArrays, *, name: str) -> None:
|
|
268
|
+
"""Check the structural invariants a well-formed DF11 group satisfies.
|
|
269
|
+
|
|
270
|
+
Raises:
|
|
271
|
+
DFloatFormatError: Any invariant is violated; the message names the group.
|
|
272
|
+
"""
|
|
273
|
+
n, n_bytes = arrays.n_elements, arrays.n_bytes
|
|
274
|
+
if n == 0 or n_bytes == 0:
|
|
275
|
+
raise DFloatFormatError(f"{name}: empty group")
|
|
276
|
+
luts = arrays.luts
|
|
277
|
+
if luts.ndim != 2 or luts.shape[1] != 256 or not 2 <= luts.shape[0] <= MAX_LUT_ROWS:
|
|
278
|
+
raise DFloatFormatError(f"{name}: luts must be [2..{MAX_LUT_ROWS}, 256], got {luts.shape}")
|
|
279
|
+
decode_rows = luts[:-1]
|
|
280
|
+
pointers = decode_rows[decode_rows >= LUT_POINTER_MIN].astype(np.int64)
|
|
281
|
+
targets = 256 - pointers
|
|
282
|
+
if pointers.size and (targets.min() < 1 or targets.max() > luts.shape[0] - 2):
|
|
283
|
+
raise DFloatFormatError(
|
|
284
|
+
f"{name}: a LUT pointer targets a row outside 1..{luts.shape[0] - 2}"
|
|
285
|
+
)
|
|
286
|
+
positions = arrays.output_positions.astype(np.int64)
|
|
287
|
+
n_blocks = arrays.n_blocks
|
|
288
|
+
if positions.size < 2 or positions.size - 1 not in (n_blocks, n_blocks - 1):
|
|
289
|
+
raise DFloatFormatError(
|
|
290
|
+
f"{name}: output_positions has {positions.size} entries for {n_blocks} blocks"
|
|
291
|
+
)
|
|
292
|
+
if positions[0] != 0:
|
|
293
|
+
raise DFloatFormatError(f"{name}: first output position is {positions[0]}, expected 0")
|
|
294
|
+
if np.any(np.diff(positions) < 0):
|
|
295
|
+
raise DFloatFormatError(f"{name}: output_positions is not monotonic")
|
|
296
|
+
if positions[-1] != n:
|
|
297
|
+
raise DFloatFormatError(f"{name}: last output position {positions[-1]} != {n} elements")
|
|
298
|
+
need_gap_bytes = -(-5 * arrays.n_threads // 8)
|
|
299
|
+
if arrays.gaps.size < need_gap_bytes:
|
|
300
|
+
raise DFloatFormatError(
|
|
301
|
+
f"{name}: gaps has {arrays.gaps.size} bytes, needs {need_gap_bytes}"
|
|
302
|
+
)
|
|
303
|
+
split = arrays.split_positions.astype(np.int64)
|
|
304
|
+
if split.size and (split[0] <= 0 or split[-1] >= n or np.any(np.diff(split) <= 0)):
|
|
305
|
+
raise DFloatFormatError(
|
|
306
|
+
f"{name}: split_positions must be strictly increasing inside (0, {n})"
|
|
307
|
+
)
|
|
308
|
+
|
|
309
|
+
|
|
310
|
+
GROUP_FIELDS: tuple[str, ...] = (
|
|
311
|
+
"encoded_exponent",
|
|
312
|
+
"sign_mantissa",
|
|
313
|
+
"luts",
|
|
314
|
+
"gaps",
|
|
315
|
+
"output_positions",
|
|
316
|
+
"split_positions",
|
|
317
|
+
)
|
|
318
|
+
GROUP_FIELD_TYPES: Mapping[str, tuple[str, int]] = {
|
|
319
|
+
"encoded_exponent": ("U8", 1),
|
|
320
|
+
"sign_mantissa": ("U8", 1),
|
|
321
|
+
"luts": ("U8", 2),
|
|
322
|
+
"gaps": ("U8", 1),
|
|
323
|
+
"output_positions": ("U8", 1),
|
|
324
|
+
"split_positions": ("I64", 1),
|
|
325
|
+
}
|
|
326
|
+
GROUP_NAME = re.compile(r"[A-Za-z0-9_.\-]{1,256}")
|
|
327
|
+
|
|
328
|
+
|
|
329
|
+
def validate_group_name(name: str) -> None:
|
|
330
|
+
"""Refuse group names that could escape a directory or blow up regex matching."""
|
|
331
|
+
if not GROUP_NAME.fullmatch(name) or ".." in name:
|
|
332
|
+
raise DFloatFormatError(f"unsafe group name {name[:80]!r}")
|
|
333
|
+
|
|
334
|
+
|
|
335
|
+
def matrix_names_for(group: str, pattern_dict: Mapping[str, tuple[str, ...]]) -> tuple[str, ...]:
|
|
336
|
+
"""Names of the weight matrices a group decodes into, in concatenation order.
|
|
337
|
+
|
|
338
|
+
Raises:
|
|
339
|
+
DFloatFormatError: No pattern, or more than one pattern, fully matches the group name.
|
|
340
|
+
"""
|
|
341
|
+
label = short_repr(group)
|
|
342
|
+
for pattern in pattern_dict: # a hand-built DF11Config never went through parse_df11_config
|
|
343
|
+
_check_pattern(pattern, source=f"group {label}")
|
|
344
|
+
matches = [subs for pattern, subs in pattern_dict.items() if re.fullmatch(pattern, group)]
|
|
345
|
+
if not matches:
|
|
346
|
+
raise DFloatFormatError(f"group {label}: no pattern in pattern_dict matches it")
|
|
347
|
+
if len(matches) > 1:
|
|
348
|
+
raise DFloatFormatError(f"group {label}: more than one pattern in pattern_dict matches it")
|
|
349
|
+
subs = matches[0]
|
|
350
|
+
if not subs:
|
|
351
|
+
return (f"{group}.weight",)
|
|
352
|
+
return tuple(f"{group}.{sub}.weight" for sub in subs)
|
|
353
|
+
|
|
354
|
+
|
|
355
|
+
@dataclass(frozen=True, slots=True, kw_only=True)
|
|
356
|
+
class DF11Group:
|
|
357
|
+
"""One compressed group: where its six tensors live and which matrices it decodes into."""
|
|
358
|
+
|
|
359
|
+
name: str
|
|
360
|
+
matrix_names: tuple[str, ...]
|
|
361
|
+
path: Path
|
|
362
|
+
tensors: Mapping[str, TensorInfo]
|
|
363
|
+
|
|
364
|
+
def load(self) -> GroupArrays:
|
|
365
|
+
"""Memory-map the group's tensors (dtypes checked at discovery) and validate them."""
|
|
366
|
+
raw = {field: read_array(self.path, self.tensors[field]) for field in GROUP_FIELDS}
|
|
367
|
+
positions = np.ascontiguousarray(raw["output_positions"])
|
|
368
|
+
if positions.size % 4:
|
|
369
|
+
raise DFloatFormatError(
|
|
370
|
+
f"{self.name}: output_positions byte length is not a multiple of 4"
|
|
371
|
+
)
|
|
372
|
+
arrays = GroupArrays(
|
|
373
|
+
encoded_exponent=np.asarray(raw["encoded_exponent"]),
|
|
374
|
+
sign_mantissa=np.asarray(raw["sign_mantissa"]),
|
|
375
|
+
luts=np.asarray(raw["luts"]),
|
|
376
|
+
gaps=np.asarray(raw["gaps"]),
|
|
377
|
+
output_positions=positions.view("<u4").astype(np.uint32),
|
|
378
|
+
split_positions=np.asarray(raw["split_positions"]),
|
|
379
|
+
)
|
|
380
|
+
validate_group_arrays(arrays, name=self.name)
|
|
381
|
+
return arrays
|
|
382
|
+
|
|
383
|
+
|
|
384
|
+
def load_group_mx(group: DF11Group) -> MxGroup:
|
|
385
|
+
"""Load one group through the bounds-checked reader and hand it to the backends as MLX arrays."""
|
|
386
|
+
return group.load().to_mx(name=group.name)
|
|
387
|
+
|
|
388
|
+
|
|
389
|
+
@dataclass(frozen=True, slots=True, kw_only=True)
|
|
390
|
+
class DF11Checkpoint:
|
|
391
|
+
"""A DF11 model directory: its config, compressed groups, and uncompressed extra tensors."""
|
|
392
|
+
|
|
393
|
+
root: Path
|
|
394
|
+
config: DF11Config
|
|
395
|
+
groups: Mapping[str, DF11Group]
|
|
396
|
+
extras: Mapping[str, tuple[Path, TensorInfo]]
|
|
397
|
+
|
|
398
|
+
|
|
399
|
+
def _regular_file(path: Path) -> bool:
|
|
400
|
+
try:
|
|
401
|
+
return stat.S_ISREG(path.stat().st_mode) # follows symlinks: HF blobs are regular files
|
|
402
|
+
except OSError:
|
|
403
|
+
return False
|
|
404
|
+
|
|
405
|
+
|
|
406
|
+
def open_checkpoint(path: str | os.PathLike[str]) -> DF11Checkpoint:
|
|
407
|
+
"""Discover the groups and extras of a DF11 checkpoint directory, reading headers only.
|
|
408
|
+
|
|
409
|
+
Raises:
|
|
410
|
+
DFloatFormatError: The config is unusable; a shard is not a regular file; a group name is
|
|
411
|
+
unsafe, a group is incomplete, split across files, or has wrong dtypes/ranks; a tensor
|
|
412
|
+
name appears in two files; or a group's matrix count disagrees with its pattern.
|
|
413
|
+
"""
|
|
414
|
+
root = Path(path).expanduser()
|
|
415
|
+
config = read_df11_config(root)
|
|
416
|
+
owner: dict[str, Path] = {}
|
|
417
|
+
headers: dict[Path, dict[str, TensorInfo]] = {}
|
|
418
|
+
for file in sorted(root.glob("*.safetensors")):
|
|
419
|
+
if not _regular_file(file):
|
|
420
|
+
raise DFloatFormatError(f"{file.name}: not a regular file")
|
|
421
|
+
header = read_header(file)
|
|
422
|
+
headers[file] = header
|
|
423
|
+
for name in header:
|
|
424
|
+
if name in owner:
|
|
425
|
+
raise DFloatFormatError(
|
|
426
|
+
f"tensor {short_repr(name)} appears in both {owner[name].name} and {file.name}"
|
|
427
|
+
)
|
|
428
|
+
owner[name] = file
|
|
429
|
+
group_names = sorted({n.rsplit(".", 1)[0] for n in owner if n.endswith(".encoded_exponent")})
|
|
430
|
+
groups: dict[str, DF11Group] = {}
|
|
431
|
+
claimed: set[str] = set()
|
|
432
|
+
for group in group_names:
|
|
433
|
+
validate_group_name(group)
|
|
434
|
+
home = owner[f"{group}.encoded_exponent"]
|
|
435
|
+
tensors: dict[str, TensorInfo] = {}
|
|
436
|
+
for field in GROUP_FIELDS:
|
|
437
|
+
full = f"{group}.{field}"
|
|
438
|
+
if full not in owner:
|
|
439
|
+
raise DFloatFormatError(f"group {group!r}: missing {field}")
|
|
440
|
+
if owner[full] != home:
|
|
441
|
+
raise DFloatFormatError(f"group {group!r}: its tensors are split across files")
|
|
442
|
+
info = headers[home][full]
|
|
443
|
+
want_dtype, want_rank = GROUP_FIELD_TYPES[field]
|
|
444
|
+
if info.dtype != want_dtype:
|
|
445
|
+
raise DFloatFormatError(
|
|
446
|
+
f"group {group!r}: {field} has dtype {info.dtype}, expected {want_dtype}"
|
|
447
|
+
)
|
|
448
|
+
if len(info.shape) != want_rank:
|
|
449
|
+
raise DFloatFormatError(
|
|
450
|
+
f"group {group!r}: {field} has rank {len(info.shape)}, expected {want_rank}"
|
|
451
|
+
)
|
|
452
|
+
tensors[field] = info
|
|
453
|
+
claimed.add(full)
|
|
454
|
+
names = matrix_names_for(group, config.pattern_dict)
|
|
455
|
+
n_matrices = tensors["split_positions"].shape[0] + 1
|
|
456
|
+
if n_matrices != len(names):
|
|
457
|
+
raise DFloatFormatError(
|
|
458
|
+
f"group {group!r}: holds {n_matrices} matrices but its pattern names {len(names)}"
|
|
459
|
+
)
|
|
460
|
+
groups[group] = DF11Group(name=group, matrix_names=names, path=home, tensors=tensors)
|
|
461
|
+
extras = {n: (f, headers[f][n]) for n, f in owner.items() if n not in claimed}
|
|
462
|
+
return DF11Checkpoint(root=root, config=config, groups=groups, extras=extras)
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
"""Generic block-boundary integration: placeholders, weight providers, the seam. No mflux here."""
|
|
@@ -0,0 +1,122 @@
|
|
|
1
|
+
"""Coverage bookkeeping: the extras plan, block-extra filtering, the resident set, full-block decode."""
|
|
2
|
+
|
|
3
|
+
from collections.abc import Iterable, Mapping
|
|
4
|
+
from pathlib import Path
|
|
5
|
+
|
|
6
|
+
import mlx.core as mx
|
|
7
|
+
|
|
8
|
+
from mlx_dfloat._safetensors import TensorInfo
|
|
9
|
+
from mlx_dfloat.errors import DFloatFormatError, DFloatIntegrationError
|
|
10
|
+
from mlx_dfloat.format import DF11Checkpoint, MxGroup, load_group_mx
|
|
11
|
+
from mlx_dfloat.integrate import seam
|
|
12
|
+
from mlx_dfloat.integrate.names import NameMap, Shapes
|
|
13
|
+
from mlx_dfloat.integrate.providers import WeightProvider, read_bf16
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def load_resident_set(
|
|
17
|
+
ckpt: DF11Checkpoint, names: Iterable[str] | None = None
|
|
18
|
+
) -> dict[str, MxGroup]:
|
|
19
|
+
"""Load every group (or the named ones) as evaluated ``MxGroup``s: the resident set of a run.
|
|
20
|
+
|
|
21
|
+
Raises:
|
|
22
|
+
DFloatIntegrationError: A requested name is not a group of the checkpoint.
|
|
23
|
+
"""
|
|
24
|
+
wanted = list(ckpt.groups) if names is None else list(names)
|
|
25
|
+
missing = [n for n in wanted if n not in ckpt.groups]
|
|
26
|
+
if missing:
|
|
27
|
+
raise DFloatIntegrationError(f"not groups of the checkpoint: {missing}")
|
|
28
|
+
return {name: load_group_mx(ckpt.groups[name]) for name in wanted}
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def extras_plan(
|
|
32
|
+
ckpt: DF11Checkpoint,
|
|
33
|
+
name_map: NameMap,
|
|
34
|
+
*,
|
|
35
|
+
counts: Mapping[str, int],
|
|
36
|
+
dropped: frozenset[str] = frozenset(),
|
|
37
|
+
extras: Mapping[str, tuple[Path, TensorInfo]] | None = None,
|
|
38
|
+
) -> list[tuple[str, Path, TensorInfo]]:
|
|
39
|
+
"""The extras to load as (module parameter name, file, tensor info), sorted by checkpoint name.
|
|
40
|
+
|
|
41
|
+
Block extras at an index at or beyond ``counts[kind]`` and names in ``dropped`` are left out.
|
|
42
|
+
``extras`` overrides the checkpoint's own extras (a BF16 base's index, for the reference side of
|
|
43
|
+
an image identity check) and defaults to ``ckpt.extras``.
|
|
44
|
+
|
|
45
|
+
Raises:
|
|
46
|
+
DFloatFormatError: Two checkpoint names map to the same module parameter name.
|
|
47
|
+
"""
|
|
48
|
+
plan: list[tuple[str, Path, TensorInfo]] = []
|
|
49
|
+
source_of: dict[str, str] = {}
|
|
50
|
+
source = ckpt.extras if extras is None else extras
|
|
51
|
+
for name, (path, info) in sorted(source.items()):
|
|
52
|
+
if name in dropped:
|
|
53
|
+
continue
|
|
54
|
+
block = _block_of(name, name_map)
|
|
55
|
+
if block is not None:
|
|
56
|
+
kind, idx = block
|
|
57
|
+
if idx >= counts.get(kind, 0):
|
|
58
|
+
continue
|
|
59
|
+
param = name_map.param_name(name)
|
|
60
|
+
if param in source_of:
|
|
61
|
+
raise DFloatFormatError(
|
|
62
|
+
f"extras {source_of[param]!r} and {name!r} both map to the parameter {param!r}"
|
|
63
|
+
)
|
|
64
|
+
source_of[param] = name
|
|
65
|
+
plan.append((param, path, info))
|
|
66
|
+
return plan
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
def _block_of(name: str, name_map: NameMap) -> tuple[str, int] | None:
|
|
70
|
+
head, _dot, _rest = name.partition(".")
|
|
71
|
+
if head not in name_map.kinds:
|
|
72
|
+
return None
|
|
73
|
+
parts = name.split(".")
|
|
74
|
+
index = parts[1] if len(parts) > 2 else ""
|
|
75
|
+
return (head, int(index)) if index.isascii() and index.isdigit() else None
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
def read_extra(path: Path, info: TensorInfo) -> mx.array:
|
|
79
|
+
"""Read one BF16 extra as a bf16 ``mx.array``.
|
|
80
|
+
|
|
81
|
+
Raises:
|
|
82
|
+
DFloatFormatError: The tensor is not BF16.
|
|
83
|
+
"""
|
|
84
|
+
return read_bf16(path, info)
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
def check_extras_cover(
|
|
88
|
+
params: Iterable[str], extras: Iterable[str], matrices: Iterable[str]
|
|
89
|
+
) -> None:
|
|
90
|
+
"""Every non-matrix parameter gets exactly one extra, and every planned extra has a parameter.
|
|
91
|
+
|
|
92
|
+
Raises:
|
|
93
|
+
DFloatIntegrationError: A parameter nobody sets, or an extra with no target.
|
|
94
|
+
"""
|
|
95
|
+
param_set, extra_set, matrix_set = set(params), set(extras), set(matrices)
|
|
96
|
+
missing = sorted(param_set - matrix_set - extra_set)
|
|
97
|
+
unexpected = sorted(extra_set - param_set)
|
|
98
|
+
if missing or unexpected:
|
|
99
|
+
raise DFloatIntegrationError(
|
|
100
|
+
f"extras do not cover the transformer: parameters without an extra {missing}; "
|
|
101
|
+
f"extras without a parameter {unexpected}"
|
|
102
|
+
)
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
def decode_resident(provider: WeightProvider, shapes: Shapes) -> dict[str, dict[str, mx.array]]:
|
|
106
|
+
"""Decode every block's group once into resident BF16 dicts, one block evaluated before the next.
|
|
107
|
+
|
|
108
|
+
One lazy eval over every block would allocate every decode up front (the run-ahead the per-block
|
|
109
|
+
eval policy exists to prevent), so each block's dict is evaluated as soon as it is cut. The
|
|
110
|
+
deferred status words are checked at the end.
|
|
111
|
+
|
|
112
|
+
Raises:
|
|
113
|
+
DFloatFormatError: A block's decode reported an error (``WeightProvider.verify``).
|
|
114
|
+
DFloatIntegrationError: A block is not resident, or a matrix's size does not match its shape.
|
|
115
|
+
"""
|
|
116
|
+
per_block: dict[str, dict[str, mx.array]] = {}
|
|
117
|
+
for name, per in shapes.items():
|
|
118
|
+
weights = provider.weights_for(name, per)
|
|
119
|
+
seam._eval(weights)
|
|
120
|
+
per_block[name] = weights
|
|
121
|
+
provider.verify()
|
|
122
|
+
return per_block
|
|
@@ -0,0 +1,55 @@
|
|
|
1
|
+
"""Group-size arithmetic and a phase-structured fit estimate (predictions, labelled as such)."""
|
|
2
|
+
|
|
3
|
+
from collections.abc import Mapping
|
|
4
|
+
from dataclasses import dataclass
|
|
5
|
+
from typing import Any
|
|
6
|
+
|
|
7
|
+
import mlx.core as mx
|
|
8
|
+
|
|
9
|
+
from mlx_dfloat.integrate.names import NameMap
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def decoded_bytes(group: Any) -> int:
|
|
13
|
+
"""Bytes of BF16 a group decodes to: two per element, one element per sign-mantissa byte."""
|
|
14
|
+
return 2 * int(group.tensors["sign_mantissa"].nbytes)
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def largest_decoded_bytes(ckpt: Any, name_map: NameMap) -> dict[str, int]:
|
|
18
|
+
"""The largest decoded group per block kind, from the checkpoint headers (nothing is loaded)."""
|
|
19
|
+
largest: dict[str, int] = {}
|
|
20
|
+
for name, group in ckpt.groups.items():
|
|
21
|
+
kind = name_map.kind_of(name)
|
|
22
|
+
largest[kind] = max(largest.get(kind, 0), decoded_bytes(group))
|
|
23
|
+
return largest
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def budget_bytes(*, reserve_bytes: int = 2 * 1024**3) -> int:
|
|
27
|
+
"""The fit rule's budget: the device's recommended working set minus a reserve for the OS."""
|
|
28
|
+
return int(mx.device_info()["max_recommended_working_set_size"]) - reserve_bytes
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
@dataclass(frozen=True, slots=True, kw_only=True)
|
|
32
|
+
class FitEstimate:
|
|
33
|
+
"""A predicted peak per phase against a budget; a prediction, not a measurement."""
|
|
34
|
+
|
|
35
|
+
phases: dict[str, int]
|
|
36
|
+
peak_phase: str
|
|
37
|
+
peak_bytes: int
|
|
38
|
+
budget_bytes: int
|
|
39
|
+
|
|
40
|
+
@property
|
|
41
|
+
def fits(self) -> bool:
|
|
42
|
+
"""Whether the peak phase stays within the budget."""
|
|
43
|
+
return self.peak_bytes <= self.budget_bytes
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def fit_estimate(phases: Mapping[str, Mapping[str, int]], *, budget_bytes: int) -> FitEstimate:
|
|
47
|
+
"""Sum each phase's terms; the peak is the largest phase."""
|
|
48
|
+
totals = {phase: sum(terms.values()) for phase, terms in phases.items()}
|
|
49
|
+
peak_phase = max(totals, key=totals.__getitem__)
|
|
50
|
+
return FitEstimate(
|
|
51
|
+
phases=totals,
|
|
52
|
+
peak_phase=peak_phase,
|
|
53
|
+
peak_bytes=totals[peak_phase],
|
|
54
|
+
budget_bytes=budget_bytes,
|
|
55
|
+
)
|