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