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/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
+ )