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,121 @@
1
+ """Name maps: how a checkpoint's matrix names land on a module's attributes."""
2
+
3
+ import re
4
+ from collections.abc import Mapping
5
+ from dataclasses import dataclass
6
+ from typing import Any, Protocol
7
+
8
+ import mlx.nn as nn
9
+
10
+ from mlx_dfloat.errors import DFloatIntegrationError
11
+
12
+ BlockShapes = dict[str, tuple[int, ...]]
13
+ Shapes = dict[str, BlockShapes]
14
+
15
+
16
+ @dataclass(frozen=True, slots=True, kw_only=True)
17
+ class Placement:
18
+ """Where one checkpoint matrix lands: the block name and the dotted attribute path of its module."""
19
+
20
+ block: str
21
+ attr: str
22
+
23
+
24
+ class NameMap(Protocol):
25
+ """The adapter-specific naming an integration needs; everything else in this package is generic."""
26
+
27
+ kinds: tuple[str, ...]
28
+
29
+ def attrs_of(self, kind: str) -> tuple[str, ...]:
30
+ """Attribute paths of every matrix module of a block of ``kind``, in the map's own order.
31
+
32
+ The order is the map's construction order, not the checkpoint's: a checkpoint group's own
33
+ concatenation order lives in its ``matrix_names``, recovered per matrix through ``place()``.
34
+ """
35
+ ...
36
+
37
+ def place(self, matrix_name: str) -> Placement:
38
+ """The block and attribute path of a checkpoint matrix name."""
39
+ ...
40
+
41
+ def kind_of(self, block_name: str) -> str:
42
+ """The kind of a block name."""
43
+ ...
44
+
45
+ def param_name(self, checkpoint_name: str) -> str:
46
+ """The module parameter name of a non-matrix checkpoint tensor (biases, norms, embedders)."""
47
+ ...
48
+
49
+ def is_matrix_module(self, module: Any) -> bool:
50
+ """Whether ``module`` is one the seam swaps weights on."""
51
+ ...
52
+
53
+
54
+ class StaticNameMap:
55
+ """A name map from explicit tables: kind -> {checkpoint sub-path: attribute path}.
56
+
57
+ Block names are ``<kind>.<index>``; a matrix name is ``<block>.<sub-path>.weight``; matrix modules are
58
+ ``nn.Linear``.
59
+ """
60
+
61
+ def __init__(self, tables: Mapping[str, Mapping[str, str]]) -> None:
62
+ """Keep one table per kind, checkpoint sub-path -> attribute path, in the given order."""
63
+ self._tables = {k: dict(v) for k, v in tables.items()}
64
+ self.kinds = tuple(self._tables)
65
+ # `[0-9]`, not `\d`: `\d` matches any Unicode digit, and int("٣") == 3 would alias block 3.
66
+ self._block = re.compile(
67
+ rf"^({'|'.join(re.escape(k) for k in self.kinds)})\.([0-9]+)\.(.+)$"
68
+ )
69
+
70
+ def attrs_of(self, kind: str) -> tuple[str, ...]:
71
+ """Attribute paths of every matrix module of a block of ``kind``, in table order."""
72
+ return tuple(self._tables[kind].values())
73
+
74
+ def _split(self, name: str) -> tuple[str, int, str] | None:
75
+ match = self._block.match(name)
76
+ return None if match is None else (match.group(1), int(match.group(2)), match.group(3))
77
+
78
+ def place(self, matrix_name: str) -> Placement:
79
+ """The block and attribute path of a checkpoint matrix name.
80
+
81
+ Raises:
82
+ DFloatIntegrationError: The name is not a ``<kind>.<index>.<sub-path>.weight`` name the
83
+ table covers.
84
+ """
85
+ parsed = self._split(matrix_name)
86
+ if parsed is None or not parsed[2].endswith(".weight"):
87
+ raise DFloatIntegrationError(f"{matrix_name!r} is not a block matrix name")
88
+ kind, idx, rest = parsed
89
+ attr = self._tables[kind].get(rest.removesuffix(".weight"))
90
+ if attr is None:
91
+ raise DFloatIntegrationError(f"{matrix_name!r}: not a {kind} matrix the map covers")
92
+ return Placement(block=f"{kind}.{idx}", attr=attr)
93
+
94
+ def kind_of(self, block_name: str) -> str:
95
+ """The kind of a ``<kind>.<index>`` block name.
96
+
97
+ Raises:
98
+ DFloatIntegrationError: ``block_name`` is not a block name of a known kind.
99
+ """
100
+ kind, _dot, idx = block_name.partition(".")
101
+ if kind not in self._tables or not (idx.isascii() and idx.isdigit()):
102
+ raise DFloatIntegrationError(f"{block_name!r} is not a block name")
103
+ return kind
104
+
105
+ def param_name(self, checkpoint_name: str) -> str:
106
+ """The module parameter name of a non-matrix checkpoint tensor.
107
+
108
+ A block extra is renamed through the same table as its block's matrices; every other name
109
+ (already the module's own) is returned unchanged.
110
+ """
111
+ parsed = self._split(checkpoint_name)
112
+ if parsed is None:
113
+ return checkpoint_name
114
+ kind, idx, rest = parsed
115
+ head, _dot, leaf = rest.rpartition(".")
116
+ mapped = self._tables[kind].get(head)
117
+ return checkpoint_name if mapped is None else f"{kind}.{idx}.{mapped}.{leaf}"
118
+
119
+ def is_matrix_module(self, module: Any) -> bool:
120
+ """Whether ``module`` is an ``nn.Linear``, the only matrix module this map swaps weights on."""
121
+ return isinstance(module, nn.Linear) # type: ignore[attr-defined]
@@ -0,0 +1,72 @@
1
+ """Zero-size placeholders in place of every block matrix, and the shapes they had."""
2
+
3
+ from collections.abc import Iterable, Sequence
4
+ from typing import Any
5
+
6
+ import mlx.core as mx
7
+
8
+ from mlx_dfloat.errors import DFloatIntegrationError
9
+ from mlx_dfloat.integrate.names import BlockShapes, NameMap, Shapes
10
+
11
+ PLACEHOLDER = mx.zeros((0,), dtype=mx.bfloat16)
12
+
13
+
14
+ def _is_index(part: str) -> bool:
15
+ """An ASCII-digit component (``"0"``); ``"²".isdigit()`` is true too, but ``int("²")`` raises."""
16
+ return part.isascii() and part.isdigit()
17
+
18
+
19
+ def get_attr_path(module: Any, path: str) -> Any:
20
+ """Resolve a dotted path on a module; a digit component indexes a list (``attn.to_out.0``).
21
+
22
+ Raises:
23
+ DFloatIntegrationError: A component does not exist.
24
+ """
25
+ node = module
26
+ for part in path.split("."):
27
+ try:
28
+ node = node[int(part)] if _is_index(part) else getattr(node, part)
29
+ except (AttributeError, IndexError, KeyError, TypeError) as exc:
30
+ raise DFloatIntegrationError(f"{path!r}: no {part!r} on {type(node).__name__}") from exc
31
+ return node
32
+
33
+
34
+ def set_attr_path(module: Any, path: str, value: Any) -> None:
35
+ """Assign ``value`` at a dotted path (see ``get_attr_path``)."""
36
+ head, _dot, leaf = path.rpartition(".")
37
+ parent = get_attr_path(module, head) if head else module
38
+ if _is_index(leaf):
39
+ parent[int(leaf)] = value
40
+ else:
41
+ setattr(parent, leaf, value)
42
+
43
+
44
+ def install_placeholders(
45
+ block_lists: Iterable[tuple[str, Sequence[Any]]], name_map: NameMap
46
+ ) -> Shapes:
47
+ """Replace every block matrix weight with ``PLACEHOLDER``; return block name -> attr -> shape.
48
+
49
+ Every mapped attribute must be a matrix module of the block and every matrix module of the block must
50
+ be mapped, so an added or renamed layer is caught here, not at the first matmul.
51
+
52
+ Raises:
53
+ DFloatIntegrationError: A block's matrix-module set differs from the map.
54
+ """
55
+ shapes: Shapes = {}
56
+ for kind, blocks in block_lists:
57
+ mapped = set(name_map.attrs_of(kind))
58
+ for idx, block in enumerate(blocks):
59
+ name = f"{kind}.{idx}"
60
+ present = {p for p, m in block.named_modules() if name_map.is_matrix_module(m)}
61
+ if present != mapped:
62
+ raise DFloatIntegrationError(
63
+ f"{name}: mapped matrices missing from the block: {sorted(mapped - present)}; "
64
+ f"Linear layers the map does not cover: {sorted(present - mapped)}"
65
+ )
66
+ per: BlockShapes = {}
67
+ for attr in name_map.attrs_of(kind):
68
+ module = get_attr_path(block, attr)
69
+ per[attr] = tuple(int(d) for d in module.weight.shape)
70
+ module.weight = PLACEHOLDER
71
+ shapes[name] = per
72
+ return shapes
@@ -0,0 +1,271 @@
1
+ """Weight providers: what the seam asks for one block's matrices."""
2
+
3
+ from collections.abc import Callable, Mapping, Sequence
4
+ from functools import partial
5
+ from pathlib import Path
6
+ from typing import Protocol, cast
7
+
8
+ import mlx.core as mx
9
+ import numpy as np
10
+
11
+ from mlx_dfloat._safetensors import TensorInfo, read_array
12
+ from mlx_dfloat.decode import DecodeResult, check_status, decode_group, split_matrices
13
+ from mlx_dfloat.errors import DFloatFormatError, DFloatIntegrationError
14
+ from mlx_dfloat.format import MxGroup
15
+ from mlx_dfloat.integrate.names import BlockShapes, NameMap
16
+
17
+ # "per-block": evaluate each block's output. "depth2": async_eval it, then eval the previous
18
+ # block's (hides host-side allocation while the GPU runs the block; not a decode/compute overlap).
19
+ # "none": no evaluation inside the step (a non-launching provider only, or a few blocks).
20
+ EVAL_POLICIES: tuple[str, ...] = ("per-block", "depth2", "none")
21
+
22
+
23
+ class WeightProvider(Protocol):
24
+ """Hands the seam one block's matrices as ``{attribute path: bf16 array of the block's shape}``."""
25
+
26
+ launches: int
27
+ launching: bool
28
+ policies: tuple[str, ...]
29
+
30
+ def weights_for(self, block_name: str, shapes: BlockShapes) -> dict[str, mx.array]:
31
+ """Weights for ``block_name``; ``shapes`` is that block's entry from ``install_placeholders``."""
32
+ ...
33
+
34
+ def verify(self) -> None:
35
+ """Check what ``weights_for`` deferred (decode status words); call it after the step's eval."""
36
+ ...
37
+
38
+ def reset(self) -> None:
39
+ """Drop per-step state after a step that raised."""
40
+ ...
41
+
42
+
43
+ def read_bf16(path: Path, info: TensorInfo) -> mx.array:
44
+ """One BF16 tensor of a safetensors file as a bf16 array (a bit view, never a cast).
45
+
46
+ Raises:
47
+ DFloatFormatError: The tensor is not BF16.
48
+ """
49
+ if info.dtype != "BF16":
50
+ raise DFloatFormatError(f"{info.name}: tensor is {info.dtype}, expected BF16")
51
+ return mx.array(np.ascontiguousarray(read_array(path, info))).view(mx.bfloat16)
52
+
53
+
54
+ class StreamingBF16Provider:
55
+ """Reads each block's BF16 matrices from safetensors shards as the block runs; nothing stays resident.
56
+
57
+ ``index`` maps a checkpoint matrix name to ``(shard path, tensor info)``: a base repository's
58
+ weight index plus the shards' headers. The reference side of an image identity check, and the
59
+ seam's way of streaming an uncompressed model.
60
+ """
61
+
62
+ launching = False
63
+ # Each block reads fresh arrays from the shards; under "none" the whole BF16 transformer would
64
+ # stay alive in one step's lazy graph, so only the per-block boundaries are offered.
65
+ policies: tuple[str, ...] = ("per-block", "depth2")
66
+
67
+ def __init__(
68
+ self,
69
+ index: Mapping[str, tuple[Path, TensorInfo]],
70
+ matrix_names: Mapping[str, Sequence[str]],
71
+ name_map: NameMap,
72
+ ) -> None:
73
+ """Hold the index, each block's matrix names (the checkpoint's order) and the map that places them."""
74
+ self._index = index
75
+ self._matrix_names = matrix_names
76
+ self._name_map = name_map
77
+ self.launches = 0
78
+ self.reads = 0
79
+
80
+ def verify(self) -> None:
81
+ """Nothing deferred: no decode happened."""
82
+
83
+ def reset(self) -> None:
84
+ """No per-step state."""
85
+
86
+ def weights_for(self, block_name: str, shapes: BlockShapes) -> dict[str, mx.array]:
87
+ """Read the block's matrices from their shards (fresh arrays; the caller's step frees them).
88
+
89
+ Raises:
90
+ DFloatIntegrationError: No matrix names for the block, a name missing from the index, a
91
+ name that is not a matrix of this block, or a tensor of the wrong shape.
92
+ DFloatFormatError: A tensor that is not BF16.
93
+ """
94
+ names = self._matrix_names.get(block_name)
95
+ if names is None:
96
+ raise DFloatIntegrationError(f"{block_name}: no matrix names")
97
+ weights: dict[str, mx.array] = {}
98
+ for matrix_name in names:
99
+ entry = self._index.get(matrix_name)
100
+ if entry is None:
101
+ raise DFloatIntegrationError(f"{matrix_name}: not in the BF16 index")
102
+ placement = self._name_map.place(matrix_name)
103
+ if placement.block != block_name or placement.attr not in shapes:
104
+ raise DFloatIntegrationError(
105
+ f"{matrix_name}: not a matrix of {block_name} with a known shape"
106
+ )
107
+ array = read_bf16(*entry)
108
+ self.reads += 1
109
+ if tuple(array.shape) != tuple(shapes[placement.attr]):
110
+ raise DFloatIntegrationError(
111
+ f"{matrix_name}: shard tensor has shape {tuple(array.shape)}, "
112
+ f"the block needs {shapes[placement.attr]}"
113
+ )
114
+ weights[placement.attr] = array
115
+ return weights
116
+
117
+
118
+ class DF11Provider:
119
+ """Decodes each block's DF11 group when asked: one decode launch per call.
120
+
121
+ The status words are not read inside ``weights_for`` (a host read there would sync before every
122
+ block); they queue in ``pending`` and ``verify`` reads them after the step's final eval.
123
+ """
124
+
125
+ launching = True
126
+ policies = EVAL_POLICIES
127
+
128
+ def __init__(
129
+ self,
130
+ resident: Mapping[str, MxGroup],
131
+ matrix_names: Mapping[str, Sequence[str]],
132
+ name_map: NameMap,
133
+ *,
134
+ decode: Callable[[MxGroup], DecodeResult] | None = None,
135
+ ) -> None:
136
+ """Hold the resident groups, each group's matrix names, and the name map that places them.
137
+
138
+ ``decode`` defaults to the Metal backend; tests inject a counting reference decode.
139
+ """
140
+ self._resident = resident
141
+ self._matrix_names = matrix_names
142
+ self._name_map = name_map
143
+ self._decode = decode if decode is not None else partial(decode_group, backend="metal")
144
+ self.launches = 0
145
+ self.pending: list[tuple[str, mx.array]] = []
146
+
147
+ def verify(self) -> None:
148
+ """Read every pending status word on the host and clear the list.
149
+
150
+ Raises:
151
+ DFloatFormatError: A block's decode reported an error; the message names the block.
152
+ """
153
+ pending, self.pending = self.pending, []
154
+ for block_name, status in pending:
155
+ check_status(status, name=block_name)
156
+
157
+ def reset(self) -> None:
158
+ """Forget queued status words (their decodes belong to a step that did not finish)."""
159
+ self.pending = []
160
+
161
+ def weights_for(self, block_name: str, shapes: BlockShapes) -> dict[str, mx.array]:
162
+ """Decode the block's group and cut it into views of the block's shapes (no copies, no host read).
163
+
164
+ Raises:
165
+ DFloatIntegrationError: The block is not resident or has no matrix names (both refused before
166
+ any decode), or the names do not fit the decoded group: their count, their block, or a
167
+ matrix's size against its shape.
168
+ """
169
+ group = self._resident.get(block_name)
170
+ if group is None:
171
+ raise DFloatIntegrationError(f"{block_name}: no resident DF11 group")
172
+ names = self._matrix_names.get(block_name)
173
+ if names is None:
174
+ raise DFloatIntegrationError(f"{block_name}: no matrix names for its DF11 group")
175
+ result = self._decode(group)
176
+ self.pending.append((block_name, result.status))
177
+ parts = split_matrices(result.bits, group.split_positions)
178
+ if len(names) != len(parts):
179
+ raise DFloatIntegrationError(
180
+ f"{block_name}: {len(parts)} matrices decoded for {len(names)} names"
181
+ )
182
+ weights: dict[str, mx.array] = {}
183
+ for matrix_name, part in zip(names, parts, strict=True):
184
+ placement = self._name_map.place(matrix_name)
185
+ if placement.block != block_name or placement.attr not in shapes:
186
+ raise DFloatIntegrationError(
187
+ f"{matrix_name}: not a matrix of {block_name} with a known shape"
188
+ )
189
+ shape = shapes[placement.attr]
190
+ n = 1
191
+ for d in shape:
192
+ n *= d
193
+ if part.size != n:
194
+ raise DFloatIntegrationError(
195
+ f"{matrix_name}: decoded {part.size} elements, expected {shape} = {n}"
196
+ )
197
+ weights[placement.attr] = part.view(mx.bfloat16).reshape(shape)
198
+ self.launches += 1
199
+ return weights
200
+
201
+
202
+ class ReuseProvider:
203
+ """One pre-decoded dict per block kind, returned for every block of that kind; never launches."""
204
+
205
+ launching = False
206
+ policies = EVAL_POLICIES
207
+
208
+ def __init__(self, per_kind: Mapping[str, Mapping[str, mx.array]], name_map: NameMap) -> None:
209
+ """Keep one dict per kind, handed back as-is (by identity) for every block of that kind."""
210
+ self._per_kind: dict[str, Mapping[str, mx.array]] = dict(per_kind)
211
+ self._name_map = name_map
212
+ self.launches = 0
213
+
214
+ def verify(self) -> None:
215
+ """Nothing deferred: no decode happened."""
216
+
217
+ def reset(self) -> None:
218
+ """No per-step state."""
219
+
220
+ def weights_for(self, block_name: str, shapes: BlockShapes) -> dict[str, mx.array]:
221
+ """The kind's dict, checked against the block's shapes.
222
+
223
+ Raises:
224
+ DFloatIntegrationError: No dict for the block's kind, or a shape differs.
225
+ """
226
+ kind = self._name_map.kind_of(block_name)
227
+ weights = self._per_kind.get(kind)
228
+ if weights is None:
229
+ raise DFloatIntegrationError(f"{block_name}: no reusable block of kind {kind!r}")
230
+ _check_shapes(block_name, weights, shapes)
231
+ return cast("dict[str, mx.array]", weights)
232
+
233
+
234
+ class ResidentProvider:
235
+ """Per-block pre-decoded dicts (every block resident at once); never launches."""
236
+
237
+ launching = False
238
+ policies = EVAL_POLICIES
239
+
240
+ def __init__(self, per_block: Mapping[str, Mapping[str, mx.array]]) -> None:
241
+ """Keep the per-block dicts, handed back as-is (by identity) for their own block."""
242
+ self._per_block: dict[str, Mapping[str, mx.array]] = dict(per_block)
243
+ self.launches = 0
244
+
245
+ def verify(self) -> None:
246
+ """Nothing deferred: no decode happened."""
247
+
248
+ def reset(self) -> None:
249
+ """No per-step state."""
250
+
251
+ def weights_for(self, block_name: str, shapes: BlockShapes) -> dict[str, mx.array]:
252
+ """The block's own dict, checked against its shapes.
253
+
254
+ Raises:
255
+ DFloatIntegrationError: The block is not resident, or a shape differs.
256
+ """
257
+ weights = self._per_block.get(block_name)
258
+ if weights is None:
259
+ raise DFloatIntegrationError(f"{block_name}: not resident")
260
+ _check_shapes(block_name, weights, shapes)
261
+ return cast("dict[str, mx.array]", weights)
262
+
263
+
264
+ def _check_shapes(block_name: str, weights: Mapping[str, mx.array], shapes: BlockShapes) -> None:
265
+ for attr, shape in shapes.items():
266
+ if attr not in weights:
267
+ raise DFloatIntegrationError(f"{block_name}: no weight for {attr!r}")
268
+ if tuple(weights[attr].shape) != tuple(shape):
269
+ raise DFloatIntegrationError(
270
+ f"{block_name}.{attr}: weight has shape {tuple(weights[attr].shape)}, the block needs {shape}"
271
+ )
@@ -0,0 +1,196 @@
1
+ """The block seam: assign one block's weights, run it, evaluate per policy, restore the placeholders."""
2
+
3
+ import time
4
+ from collections.abc import Callable
5
+ from dataclasses import dataclass
6
+ from typing import Any, Protocol
7
+
8
+ import mlx.core as mx
9
+
10
+ from mlx_dfloat.errors import DFloatIntegrationError
11
+ from mlx_dfloat.integrate.names import Shapes
12
+ from mlx_dfloat.integrate.placeholders import PLACEHOLDER, get_attr_path
13
+ from mlx_dfloat.integrate.providers import EVAL_POLICIES, WeightProvider
14
+
15
+ # "none" with a launching provider keeps every decoded group alive
16
+ MAX_NONE_POLICY_LAUNCHING_BLOCKS = 2
17
+
18
+ # The seam calls MLX through these names so a test can record the evaluation order.
19
+ _eval = mx.eval
20
+ _async_eval = mx.async_eval
21
+
22
+
23
+ class TraceHook(Protocol):
24
+ """Receives one event per block per step and the step boundaries; four clock reads per block, no sync."""
25
+
26
+ def begin_step(self) -> None:
27
+ """Open a new step (called once, before its first block runs)."""
28
+ ...
29
+
30
+ def end_step(self) -> None:
31
+ """Close the step (called once, after its last block's eval)."""
32
+ ...
33
+
34
+ def record(
35
+ self,
36
+ block: str,
37
+ *,
38
+ decode_start: float,
39
+ decode_end: float,
40
+ encode_end: float,
41
+ eval_end: float,
42
+ ) -> None:
43
+ """Record one block's four clock reads: decode, the mflux call, and the policy's eval."""
44
+ ...
45
+
46
+
47
+ @dataclass(slots=True)
48
+ class SeamState:
49
+ """Mutable per-attachment state the seam threads through a step: never rebuilt mid-step."""
50
+
51
+ provider: WeightProvider
52
+ shapes: Shapes
53
+ policy: str
54
+ tracer: TraceHook | None = None
55
+ verify_in_call: bool = False
56
+ prev: Any = None # depth2: the previous block's output, evaluated after the next is queued
57
+
58
+
59
+ def attach_state(
60
+ provider: WeightProvider,
61
+ shapes: Shapes,
62
+ *,
63
+ eval_policy: str = "per-block",
64
+ tracer: TraceHook | None = None,
65
+ verify_in_call: bool = False,
66
+ ) -> SeamState:
67
+ """Validate the policy for this provider and build the seam's state.
68
+
69
+ Raises:
70
+ DFloatIntegrationError: Unknown policy; a policy the provider does not run under; ``"none"`` with a
71
+ launching provider over more than ``MAX_NONE_POLICY_LAUNCHING_BLOCKS`` blocks; or
72
+ ``verify_in_call`` with ``"none"`` and a launching provider (reading the status words at the
73
+ end of the call would force every decode before the caller's own eval).
74
+ """
75
+ if eval_policy not in EVAL_POLICIES:
76
+ raise DFloatIntegrationError(
77
+ f"unknown eval policy {eval_policy!r}; choose from {EVAL_POLICIES}"
78
+ )
79
+ if eval_policy not in provider.policies:
80
+ raise DFloatIntegrationError(
81
+ f"{type(provider).__name__} runs only under {provider.policies}, not {eval_policy!r}"
82
+ )
83
+ if (
84
+ eval_policy == "none"
85
+ and provider.launching
86
+ and len(shapes) > MAX_NONE_POLICY_LAUNCHING_BLOCKS
87
+ ):
88
+ raise DFloatIntegrationError(
89
+ f"eval policy 'none' with a launching provider over {len(shapes)} blocks would keep every decoded "
90
+ f"group resident until the final eval; 'none' is for non-launching providers (or at most "
91
+ f"{MAX_NONE_POLICY_LAUNCHING_BLOCKS} blocks)"
92
+ )
93
+ if verify_in_call and eval_policy == "none" and provider.launching:
94
+ raise DFloatIntegrationError(
95
+ "verify_in_call with eval policy 'none' and a launching provider would read the status "
96
+ "words before the caller's eval, forcing every decode of the step; verify_step() after "
97
+ "the step's eval instead"
98
+ )
99
+ return SeamState(
100
+ provider=provider,
101
+ shapes=shapes,
102
+ policy=eval_policy,
103
+ tracer=tracer,
104
+ verify_in_call=verify_in_call,
105
+ )
106
+
107
+
108
+ def begin_step(state: SeamState) -> None:
109
+ """Refuse a step while the last one's status words are unchecked; open the tracer's step.
110
+
111
+ Raises:
112
+ DFloatIntegrationError: The provider still holds status words from an earlier step (nobody
113
+ called ``verify_step()`` after it); they are kept, so a ``verify_step()`` can still check them.
114
+ """
115
+ pending = getattr(state.provider, "pending", None)
116
+ if pending:
117
+ raise DFloatIntegrationError(
118
+ f"{len(pending)} decode status words from an earlier step are unchecked; call "
119
+ "verify_step() after each step's eval, or attach with verify_in_call=True"
120
+ )
121
+ if state.tracer is not None:
122
+ state.tracer.begin_step()
123
+
124
+
125
+ def end_step(state: SeamState, out: Any) -> Any:
126
+ """Drain depth-2's tail, verify when asked, close the trace; returns ``out``."""
127
+ if state.prev is not None:
128
+ _eval(state.prev)
129
+ state.prev = None
130
+ if state.verify_in_call:
131
+ state.provider.verify()
132
+ if state.tracer is not None:
133
+ state.tracer.end_step()
134
+ return out
135
+
136
+
137
+ def abort_step(state: SeamState) -> None:
138
+ """After a step that raised: no stale look-ahead, no stale status words, no depth-2 tail."""
139
+ state.prev = None
140
+ state.provider.reset()
141
+
142
+
143
+ def verify_step(state: SeamState) -> None:
144
+ """Run the provider's deferred checks (an explicit call, for a step run without ``verify_in_call``)."""
145
+ state.provider.verify()
146
+
147
+
148
+ def _seam_eval(state: SeamState, out: Any) -> None:
149
+ if state.policy == "per-block":
150
+ _eval(out)
151
+ elif state.policy == "depth2":
152
+ _async_eval(out) # type: ignore[no-untyped-call] # mlx's async_eval stub carries no annotations
153
+ if state.prev is not None:
154
+ _eval(state.prev)
155
+ state.prev = out
156
+
157
+
158
+ def run_block(state: SeamState, block_name: str, block: Any, run: Callable[[], Any]) -> Any:
159
+ """Assign the block's weights, run it, evaluate per policy, restore the placeholders (also on a raise).
160
+
161
+ Raises:
162
+ DFloatIntegrationError: The provider's dict does not cover the block's matrices, or a weight's
163
+ shape does not match the block.
164
+ """
165
+ shapes = state.shapes[block_name]
166
+ t_decode_start = time.perf_counter()
167
+ weights = state.provider.weights_for(block_name, shapes)
168
+ t_decode_end = time.perf_counter()
169
+ if weights.keys() != shapes.keys():
170
+ raise DFloatIntegrationError(
171
+ f"{block_name}: provider returned {sorted(weights)}; the block needs {sorted(shapes)} "
172
+ f"(missing {sorted(shapes.keys() - weights.keys())})"
173
+ )
174
+ for attr, shape in shapes.items():
175
+ if tuple(weights[attr].shape) != tuple(shape):
176
+ raise DFloatIntegrationError(
177
+ f"{block_name}.{attr}: weight has shape {tuple(weights[attr].shape)}, the block needs {shape}"
178
+ )
179
+ try:
180
+ for attr, weight in weights.items():
181
+ get_attr_path(block, attr).weight = weight
182
+ out = run()
183
+ t_encode_end = time.perf_counter()
184
+ _seam_eval(state, out)
185
+ if state.tracer is not None:
186
+ state.tracer.record(
187
+ block_name,
188
+ decode_start=t_decode_start,
189
+ decode_end=t_decode_end,
190
+ encode_end=t_encode_end,
191
+ eval_end=time.perf_counter(),
192
+ )
193
+ return out
194
+ finally:
195
+ for attr in shapes:
196
+ get_attr_path(block, attr).weight = PLACEHOLDER
@@ -0,0 +1,31 @@
1
+ """mflux adapters. Everything here imports mflux lazily; ``import mlx_dfloat`` never needs it."""
2
+
3
+ import importlib.util
4
+ from typing import Any
5
+
6
+ from mlx_dfloat.errors import DFloatDependencyError
7
+
8
+
9
+ def require_mflux() -> None:
10
+ """Raise a package-rooted error when the optional ``mflux`` extra is missing.
11
+
12
+ Raises:
13
+ DFloatDependencyError: ``mflux`` cannot be imported.
14
+ """
15
+ if importlib.util.find_spec("mflux") is None:
16
+ raise DFloatDependencyError(
17
+ "the mflux adapters need the optional extra: install mlx-dfloat[mflux]"
18
+ )
19
+
20
+
21
+ def __getattr__(name: str) -> Any:
22
+ """``DFloatFlux1`` is imported on first use, so ``import mlx_dfloat.mflux`` never needs mflux."""
23
+ if name == "DFloatFlux1":
24
+ from mlx_dfloat.mflux.flux1.model import DFloatFlux1
25
+
26
+ return DFloatFlux1
27
+ raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
28
+
29
+
30
+ # DFloatFlux1 stays out of __all__: a star-import must not need mflux; `mlx_dfloat.mflux.DFloatFlux1` still works.
31
+ __all__ = ["require_mflux"]