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,216 @@
|
|
|
1
|
+
"""The launch gate every bench run passes first.
|
|
2
|
+
|
|
3
|
+
Every probe is one macOS command whose text a pure parser reads; ``check`` is a pure function of
|
|
4
|
+
the parsed sample. An unreadable probe parses to None, and ``check`` reports it as a failed gate:
|
|
5
|
+
the gate refuses rather than guesses. ``pmset -g therm`` prints no ``CPU_Speed_Limit`` line when
|
|
6
|
+
macOS has recorded no limit; that reads as None and passes the speed-limit check.
|
|
7
|
+
|
|
8
|
+
The busy gate lists heavy processes by pid and executable name only, never their arguments.
|
|
9
|
+
Command lines containing a substring from ``MLX_DFLOAT_PREFLIGHT_EXCLUDE`` (comma-separated, empty
|
|
10
|
+
by default) are not counted as busy.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
import dataclasses
|
|
14
|
+
import os
|
|
15
|
+
import re
|
|
16
|
+
import shutil
|
|
17
|
+
import subprocess
|
|
18
|
+
from collections.abc import Callable, Mapping, Sequence
|
|
19
|
+
from pathlib import Path
|
|
20
|
+
|
|
21
|
+
from mlx_dfloat.bench.capped import GIB
|
|
22
|
+
|
|
23
|
+
_PERCENT = re.compile(r"\t(\d+)%;")
|
|
24
|
+
_WATTS = re.compile(r"Wattage\s*=\s*(\d+)W")
|
|
25
|
+
_SPEED = re.compile(r"CPU_Speed_Limit\s*=\s*(\d+)")
|
|
26
|
+
_CLAMSHELL = re.compile(r'"AppleClamshellState"\s*=\s*(Yes|No)')
|
|
27
|
+
_FREE = re.compile(r"free percentage:\s*(\d+)%")
|
|
28
|
+
_PS_LINE = re.compile(r"^\s*(\d+)\s+(\d+)\s+(.*)$")
|
|
29
|
+
HEAVY_PATTERNS: tuple[str, ...] = (
|
|
30
|
+
"bench_",
|
|
31
|
+
"sweep",
|
|
32
|
+
"calibrat",
|
|
33
|
+
"pytest",
|
|
34
|
+
"mflux",
|
|
35
|
+
"mlx_lm",
|
|
36
|
+
"mlx-guard",
|
|
37
|
+
"generate",
|
|
38
|
+
"verify_checkpoint",
|
|
39
|
+
"verify_remote_group",
|
|
40
|
+
"bench_decode_kernel",
|
|
41
|
+
"verify_image",
|
|
42
|
+
"mlx-dfloat",
|
|
43
|
+
)
|
|
44
|
+
EXCLUDE_ENV = "MLX_DFLOAT_PREFLIGHT_EXCLUDE"
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def preflight_exclude(environ: Mapping[str, str] = os.environ) -> tuple[str, ...]:
|
|
48
|
+
"""The substrings ``MLX_DFLOAT_PREFLIGHT_EXCLUDE`` names (comma-separated); empty when unset."""
|
|
49
|
+
return tuple(part.strip() for part in environ.get(EXCLUDE_ENV, "").split(",") if part.strip())
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
@dataclasses.dataclass(frozen=True, slots=True, kw_only=True)
|
|
53
|
+
class Preflight:
|
|
54
|
+
"""One parsed sample of the machine state before a launch."""
|
|
55
|
+
|
|
56
|
+
ac_power: bool | None
|
|
57
|
+
battery_percent: int | None
|
|
58
|
+
charging: str | None
|
|
59
|
+
charger_watts: int | None
|
|
60
|
+
cpu_speed_limit: int | None
|
|
61
|
+
lid_open: bool | None
|
|
62
|
+
free_disk_bytes: int | None
|
|
63
|
+
memory_free_percent: int | None
|
|
64
|
+
busy_processes: tuple[str, ...] | None
|
|
65
|
+
|
|
66
|
+
def as_dict(self) -> dict[str, object]:
|
|
67
|
+
"""Every field, for the run record."""
|
|
68
|
+
return dataclasses.asdict(self)
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def parse_batt(text: str) -> tuple[bool | None, int | None, str | None]:
|
|
72
|
+
"""``pmset -g batt``: (on AC or None when unreadable, battery %, charging "yes" / "no" / "full" / None)."""
|
|
73
|
+
if "drawing from" not in text:
|
|
74
|
+
return None, None, None
|
|
75
|
+
ac = "'AC Power'" in text
|
|
76
|
+
match = _PERCENT.search(text)
|
|
77
|
+
pct = int(match.group(1)) if match else None
|
|
78
|
+
if pct is None:
|
|
79
|
+
state = None
|
|
80
|
+
elif "not charging" in text or "discharging" in text:
|
|
81
|
+
state = "no"
|
|
82
|
+
elif "charged" in text:
|
|
83
|
+
state = "full"
|
|
84
|
+
elif "charging" in text:
|
|
85
|
+
state = "yes"
|
|
86
|
+
else:
|
|
87
|
+
state = None
|
|
88
|
+
return ac, pct, state
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
def parse_ac(text: str) -> int | None:
|
|
92
|
+
"""``pmset -g ac``: the adapter wattage, or None (no Wattage line, e.g. "No adapter attached.")."""
|
|
93
|
+
match = _WATTS.search(text)
|
|
94
|
+
return int(match.group(1)) if match else None
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
def parse_therm(text: str) -> int | None:
|
|
98
|
+
"""``pmset -g therm``: the CPU speed limit, or None when none is recorded (or the probe failed)."""
|
|
99
|
+
match = _SPEED.search(text)
|
|
100
|
+
return int(match.group(1)) if match else None
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
def parse_clamshell(text: str) -> bool | None:
|
|
104
|
+
"""``ioreg`` AppleClamshellState: True when the lid is open, False closed, None unreadable."""
|
|
105
|
+
match = _CLAMSHELL.search(text)
|
|
106
|
+
return None if match is None else match.group(1) == "No"
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
def parse_memory_pressure(text: str) -> int | None:
|
|
110
|
+
"""``memory_pressure``: the system-wide free percentage, or None."""
|
|
111
|
+
match = _FREE.search(text)
|
|
112
|
+
return int(match.group(1)) if match else None
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
def parse_ps(
|
|
116
|
+
text: str,
|
|
117
|
+
*,
|
|
118
|
+
patterns: Sequence[str],
|
|
119
|
+
min_rss_bytes: int = GIB,
|
|
120
|
+
exclude: Sequence[str] = (),
|
|
121
|
+
) -> tuple[str, ...]:
|
|
122
|
+
"""``ps -Ao pid=,rss=,command=`` (rss in KiB): the matching processes at or above ``min_rss_bytes``.
|
|
123
|
+
|
|
124
|
+
Each is listed as ``"<executable name> (pid <pid>)"``; its arguments are never recorded. A
|
|
125
|
+
command line containing any ``exclude`` substring is skipped.
|
|
126
|
+
"""
|
|
127
|
+
wanted = re.compile("|".join(re.escape(p) for p in patterns), re.IGNORECASE)
|
|
128
|
+
out: list[str] = []
|
|
129
|
+
for line in text.splitlines():
|
|
130
|
+
m = _PS_LINE.match(line)
|
|
131
|
+
if not m:
|
|
132
|
+
continue
|
|
133
|
+
rss_bytes, command = int(m.group(2)) * 1024, m.group(3).strip()
|
|
134
|
+
if any(x in command for x in exclude) or not wanted.search(command):
|
|
135
|
+
continue
|
|
136
|
+
if rss_bytes >= min_rss_bytes:
|
|
137
|
+
out.append(f"{Path(command.split()[0]).name} (pid {m.group(1)})")
|
|
138
|
+
return tuple(out)
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
def check(
|
|
142
|
+
p: Preflight,
|
|
143
|
+
*,
|
|
144
|
+
min_battery: int = 40,
|
|
145
|
+
min_not_charging: int = 50,
|
|
146
|
+
min_free_disk_bytes: int = 20 * GIB,
|
|
147
|
+
min_memory_free_percent: int = 20,
|
|
148
|
+
) -> list[str]:
|
|
149
|
+
"""The failed gates, by name; empty means go. A None that a gate needs is ``unreadable:<field>``.
|
|
150
|
+
|
|
151
|
+
``not_charging`` (on AC, battery not charging) fires only below ``min_not_charging`` percent:
|
|
152
|
+
macOS Optimized Battery Charging holds a MacBook on AC at 80 % without charging, and a battery
|
|
153
|
+
at half charge or more on AC may run.
|
|
154
|
+
"""
|
|
155
|
+
failed: list[str] = []
|
|
156
|
+
pct = p.battery_percent
|
|
157
|
+
if p.ac_power is None:
|
|
158
|
+
failed.append("unreadable:ac_power")
|
|
159
|
+
elif not p.ac_power:
|
|
160
|
+
failed.append("ac_power")
|
|
161
|
+
if pct is not None and pct < min_battery:
|
|
162
|
+
failed.append("battery")
|
|
163
|
+
if pct is not None and p.ac_power and p.charging == "no" and pct < min_not_charging:
|
|
164
|
+
failed.append("not_charging")
|
|
165
|
+
if p.cpu_speed_limit is not None and p.cpu_speed_limit < 100:
|
|
166
|
+
failed.append("cpu_speed_limit")
|
|
167
|
+
if pct is not None:
|
|
168
|
+
if p.lid_open is None:
|
|
169
|
+
failed.append("unreadable:lid_open")
|
|
170
|
+
elif not p.lid_open:
|
|
171
|
+
failed.append("lid")
|
|
172
|
+
if p.free_disk_bytes is None:
|
|
173
|
+
failed.append("unreadable:free_disk_bytes")
|
|
174
|
+
elif p.free_disk_bytes < min_free_disk_bytes:
|
|
175
|
+
failed.append("free_disk")
|
|
176
|
+
if p.memory_free_percent is None:
|
|
177
|
+
failed.append("unreadable:memory_free_percent")
|
|
178
|
+
elif p.memory_free_percent < min_memory_free_percent:
|
|
179
|
+
failed.append("memory_free")
|
|
180
|
+
if p.busy_processes is None:
|
|
181
|
+
failed.append("unreadable:busy_processes")
|
|
182
|
+
elif p.busy_processes:
|
|
183
|
+
failed.append("busy")
|
|
184
|
+
return failed
|
|
185
|
+
|
|
186
|
+
|
|
187
|
+
def _run(*argv: str) -> str:
|
|
188
|
+
try:
|
|
189
|
+
return subprocess.run(list(argv), capture_output=True, text=True, check=False).stdout
|
|
190
|
+
except OSError:
|
|
191
|
+
return ""
|
|
192
|
+
|
|
193
|
+
|
|
194
|
+
def sample(disk_path: Path = Path("/"), *, run: Callable[..., str] = _run) -> Preflight:
|
|
195
|
+
"""Probe the machine once. Never raises; an unreadable probe is None."""
|
|
196
|
+
ac, pct, state = parse_batt(run("pmset", "-g", "batt"))
|
|
197
|
+
try:
|
|
198
|
+
free: int | None = shutil.disk_usage(disk_path).free
|
|
199
|
+
except OSError:
|
|
200
|
+
free = None
|
|
201
|
+
ps_text = run("ps", "-Ao", "pid=,rss=,command=")
|
|
202
|
+
return Preflight(
|
|
203
|
+
ac_power=ac,
|
|
204
|
+
battery_percent=pct,
|
|
205
|
+
charging=state,
|
|
206
|
+
charger_watts=parse_ac(run("pmset", "-g", "ac")),
|
|
207
|
+
cpu_speed_limit=parse_therm(run("pmset", "-g", "therm")),
|
|
208
|
+
lid_open=parse_clamshell(run("ioreg", "-r", "-k", "AppleClamshellState", "-d", "4")),
|
|
209
|
+
free_disk_bytes=free,
|
|
210
|
+
memory_free_percent=parse_memory_pressure(run("memory_pressure")),
|
|
211
|
+
busy_processes=(
|
|
212
|
+
parse_ps(ps_text, patterns=HEAVY_PATTERNS, exclude=preflight_exclude())
|
|
213
|
+
if ps_text
|
|
214
|
+
else None
|
|
215
|
+
),
|
|
216
|
+
)
|
|
@@ -0,0 +1,227 @@
|
|
|
1
|
+
"""Read the step bench's per-child result files and pool them into a summary.
|
|
2
|
+
|
|
3
|
+
Each child writes ``round{r}-{mode}.json``. Every pair pools its two conditions over the rounds in
|
|
4
|
+
which both completed; the overhead comes from the pooled medians (``t_df11 / t_control - 1``).
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
import json
|
|
8
|
+
import statistics
|
|
9
|
+
from collections.abc import Mapping, Sequence
|
|
10
|
+
from dataclasses import dataclass
|
|
11
|
+
from pathlib import Path
|
|
12
|
+
from typing import Any
|
|
13
|
+
|
|
14
|
+
from mlx_dfloat.bench.scenario import CONDITIONS
|
|
15
|
+
from mlx_dfloat.errors import DFloatFormatError
|
|
16
|
+
|
|
17
|
+
PAIRS: tuple[tuple[str, str, str], ...] = (
|
|
18
|
+
("per-block", "df11", "control"),
|
|
19
|
+
("depth2", "df11-depth2", "control-depth2"),
|
|
20
|
+
)
|
|
21
|
+
EXTRA_PAIRS: tuple[tuple[str, str, str], ...] = (
|
|
22
|
+
("eval-policy", "control", "control-noeval"),
|
|
23
|
+
("q8", "df11", "q8"),
|
|
24
|
+
)
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
@dataclass(frozen=True, slots=True, kw_only=True)
|
|
28
|
+
class ConditionResult:
|
|
29
|
+
"""One completed child: a condition's timed steps in one round, with its peaks."""
|
|
30
|
+
|
|
31
|
+
condition: str
|
|
32
|
+
round: int
|
|
33
|
+
step_s: tuple[float, ...]
|
|
34
|
+
launches_per_step: tuple[int, ...]
|
|
35
|
+
launches_expected: int
|
|
36
|
+
step_footprint_peak: int
|
|
37
|
+
step_mlx_peak: int
|
|
38
|
+
step_watched_peak: int
|
|
39
|
+
footprint_peak: int
|
|
40
|
+
mlx_peak: int
|
|
41
|
+
watched_peak: int
|
|
42
|
+
scenario_hash: str
|
|
43
|
+
label: str
|
|
44
|
+
limits: dict[str, object]
|
|
45
|
+
|
|
46
|
+
@property
|
|
47
|
+
def median(self) -> float:
|
|
48
|
+
"""Median of the timed steps, in seconds."""
|
|
49
|
+
return statistics.median(self.step_s)
|
|
50
|
+
|
|
51
|
+
@property
|
|
52
|
+
def spread(self) -> float:
|
|
53
|
+
"""(max - min) / median over the timed steps."""
|
|
54
|
+
return (max(self.step_s) - min(self.step_s)) / self.median
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def _field(data: Mapping[str, Any], name: str) -> Any:
|
|
58
|
+
if name not in data:
|
|
59
|
+
raise DFloatFormatError(f"child result: missing field {name!r}")
|
|
60
|
+
return data[name]
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
def result_from_json(data: Mapping[str, Any]) -> ConditionResult:
|
|
64
|
+
"""Build a result from a child JSON.
|
|
65
|
+
|
|
66
|
+
Raises:
|
|
67
|
+
DFloatFormatError: A missing field, a nonzero ``exit_code`` or an empty ``step_s``.
|
|
68
|
+
"""
|
|
69
|
+
mode = str(_field(data, "mode"))
|
|
70
|
+
exit_code = _field(data, "exit_code")
|
|
71
|
+
if exit_code != 0:
|
|
72
|
+
raise DFloatFormatError(f"{mode}: exit_code {exit_code}: the child did not complete")
|
|
73
|
+
step_s = tuple(float(s) for s in _field(data, "step_s"))
|
|
74
|
+
if not step_s:
|
|
75
|
+
raise DFloatFormatError(f"{mode}: step_s is empty: the child has no timed steps")
|
|
76
|
+
warmup = int(_field(data, "warmup"))
|
|
77
|
+
return ConditionResult(
|
|
78
|
+
condition=mode,
|
|
79
|
+
round=int(_field(data, "round")),
|
|
80
|
+
step_s=step_s,
|
|
81
|
+
launches_per_step=tuple(int(n) for n in _field(data, "launches_per_step")[warmup:]),
|
|
82
|
+
launches_expected=int(_field(data, "launches_expected_per_step")),
|
|
83
|
+
step_footprint_peak=int(_field(data, "step_footprint_peak_bytes")),
|
|
84
|
+
step_mlx_peak=int(_field(data, "step_mlx_peak_bytes")),
|
|
85
|
+
step_watched_peak=int(_field(data, "step_watched_peak_bytes")),
|
|
86
|
+
footprint_peak=int(_field(data, "footprint_peak_bytes")),
|
|
87
|
+
mlx_peak=int(_field(data, "mlx_peak_memory_bytes")),
|
|
88
|
+
watched_peak=int(_field(data, "watched_peak_bytes")),
|
|
89
|
+
scenario_hash=str(_field(data, "scenario_hash")),
|
|
90
|
+
label=str(_field(data, "label")),
|
|
91
|
+
limits=dict(_field(data, "limits")),
|
|
92
|
+
)
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
def _one_scenario(results: Sequence[ConditionResult]) -> str:
|
|
96
|
+
hashes = sorted({r.scenario_hash for r in results})
|
|
97
|
+
if len(hashes) > 1:
|
|
98
|
+
raise DFloatFormatError(
|
|
99
|
+
f"results carry {len(hashes)} different scenario hashes ({', '.join(h[:12] for h in hashes)}): "
|
|
100
|
+
"a stale child from another recipe is in this set"
|
|
101
|
+
)
|
|
102
|
+
return hashes[0] if hashes else ""
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
def _order(r: ConditionResult) -> tuple[int, int]:
|
|
106
|
+
idx = CONDITIONS.index(r.condition) if r.condition in CONDITIONS else len(CONDITIONS)
|
|
107
|
+
return (r.round, idx)
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
def load_results(directory: Path | str) -> list[ConditionResult]:
|
|
111
|
+
"""Read the complete ``round*-*.json`` children of ``directory``, sorted by (round, condition).
|
|
112
|
+
|
|
113
|
+
Failed or partial children are skipped.
|
|
114
|
+
|
|
115
|
+
Raises:
|
|
116
|
+
DFloatFormatError: The complete children carry more than one scenario hash.
|
|
117
|
+
"""
|
|
118
|
+
results: list[ConditionResult] = []
|
|
119
|
+
for path in sorted(Path(directory).glob("round*-*.json")):
|
|
120
|
+
try:
|
|
121
|
+
data = json.loads(path.read_text())
|
|
122
|
+
except (OSError, ValueError) as exc:
|
|
123
|
+
raise DFloatFormatError(f"{path}: cannot read the child result: {exc}") from exc
|
|
124
|
+
if not isinstance(data, dict) or data.get("exit_code") != 0 or not data.get("step_s"):
|
|
125
|
+
continue
|
|
126
|
+
results.append(result_from_json(data))
|
|
127
|
+
_one_scenario(results)
|
|
128
|
+
return sorted(results, key=_order)
|
|
129
|
+
|
|
130
|
+
|
|
131
|
+
@dataclass(frozen=True, slots=True, kw_only=True)
|
|
132
|
+
class Summary:
|
|
133
|
+
"""Pooled results: per-condition stats, pair overheads, the eval-policy cost, the q8 ratio."""
|
|
134
|
+
|
|
135
|
+
scenario_hash: str
|
|
136
|
+
rounds_seen: int
|
|
137
|
+
conditions: dict[str, dict[str, float | int]]
|
|
138
|
+
overhead: dict[str, float]
|
|
139
|
+
eval_cost_s: float | None
|
|
140
|
+
q8_ratio: float | None
|
|
141
|
+
paired_rounds: dict[str, list[int]]
|
|
142
|
+
pair_n: dict[str, int]
|
|
143
|
+
|
|
144
|
+
|
|
145
|
+
def summarise(results: Sequence[ConditionResult]) -> Summary:
|
|
146
|
+
"""Pool each pair over the rounds where both of its conditions completed.
|
|
147
|
+
|
|
148
|
+
Raises:
|
|
149
|
+
DFloatFormatError: The results carry more than one scenario hash.
|
|
150
|
+
"""
|
|
151
|
+
scenario = _one_scenario(results)
|
|
152
|
+
by: dict[str, dict[int, ConditionResult]] = {}
|
|
153
|
+
for r in results:
|
|
154
|
+
by.setdefault(r.condition, {})[r.round] = r
|
|
155
|
+
|
|
156
|
+
paired: dict[str, list[int]] = {}
|
|
157
|
+
pooled_steps: dict[str, dict[str, list[float]]] = {}
|
|
158
|
+
conditions: dict[str, dict[str, float | int]] = {}
|
|
159
|
+
for label, a, b in (*PAIRS, *EXTRA_PAIRS):
|
|
160
|
+
shared = sorted(set(by.get(a, {})) & set(by.get(b, {})))
|
|
161
|
+
paired[label] = shared
|
|
162
|
+
pooled_steps[label] = {}
|
|
163
|
+
for cond in (a, b):
|
|
164
|
+
rows = [by[cond][rnd] for rnd in shared]
|
|
165
|
+
steps = [s for row in rows for s in row.step_s]
|
|
166
|
+
pooled_steps[label][cond] = steps
|
|
167
|
+
if shared and cond not in conditions:
|
|
168
|
+
m = statistics.median(steps)
|
|
169
|
+
conditions[cond] = {
|
|
170
|
+
"median": m,
|
|
171
|
+
"spread": (max(steps) - min(steps)) / m,
|
|
172
|
+
"n": len(steps),
|
|
173
|
+
"step_watched_peak": max(row.step_watched_peak for row in rows),
|
|
174
|
+
"footprint_peak": max(row.footprint_peak for row in rows),
|
|
175
|
+
"mlx_peak": max(row.mlx_peak for row in rows),
|
|
176
|
+
}
|
|
177
|
+
|
|
178
|
+
def med(label: str, cond: str) -> float:
|
|
179
|
+
return statistics.median(pooled_steps[label][cond])
|
|
180
|
+
|
|
181
|
+
overhead = {label: med(label, d) / med(label, c) - 1 for label, d, c in PAIRS if paired[label]}
|
|
182
|
+
eval_cost = (
|
|
183
|
+
med("eval-policy", "control") - med("eval-policy", "control-noeval")
|
|
184
|
+
if paired["eval-policy"]
|
|
185
|
+
else None
|
|
186
|
+
)
|
|
187
|
+
q8_ratio = med("q8", "df11") / med("q8", "q8") if paired["q8"] else None
|
|
188
|
+
pair_n = {
|
|
189
|
+
label: sum(len(v) for v in steps.values())
|
|
190
|
+
for label, steps in pooled_steps.items()
|
|
191
|
+
if paired[label]
|
|
192
|
+
}
|
|
193
|
+
return Summary(
|
|
194
|
+
scenario_hash=scenario,
|
|
195
|
+
rounds_seen=len({r.round for r in results}),
|
|
196
|
+
conditions=conditions,
|
|
197
|
+
overhead=overhead,
|
|
198
|
+
eval_cost_s=eval_cost,
|
|
199
|
+
q8_ratio=q8_ratio,
|
|
200
|
+
paired_rounds={k: v for k, v in paired.items() if v},
|
|
201
|
+
pair_n=pair_n,
|
|
202
|
+
)
|
|
203
|
+
|
|
204
|
+
|
|
205
|
+
def expected_missing(
|
|
206
|
+
results: Sequence[ConditionResult], *, conditions: Sequence[str], rounds: int
|
|
207
|
+
) -> list[str]:
|
|
208
|
+
"""List ``"round r: condition"`` for every absent result, round-major in ``conditions`` order."""
|
|
209
|
+
have = {(r.round, r.condition) for r in results}
|
|
210
|
+
return [
|
|
211
|
+
f"round {rnd}: {c}"
|
|
212
|
+
for rnd in range(1, rounds + 1)
|
|
213
|
+
for c in conditions
|
|
214
|
+
if (rnd, c) not in have
|
|
215
|
+
]
|
|
216
|
+
|
|
217
|
+
|
|
218
|
+
__all__ = [
|
|
219
|
+
"EXTRA_PAIRS",
|
|
220
|
+
"PAIRS",
|
|
221
|
+
"ConditionResult",
|
|
222
|
+
"Summary",
|
|
223
|
+
"expected_missing",
|
|
224
|
+
"load_results",
|
|
225
|
+
"result_from_json",
|
|
226
|
+
"summarise",
|
|
227
|
+
]
|
|
@@ -0,0 +1,191 @@
|
|
|
1
|
+
"""A bench scenario: the pinned recipe a result file is keyed on.
|
|
2
|
+
|
|
3
|
+
The file is TOML. Every field is validated here with the field name in the error, and unknown keys
|
|
4
|
+
are refused, so a mis-typed key can never run the bench on a default while the result claims the
|
|
5
|
+
scenario. ``scenario_hash`` is the sha256 of the canonical JSON of the fields, so two files with the
|
|
6
|
+
same content hash the same whatever their key order or spacing.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
import dataclasses
|
|
10
|
+
import hashlib
|
|
11
|
+
import json
|
|
12
|
+
import re
|
|
13
|
+
import tomllib
|
|
14
|
+
from collections.abc import Mapping
|
|
15
|
+
from pathlib import Path
|
|
16
|
+
|
|
17
|
+
from mlx_dfloat.errors import DFloatFormatError
|
|
18
|
+
|
|
19
|
+
CONDITIONS: tuple[str, ...] = (
|
|
20
|
+
"df11",
|
|
21
|
+
"control",
|
|
22
|
+
"df11-depth2",
|
|
23
|
+
"control-depth2",
|
|
24
|
+
"control-noeval",
|
|
25
|
+
"q8",
|
|
26
|
+
)
|
|
27
|
+
MODELS: tuple[str, ...] = ("schnell", "dev")
|
|
28
|
+
_SHA = re.compile(r"^[0-9a-f]{40}$")
|
|
29
|
+
# The name is a directory under the results root: a plain lowercase name, never a path.
|
|
30
|
+
_NAME = re.compile(r"^[a-z0-9][a-z0-9._-]{0,63}$")
|
|
31
|
+
RESERVED_NAMES: tuple[str, ...] = ("tiers", "harness-proof") # the table's own input directories
|
|
32
|
+
MAX_CACHE_LIMIT = 2**63 - 1 # mx.set_cache_limit takes a signed 64-bit size
|
|
33
|
+
_STR_FIELDS = (
|
|
34
|
+
"name",
|
|
35
|
+
"model",
|
|
36
|
+
"df11_repo",
|
|
37
|
+
"df11_revision",
|
|
38
|
+
"base_repo",
|
|
39
|
+
"base_revision",
|
|
40
|
+
"prompt",
|
|
41
|
+
)
|
|
42
|
+
_INT_FIELDS = ("seed", "steps", "warmup", "size", "rounds", "cache_limit_bytes")
|
|
43
|
+
_REQUIRED = (*_STR_FIELDS, *_INT_FIELDS, "conditions", "wall_budget_s")
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
@dataclasses.dataclass(frozen=True, slots=True, kw_only=True)
|
|
47
|
+
class Scenario:
|
|
48
|
+
"""One pinned bench recipe (see the module docstring)."""
|
|
49
|
+
|
|
50
|
+
name: str
|
|
51
|
+
model: str
|
|
52
|
+
df11_repo: str
|
|
53
|
+
df11_revision: str
|
|
54
|
+
base_repo: str
|
|
55
|
+
base_revision: str
|
|
56
|
+
prompt: str
|
|
57
|
+
seed: int
|
|
58
|
+
steps: int
|
|
59
|
+
warmup: int
|
|
60
|
+
size: int
|
|
61
|
+
rounds: int
|
|
62
|
+
cache_limit_bytes: int
|
|
63
|
+
conditions: tuple[str, ...]
|
|
64
|
+
wall_budget_s: float
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def _fail(source: str, field: str, why: str) -> DFloatFormatError:
|
|
68
|
+
return DFloatFormatError(f"{source}: {field}: {why}")
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def _str(data: Mapping[str, object], field: str, source: str) -> str:
|
|
72
|
+
value = data[field]
|
|
73
|
+
if not isinstance(value, str):
|
|
74
|
+
raise _fail(source, field, f"expected a string, got {type(value).__name__}")
|
|
75
|
+
return value
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
def _int(data: Mapping[str, object], field: str, source: str) -> int:
|
|
79
|
+
value = data[field]
|
|
80
|
+
if isinstance(value, bool) or not isinstance(value, int):
|
|
81
|
+
raise _fail(source, field, f"expected an integer, got {type(value).__name__}")
|
|
82
|
+
return value
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
def _float(data: Mapping[str, object], field: str, source: str) -> float:
|
|
86
|
+
value = data[field]
|
|
87
|
+
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
|
88
|
+
raise _fail(source, field, f"expected a number, got {type(value).__name__}")
|
|
89
|
+
return float(value)
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
def scenario_from_mapping(data: Mapping[str, object], *, source: str = "<mapping>") -> Scenario:
|
|
93
|
+
"""Validate a parsed mapping into a ``Scenario``.
|
|
94
|
+
|
|
95
|
+
Raises:
|
|
96
|
+
DFloatFormatError: An unknown key, a missing field, a wrong type, a value out of range, a
|
|
97
|
+
revision that is not a full SHA, a duplicate, unknown or empty condition list, or a
|
|
98
|
+
name that is not a plain directory name (or is reserved).
|
|
99
|
+
"""
|
|
100
|
+
unknown = sorted(set(data) - set(_REQUIRED))
|
|
101
|
+
if unknown:
|
|
102
|
+
raise _fail(source, unknown[0], "unknown key")
|
|
103
|
+
for field in _REQUIRED:
|
|
104
|
+
if field not in data:
|
|
105
|
+
raise _fail(source, field, "missing")
|
|
106
|
+
strings = {f: _str(data, f, source) for f in _STR_FIELDS}
|
|
107
|
+
ints = {f: _int(data, f, source) for f in _INT_FIELDS}
|
|
108
|
+
name = strings["name"]
|
|
109
|
+
if not _NAME.fullmatch(name) or ".." in name:
|
|
110
|
+
raise _fail(
|
|
111
|
+
source,
|
|
112
|
+
"name",
|
|
113
|
+
f"{name!r} must be 1-64 characters of a-z, 0-9, '.', '_' or '-', start with a letter "
|
|
114
|
+
"or digit, and hold no '..'",
|
|
115
|
+
)
|
|
116
|
+
if name in RESERVED_NAMES:
|
|
117
|
+
raise _fail(source, "name", f"{name!r} is reserved for the results root's own directories")
|
|
118
|
+
if strings["model"] not in MODELS:
|
|
119
|
+
raise _fail(source, "model", f"choose from {MODELS}")
|
|
120
|
+
for field in ("df11_revision", "base_revision"):
|
|
121
|
+
if not _SHA.match(strings[field]):
|
|
122
|
+
raise _fail(source, field, "must be a full 40-hex commit SHA")
|
|
123
|
+
for field in ("steps", "warmup", "rounds", "cache_limit_bytes"):
|
|
124
|
+
if ints[field] < 1:
|
|
125
|
+
raise _fail(source, field, "must be >= 1")
|
|
126
|
+
if ints["cache_limit_bytes"] > MAX_CACHE_LIMIT:
|
|
127
|
+
raise _fail(source, "cache_limit_bytes", "must be below 2**63 (MLX takes a signed size)")
|
|
128
|
+
if ints["size"] < 16 or ints["size"] % 16:
|
|
129
|
+
raise _fail(source, "size", "must be a positive multiple of 16")
|
|
130
|
+
wall = _float(data, "wall_budget_s", source)
|
|
131
|
+
if wall <= 0:
|
|
132
|
+
raise _fail(source, "wall_budget_s", "must be positive")
|
|
133
|
+
raw = data["conditions"]
|
|
134
|
+
if not isinstance(raw, list) or not all(isinstance(c, str) for c in raw):
|
|
135
|
+
raise _fail(source, "conditions", "expected a list of strings")
|
|
136
|
+
conditions = tuple(raw)
|
|
137
|
+
if not conditions:
|
|
138
|
+
raise _fail(source, "conditions", "must not be empty")
|
|
139
|
+
if len(set(conditions)) != len(conditions):
|
|
140
|
+
raise _fail(source, "conditions", "duplicate entry")
|
|
141
|
+
bad = [c for c in conditions if c not in CONDITIONS]
|
|
142
|
+
if bad:
|
|
143
|
+
raise _fail(source, "conditions", f"unknown {bad[0]!r}; choose from {CONDITIONS}")
|
|
144
|
+
return Scenario(
|
|
145
|
+
name=strings["name"],
|
|
146
|
+
model=strings["model"],
|
|
147
|
+
df11_repo=strings["df11_repo"],
|
|
148
|
+
df11_revision=strings["df11_revision"],
|
|
149
|
+
base_repo=strings["base_repo"],
|
|
150
|
+
base_revision=strings["base_revision"],
|
|
151
|
+
prompt=strings["prompt"],
|
|
152
|
+
seed=ints["seed"],
|
|
153
|
+
steps=ints["steps"],
|
|
154
|
+
warmup=ints["warmup"],
|
|
155
|
+
size=ints["size"],
|
|
156
|
+
rounds=ints["rounds"],
|
|
157
|
+
cache_limit_bytes=ints["cache_limit_bytes"],
|
|
158
|
+
conditions=conditions,
|
|
159
|
+
wall_budget_s=wall,
|
|
160
|
+
)
|
|
161
|
+
|
|
162
|
+
|
|
163
|
+
def load_scenario(path: Path) -> Scenario:
|
|
164
|
+
"""Read and validate a scenario TOML file.
|
|
165
|
+
|
|
166
|
+
Raises:
|
|
167
|
+
DFloatFormatError: The file is not valid TOML or a field is invalid (the message names the file).
|
|
168
|
+
"""
|
|
169
|
+
try:
|
|
170
|
+
data = tomllib.loads(path.read_text())
|
|
171
|
+
except (OSError, tomllib.TOMLDecodeError) as exc:
|
|
172
|
+
raise DFloatFormatError(f"{path}: cannot read the scenario: {exc}") from exc
|
|
173
|
+
return scenario_from_mapping(data, source=str(path))
|
|
174
|
+
|
|
175
|
+
|
|
176
|
+
def scenario_hash(scenario: Scenario) -> str:
|
|
177
|
+
"""sha256 hex of the canonical JSON of the scenario's fields."""
|
|
178
|
+
canonical = json.dumps(dataclasses.asdict(scenario), sort_keys=True, separators=(",", ":"))
|
|
179
|
+
return hashlib.sha256(canonical.encode()).hexdigest()
|
|
180
|
+
|
|
181
|
+
|
|
182
|
+
__all__ = [
|
|
183
|
+
"CONDITIONS",
|
|
184
|
+
"MAX_CACHE_LIMIT",
|
|
185
|
+
"MODELS",
|
|
186
|
+
"RESERVED_NAMES",
|
|
187
|
+
"Scenario",
|
|
188
|
+
"load_scenario",
|
|
189
|
+
"scenario_from_mapping",
|
|
190
|
+
"scenario_hash",
|
|
191
|
+
]
|