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