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,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"]
|