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,251 @@
|
|
|
1
|
+
"""README fragments rendered from result data, and the marker splice.
|
|
2
|
+
|
|
3
|
+
Pure rendering: bytes stay ``int`` until a fragment is printed, output is deterministic
|
|
4
|
+
(the caption's date is passed in), and every fragment ends with a newline.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
import re
|
|
8
|
+
from collections.abc import Mapping, Sequence
|
|
9
|
+
from dataclasses import dataclass
|
|
10
|
+
from typing import Any
|
|
11
|
+
|
|
12
|
+
from mlx_dfloat.bench.capped import GIB
|
|
13
|
+
from mlx_dfloat.bench.results import Summary
|
|
14
|
+
from mlx_dfloat.errors import DFloatFormatError
|
|
15
|
+
|
|
16
|
+
_MODEL_NAMES = {"schnell": "FLUX.1-schnell", "dev": "FLUX.1-dev", "krea-dev": "FLUX.1-Krea-dev"}
|
|
17
|
+
_NOT_MEASURED = "not measured"
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
@dataclass(frozen=True, slots=True, kw_only=True)
|
|
21
|
+
class TierRow:
|
|
22
|
+
"""One row of the tier table."""
|
|
23
|
+
|
|
24
|
+
mac_gb: int
|
|
25
|
+
ceiling_bytes: int
|
|
26
|
+
model: str
|
|
27
|
+
df11_bytes: int
|
|
28
|
+
watched_peak_bytes: int
|
|
29
|
+
footprint_peak_bytes: int
|
|
30
|
+
mlx_peak_bytes: int
|
|
31
|
+
label: str
|
|
32
|
+
status: str
|
|
33
|
+
limits_note: str
|
|
34
|
+
source: str
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
@dataclass(frozen=True, slots=True, kw_only=True)
|
|
38
|
+
class ProofRecord:
|
|
39
|
+
"""The harness proof: one passing run and one aborted run under the same cap."""
|
|
40
|
+
|
|
41
|
+
cap_bytes: int
|
|
42
|
+
pass_size: int
|
|
43
|
+
pass_watched_peak: int
|
|
44
|
+
abort_size: int
|
|
45
|
+
abort_reason: str
|
|
46
|
+
abort_counter: str
|
|
47
|
+
abort_peak: int
|
|
48
|
+
source_dir: str
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def _need(mapping: Mapping[str, Any], key: str, where: str) -> Any:
|
|
52
|
+
if key not in mapping:
|
|
53
|
+
raise DFloatFormatError(f"{where} has no {key!r} field")
|
|
54
|
+
return mapping[key]
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def tier_row_from_generate_report(report: Mapping[str, Any], *, source: str) -> TierRow:
|
|
58
|
+
"""Build a tier-table row from an ``mlx-dfloat generate --report`` JSON.
|
|
59
|
+
|
|
60
|
+
``ceiling_bytes`` is the tier's fit budget (``limits.tier.ceiling_bytes``: the recommended
|
|
61
|
+
working set minus the reserve), which the status is judged against. On the host tier the
|
|
62
|
+
watchdog's own abort line is higher (RAM minus 4 GiB); on a CAPPED tier the two are the same.
|
|
63
|
+
|
|
64
|
+
Raises:
|
|
65
|
+
DFloatFormatError: A field is missing, or the report is a harness proof.
|
|
66
|
+
"""
|
|
67
|
+
label = _need(report, "label", "report")
|
|
68
|
+
if label == "PROOF":
|
|
69
|
+
raise DFloatFormatError("a PROOF report belongs to the harness proof, not the tier table")
|
|
70
|
+
limits = _need(report, "limits", "report")
|
|
71
|
+
tier = _need(limits, "tier", "limits")
|
|
72
|
+
sizes = _need(report, "sizes", "report")
|
|
73
|
+
ceiling = int(_need(tier, "ceiling_bytes", "limits.tier"))
|
|
74
|
+
watched = int(_need(report, "watched_peak_bytes", "report"))
|
|
75
|
+
return TierRow(
|
|
76
|
+
mac_gb=int(_need(tier, "tier_gb", "limits.tier")),
|
|
77
|
+
ceiling_bytes=ceiling,
|
|
78
|
+
model=str(_need(report, "model", "report")),
|
|
79
|
+
df11_bytes=int(_need(sizes, "compressed", "sizes")) + int(_need(sizes, "extras", "sizes")),
|
|
80
|
+
watched_peak_bytes=watched,
|
|
81
|
+
footprint_peak_bytes=int(_need(report, "footprint_peak_bytes", "report")),
|
|
82
|
+
# The watchdog's MLX peak (active + cache), not a phase's active-only mx.get_peak_memory.
|
|
83
|
+
mlx_peak_bytes=int(_need(report, "mlx_peak_bytes", "report")),
|
|
84
|
+
label=str(label),
|
|
85
|
+
status="target" if watched <= ceiling else "over",
|
|
86
|
+
limits_note="host caps"
|
|
87
|
+
if _need(limits, "applied", "limits") == "host-caps"
|
|
88
|
+
else "MLX defaults for the tier",
|
|
89
|
+
source=source,
|
|
90
|
+
)
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
def proof_from_files(
|
|
94
|
+
pass_report: Mapping[str, Any],
|
|
95
|
+
abort_artifact: Mapping[str, Any],
|
|
96
|
+
*,
|
|
97
|
+
source_dir: str,
|
|
98
|
+
) -> ProofRecord:
|
|
99
|
+
"""Pair the passing run with the aborted run of the harness proof.
|
|
100
|
+
|
|
101
|
+
The aborted run's size comes from the artifact's own record (``context.height`` and
|
|
102
|
+
``context.width``, which ``generate`` hands its watchdog), never from a default.
|
|
103
|
+
|
|
104
|
+
Raises:
|
|
105
|
+
DFloatFormatError: The pass report is not a PROOF or is not square; the artifact has no
|
|
106
|
+
run context, or its run is not square; the two ran under different caps, for different
|
|
107
|
+
models, or (when both name one) with different seeds.
|
|
108
|
+
"""
|
|
109
|
+
if _need(pass_report, "label", "pass report") != "PROOF":
|
|
110
|
+
raise DFloatFormatError("the pass report's label is not PROOF")
|
|
111
|
+
height = _need(pass_report, "height", "pass report")
|
|
112
|
+
if height != _need(pass_report, "width", "pass report"):
|
|
113
|
+
raise DFloatFormatError("the pass report is not square (height != width)")
|
|
114
|
+
cap = int(_need(pass_report, "memory_ceiling_bytes", "pass report"))
|
|
115
|
+
if cap != int(_need(abort_artifact, "ceiling", "abort artifact")):
|
|
116
|
+
raise DFloatFormatError(
|
|
117
|
+
"the proof requires the same cap: pass report and abort artifact differ"
|
|
118
|
+
)
|
|
119
|
+
context = abort_artifact.get("context")
|
|
120
|
+
if not isinstance(context, Mapping):
|
|
121
|
+
raise DFloatFormatError(
|
|
122
|
+
"the abort artifact records no run context (model, height, width): it was written "
|
|
123
|
+
"by a watchdog that was not told the run's size"
|
|
124
|
+
)
|
|
125
|
+
abort_height = _need(context, "height", "abort artifact context")
|
|
126
|
+
if abort_height != _need(context, "width", "abort artifact context"):
|
|
127
|
+
raise DFloatFormatError("the aborted run is not square (height != width)")
|
|
128
|
+
model = _need(pass_report, "model", "pass report")
|
|
129
|
+
if model != _need(context, "model", "abort artifact context"):
|
|
130
|
+
raise DFloatFormatError(
|
|
131
|
+
f"the proof requires the same model: the pass report ran {model!r}, "
|
|
132
|
+
f"the aborted run {context['model']!r}"
|
|
133
|
+
)
|
|
134
|
+
if "seed" in pass_report and "seed" in context and pass_report["seed"] != context["seed"]:
|
|
135
|
+
raise DFloatFormatError(
|
|
136
|
+
f"the proof requires the same seed: the pass report ran {pass_report['seed']!r}, "
|
|
137
|
+
f"the aborted run {context['seed']!r}"
|
|
138
|
+
)
|
|
139
|
+
return ProofRecord(
|
|
140
|
+
cap_bytes=cap,
|
|
141
|
+
pass_size=int(height),
|
|
142
|
+
pass_watched_peak=int(_need(pass_report, "watched_peak_bytes", "pass report")),
|
|
143
|
+
abort_size=int(abort_height),
|
|
144
|
+
abort_reason=str(_need(abort_artifact, "reason", "abort artifact")),
|
|
145
|
+
abort_counter=str(_need(abort_artifact, "verdict_counter", "abort artifact")),
|
|
146
|
+
abort_peak=int(_need(abort_artifact, "peak_watched", "abort artifact")),
|
|
147
|
+
source_dir=source_dir,
|
|
148
|
+
)
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
def _gib(n: int) -> str:
|
|
152
|
+
return f"{n / GIB:.2f} GiB"
|
|
153
|
+
|
|
154
|
+
|
|
155
|
+
def _pct(x: float | None) -> str:
|
|
156
|
+
return _NOT_MEASURED if x is None else f"{x * 100:+.1f} %"
|
|
157
|
+
|
|
158
|
+
|
|
159
|
+
_TIER_HEADER = (
|
|
160
|
+
"| Mac | Fit budget (budget − reserve) | Model | DF11 size | Peak (watched) | Peak footprint " # noqa: RUF001
|
|
161
|
+
"| Peak MLX (active + cache) | Label | Status | Limits | Result |"
|
|
162
|
+
)
|
|
163
|
+
|
|
164
|
+
|
|
165
|
+
def render_tier_table(rows: Sequence[TierRow]) -> str:
|
|
166
|
+
"""Render the tier table: GiB with two decimals, one row per measured tier."""
|
|
167
|
+
lines = [_TIER_HEADER, "|" + "---|" * 11]
|
|
168
|
+
lines.extend(
|
|
169
|
+
f"| {r.mac_gb} GB | {_gib(r.ceiling_bytes)} | {_MODEL_NAMES.get(r.model, r.model)} "
|
|
170
|
+
f"| {_gib(r.df11_bytes)} | {_gib(r.watched_peak_bytes)} | {_gib(r.footprint_peak_bytes)} "
|
|
171
|
+
f"| {_gib(r.mlx_peak_bytes)} | {r.label} | {r.status} | {r.limits_note} | `{r.source}` |"
|
|
172
|
+
for r in rows
|
|
173
|
+
)
|
|
174
|
+
return "\n".join(lines) + "\n"
|
|
175
|
+
|
|
176
|
+
|
|
177
|
+
def scenario_title(key: str) -> str:
|
|
178
|
+
"""``flux1-dev-1024`` as ``FLUX.1-dev, 1024²``; any other key unchanged."""
|
|
179
|
+
m = re.fullmatch(r"flux1-(.+)-(\d+)", key)
|
|
180
|
+
if m is None:
|
|
181
|
+
return key
|
|
182
|
+
return f"{_MODEL_NAMES.get(m.group(1), m.group(1))}, {m.group(2)}²"
|
|
183
|
+
|
|
184
|
+
|
|
185
|
+
def render_overhead_block(
|
|
186
|
+
summaries: Mapping[str, Summary],
|
|
187
|
+
*,
|
|
188
|
+
caption: str,
|
|
189
|
+
reproducers: Mapping[str, str],
|
|
190
|
+
cache_limit_note: str,
|
|
191
|
+
preflight_skipped: Mapping[str, Sequence[str]] | None = None,
|
|
192
|
+
) -> str:
|
|
193
|
+
"""Render each scenario's overhead line and recorded command, then the caption and cache note.
|
|
194
|
+
|
|
195
|
+
``preflight_skipped`` maps a scenario whose run skipped the launch check to the gates that
|
|
196
|
+
failed; its command line says so.
|
|
197
|
+
"""
|
|
198
|
+
skipped = preflight_skipped or {}
|
|
199
|
+
lines = []
|
|
200
|
+
for key, s in summaries.items():
|
|
201
|
+
per_block = _pct(s.overhead.get("per-block"))
|
|
202
|
+
depth2 = _pct(s.overhead.get("depth2"))
|
|
203
|
+
cost = _NOT_MEASURED if s.eval_cost_s is None else f"{s.eval_cost_s:.2f} s/step"
|
|
204
|
+
q8 = _NOT_MEASURED if s.q8_ratio is None else f"{s.q8_ratio:.2f}×" # noqa: RUF001
|
|
205
|
+
lines.append(
|
|
206
|
+
f"{scenario_title(key)}, per-block evaluation: {per_block} (depth-2: {depth2}); "
|
|
207
|
+
f"eval policy cost {cost}; DF11 (per-block) over the mflux q8 step "
|
|
208
|
+
f"(one eval per step, same cache limit): {q8}"
|
|
209
|
+
)
|
|
210
|
+
lines.append("")
|
|
211
|
+
command = f"Command: `{reproducers.get(key, '')}`"
|
|
212
|
+
if key in skipped:
|
|
213
|
+
command += f" (preflight skipped: {', '.join(skipped[key]) or 'no gate failed'})"
|
|
214
|
+
lines += [command, ""]
|
|
215
|
+
lines += [caption, "", cache_limit_note]
|
|
216
|
+
return "\n".join(lines) + "\n"
|
|
217
|
+
|
|
218
|
+
|
|
219
|
+
def render_proof_paragraph(proof: ProofRecord) -> str:
|
|
220
|
+
"""Render the harness-proof paragraph."""
|
|
221
|
+
return (
|
|
222
|
+
f"Harness proof: under one {_gib(proof.cap_bytes)} cap, a {proof.pass_size}² run passed "
|
|
223
|
+
f"with a watched peak of {_gib(proof.pass_watched_peak)}, and a {proof.abort_size}² run was "
|
|
224
|
+
f"stopped by the watchdog ({proof.abort_reason}, counter {proof.abort_counter}) at "
|
|
225
|
+
f"{_gib(proof.abort_peak)}. Records: `{proof.source_dir}`.\n"
|
|
226
|
+
)
|
|
227
|
+
|
|
228
|
+
|
|
229
|
+
def splice(text: str, block: str, fragment: str) -> str:
|
|
230
|
+
"""Replace what sits between the ``bench:<block>`` markers with ``fragment``.
|
|
231
|
+
|
|
232
|
+
Raises:
|
|
233
|
+
DFloatFormatError: Not exactly one start marker and one end marker, start first.
|
|
234
|
+
"""
|
|
235
|
+
start, end = f"<!-- bench:{block} -->", f"<!-- /bench:{block} -->"
|
|
236
|
+
if text.count(start) != 1 or text.count(end) != 1 or text.index(start) > text.index(end):
|
|
237
|
+
raise DFloatFormatError(f"expected exactly one {start} before one {end}")
|
|
238
|
+
head, rest = text.split(start, 1)
|
|
239
|
+
_, tail = rest.split(end, 1)
|
|
240
|
+
return f"{head}{start}\n{fragment}{end}{tail}"
|
|
241
|
+
|
|
242
|
+
|
|
243
|
+
def caption(provenance: Mapping[str, Any], *, date: str) -> str:
|
|
244
|
+
"""One-line hardware and version caption; a ``-dirty`` git suffix is kept."""
|
|
245
|
+
info = provenance["device_info"]
|
|
246
|
+
git = str(provenance["git"])
|
|
247
|
+
git = git[:7] + ("-dirty" if git.endswith("-dirty") else "")
|
|
248
|
+
return (
|
|
249
|
+
f"{info['device_name']}, {round(info['memory_size'] / GIB)} GB, macOS {provenance['macos']}, "
|
|
250
|
+
f"mlx {provenance['mlx']}, mflux {provenance['mflux']}, git {git}, {date}"
|
|
251
|
+
)
|
mlx_dfloat/cli.py
ADDED
|
@@ -0,0 +1,30 @@
|
|
|
1
|
+
"""The ``mlx-dfloat`` command: ``generate`` today; ``selftest`` and ``bench`` later."""
|
|
2
|
+
|
|
3
|
+
import argparse
|
|
4
|
+
from collections.abc import Sequence
|
|
5
|
+
|
|
6
|
+
from mlx_dfloat._version import __version__
|
|
7
|
+
from mlx_dfloat.mflux.flux1.cli import add_generate_parser
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def build_parser() -> argparse.ArgumentParser:
|
|
11
|
+
"""The top-level parser with its subcommands (no mflux import: each subcommand imports what it runs)."""
|
|
12
|
+
parser = argparse.ArgumentParser(
|
|
13
|
+
prog="mlx-dfloat",
|
|
14
|
+
description="Run DFloat11 losslessly compressed BF16 models on Apple Silicon.",
|
|
15
|
+
)
|
|
16
|
+
parser.add_argument("--version", action="version", version=f"mlx-dfloat {__version__}")
|
|
17
|
+
sub = parser.add_subparsers(dest="command", required=True)
|
|
18
|
+
add_generate_parser(sub)
|
|
19
|
+
return parser
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def main(argv: Sequence[str] | None = None) -> int:
|
|
23
|
+
"""Parse and run the chosen subcommand; the return value is the process exit code."""
|
|
24
|
+
args = build_parser().parse_args(argv)
|
|
25
|
+
run = args.run
|
|
26
|
+
return int(run(args))
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
if __name__ == "__main__": # pragma: no cover
|
|
30
|
+
raise SystemExit(main())
|
mlx_dfloat/decode.py
ADDED
|
@@ -0,0 +1,120 @@
|
|
|
1
|
+
"""Decode DF11 groups to BF16 bit patterns through an explicitly chosen backend."""
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass
|
|
4
|
+
from typing import Literal
|
|
5
|
+
|
|
6
|
+
import mlx.core as mx
|
|
7
|
+
import numpy as np
|
|
8
|
+
|
|
9
|
+
from mlx_dfloat import reference
|
|
10
|
+
from mlx_dfloat.errors import DFloatBackendError, DFloatFormatError
|
|
11
|
+
from mlx_dfloat.format import GroupArrays, MxGroup
|
|
12
|
+
|
|
13
|
+
Backend = Literal["reference", "metal"]
|
|
14
|
+
|
|
15
|
+
STATUS_INVALID_CODE = 1
|
|
16
|
+
STATUS_COUNT_MISMATCH = 2
|
|
17
|
+
STATUS_BROKEN_CHAIN = 4
|
|
18
|
+
STATUS_PATH_DIRECT = 8 # informational: the block wrote straight to device memory
|
|
19
|
+
STATUS_ERROR_MASK = STATUS_INVALID_CODE | STATUS_COUNT_MISMATCH | STATUS_BROKEN_CHAIN
|
|
20
|
+
_STATUS_TEXT = {
|
|
21
|
+
STATUS_INVALID_CODE: "invalid code",
|
|
22
|
+
STATUS_COUNT_MISMATCH: "code count does not match output_positions",
|
|
23
|
+
STATUS_BROKEN_CHAIN: "thread chain broken (corrupt gaps or stream)",
|
|
24
|
+
}
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
@dataclass(frozen=True, slots=True, kw_only=True)
|
|
28
|
+
class DecodeResult:
|
|
29
|
+
"""One decoded group: its BF16 bits, a per-block status word and how it was decoded.
|
|
30
|
+
|
|
31
|
+
``threadgroup_bytes`` is the static threadgroup memory the kernel instantiation reserves (0 for
|
|
32
|
+
the reference), not a per-block figure: every launched block reserves it, staged or direct.
|
|
33
|
+
The bits of a block whose status word has an error bit are undefined; call ``check`` first.
|
|
34
|
+
"""
|
|
35
|
+
|
|
36
|
+
bits: mx.array
|
|
37
|
+
status: mx.array
|
|
38
|
+
backend: Backend
|
|
39
|
+
direct_blocks: int
|
|
40
|
+
threadgroup_bytes: int
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def _to_arrays(group: MxGroup) -> GroupArrays:
|
|
44
|
+
return GroupArrays(
|
|
45
|
+
encoded_exponent=np.array(group.encoded_exponent),
|
|
46
|
+
sign_mantissa=np.array(group.sign_mantissa),
|
|
47
|
+
luts=np.array(group.luts),
|
|
48
|
+
gaps=np.array(group.gaps),
|
|
49
|
+
output_positions=np.array(group.positions).astype(np.uint32),
|
|
50
|
+
split_positions=np.array(group.split_positions, dtype=np.int64),
|
|
51
|
+
)
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def available_backends() -> tuple[Backend, ...]:
|
|
55
|
+
"""Backends that can run here; "metal" appears only after its kernel warm-up succeeds.
|
|
56
|
+
|
|
57
|
+
The warm-up compiles and checks a pipeline for every input-binding signature a group can
|
|
58
|
+
present (many blocks, one block, all-tiny), so once "metal" is listed every group's decode
|
|
59
|
+
reuses a pipeline that already decoded bit-exactly here.
|
|
60
|
+
"""
|
|
61
|
+
from mlx_dfloat import _metal_decode # lazy: importing the package must not touch the GPU
|
|
62
|
+
|
|
63
|
+
return ("reference", "metal") if _metal_decode.metal_ready() else ("reference",)
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def decode_group(group: MxGroup, *, backend: Backend) -> DecodeResult:
|
|
67
|
+
"""Decode a group. The reference runs eagerly on the CPU; "metal" returns a lazy result.
|
|
68
|
+
|
|
69
|
+
The reference backend skips only the stream-level EOF byte-count check (`check_stream_end=False`): the kernel
|
|
70
|
+
cannot see it either, and the committed test slice is deliberately truncated. The reference module's own tests
|
|
71
|
+
keep the strict path.
|
|
72
|
+
|
|
73
|
+
Raises:
|
|
74
|
+
DFloatBackendError: `backend` is not a known backend, or `backend="metal"` is requested and
|
|
75
|
+
the Metal backend cannot run here.
|
|
76
|
+
DFloatFormatError: The reference decoder finds the group structurally invalid.
|
|
77
|
+
"""
|
|
78
|
+
if backend == "reference":
|
|
79
|
+
out = reference.decode_group(_to_arrays(group), name=group.name, check_stream_end=False)
|
|
80
|
+
return DecodeResult(
|
|
81
|
+
bits=mx.array(out),
|
|
82
|
+
status=mx.zeros((group.n_launch,), dtype=mx.uint32),
|
|
83
|
+
backend="reference",
|
|
84
|
+
direct_blocks=0,
|
|
85
|
+
threadgroup_bytes=0,
|
|
86
|
+
)
|
|
87
|
+
if backend == "metal":
|
|
88
|
+
from mlx_dfloat import _metal_decode
|
|
89
|
+
|
|
90
|
+
return _metal_decode.decode(group)
|
|
91
|
+
raise DFloatBackendError(f"unknown backend {backend!r}")
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
def check(result: DecodeResult, *, name: str = "<group>") -> None:
|
|
95
|
+
"""Evaluate the status words and refuse the group if any block reports an error.
|
|
96
|
+
|
|
97
|
+
Raises:
|
|
98
|
+
DFloatFormatError: A block's status word has an error bit set (informational bits, such
|
|
99
|
+
as `STATUS_PATH_DIRECT`, are ignored).
|
|
100
|
+
"""
|
|
101
|
+
check_status(result.status, name=name)
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
def check_status(status_words: mx.array, *, name: str = "<group>") -> None:
|
|
105
|
+
"""``check`` on a bare status array: reads it on the host, so call it after the step's eval.
|
|
106
|
+
|
|
107
|
+
Raises:
|
|
108
|
+
DFloatFormatError: A block's status word has an error bit set.
|
|
109
|
+
"""
|
|
110
|
+
status = np.array(status_words) & STATUS_ERROR_MASK
|
|
111
|
+
bad = np.flatnonzero(status)
|
|
112
|
+
if bad.size:
|
|
113
|
+
b = int(bad[0])
|
|
114
|
+
reasons = ", ".join(t for flag, t in _STATUS_TEXT.items() if int(status[b]) & flag)
|
|
115
|
+
raise DFloatFormatError(f"{name}: block {b}: {reasons}")
|
|
116
|
+
|
|
117
|
+
|
|
118
|
+
def split_matrices(flat: mx.array, split_positions: tuple[int, ...]) -> list[mx.array]:
|
|
119
|
+
"""Cut a decoded group into its matrices (views, no copies)."""
|
|
120
|
+
return list(mx.split(flat, list(split_positions))) if split_positions else [flat]
|
mlx_dfloat/errors.py
ADDED
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
"""Package-rooted exceptions for mlx-dfloat."""
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
class DFloatError(Exception):
|
|
5
|
+
"""Base class for every error raised by mlx-dfloat."""
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class DFloatFormatError(DFloatError, ValueError):
|
|
9
|
+
"""A DFloat11 checkpoint is malformed or uses a layout this version cannot read."""
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class DFloatResourceError(DFloatError, MemoryError):
|
|
13
|
+
"""A decode would need more memory than the configured budget allows."""
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class DFloatBackendError(DFloatError):
|
|
17
|
+
"""A decode backend is unavailable here or its kernel cannot run."""
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class DFloatIntegrationError(DFloatError, RuntimeError):
|
|
21
|
+
"""A block-seam or name-map invariant failed: a name outside the map, a missing or extra layer, a shape mismatch."""
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
class DFloatUnsupportedError(DFloatError, ValueError):
|
|
25
|
+
"""An option this integration does not support on the DFloat11 path (quantisation, LoRA, img2img, ...)."""
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class DFloatDependencyError(DFloatError, ImportError):
|
|
29
|
+
"""An optional dependency this feature needs is not installed."""
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
class DFloatAccessError(DFloatError, PermissionError):
|
|
33
|
+
"""A model repository on the Hub that this account may not read (a gated licence not accepted, or no token)."""
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
__all__ = [
|
|
37
|
+
"DFloatAccessError",
|
|
38
|
+
"DFloatBackendError",
|
|
39
|
+
"DFloatDependencyError",
|
|
40
|
+
"DFloatError",
|
|
41
|
+
"DFloatFormatError",
|
|
42
|
+
"DFloatIntegrationError",
|
|
43
|
+
"DFloatResourceError",
|
|
44
|
+
"DFloatUnsupportedError",
|
|
45
|
+
]
|