mlx-train-perf 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.
@@ -0,0 +1,12 @@
1
+ from mlx_train_perf.core.loss import (
2
+ DenseHead,
3
+ HeadRef,
4
+ QuantizedHead,
5
+ Resolution,
6
+ linear_cross_entropy,
7
+ resolve_impl,
8
+ tied_head,
9
+ )
10
+
11
+ __all__ = ["DenseHead", "HeadRef", "QuantizedHead", "Resolution", "linear_cross_entropy",
12
+ "resolve_impl", "tied_head"]
@@ -0,0 +1,21 @@
1
+ from mlx_train_perf.errors import UnverifiedMlxError
2
+
3
+ VERIFIED_MLX_VERSIONS: frozenset[str] = frozenset({"0.31.2"})
4
+
5
+
6
+ def _installed_mlx_version() -> str:
7
+ import mlx.core as mx # noqa: PLC0415
8
+
9
+ return str(mx.__version__) # type: ignore[attr-defined] # mlx stubs omit __version__
10
+
11
+
12
+ def check_mlx_verified(*, allow_unverified: bool) -> None:
13
+ """Kernel impl gate: JIT codegen is version-sensitive (register cliffs are sharp)."""
14
+ ver = _installed_mlx_version()
15
+ if ver in VERIFIED_MLX_VERSIONS or allow_unverified:
16
+ return
17
+ raise UnverifiedMlxError(
18
+ f"mlx {ver} is not kernel-verified (verified: {sorted(VERIFIED_MLX_VERSIONS)}). "
19
+ f"Pass allow_unverified_mlx=True to proceed at your own risk (run the parity suite "
20
+ f"first: pytest --run-metal), or use impl='chunked'."
21
+ )
File without changes
@@ -0,0 +1,214 @@
1
+ """mlx-lm trainer adapter: `split_model` + `make_loss_fn`.
2
+
3
+ Verified against the installed mlx-lm==0.31.3 (mlx==0.31.2) source, 2026-07-04:
4
+
5
+ - **Trainer loss contract** (site-packages/mlx_lm/tuner/trainer.py:86-99, `default_loss`)::
6
+
7
+ inputs = batch[:, :-1]
8
+ targets = batch[:, 1:]
9
+ logits = model(inputs)
10
+ steps = mx.arange(1, targets.shape[1] + 1)
11
+ mask = mx.logical_and(steps >= lengths[:, 0:1], steps <= lengths[:, 1:])
12
+ ce = nn.losses.cross_entropy(logits, targets) * mask
13
+ ntoks = mask.sum()
14
+ ce = ce.astype(mx.float32).sum() / ntoks
15
+ return ce, ntoks
16
+
17
+ `lengths` is a TWO-COLUMN `(offset, length)` int array, one row per batch example --
18
+ built by `iterate_batches` (trainer.py:170) as `mx.array(list(zip(offsets, lengths)))`.
19
+ This module reproduces that mask and denominator exactly (mirrors the contract; does
20
+ not import `trainer.py` itself, to keep `mlx_lm` a lazy import here).
21
+
22
+ - **Injection point**: `mlx_lm.tuner.trainer.train(..., loss: callable = default_loss)`,
23
+ invoked as `nn.value_and_grad(model, loss)(model, *batch)` (trainer.py:218-240,
24
+ 240). Per the installed `mlx.nn.value_and_grad` source, its wrapped function calls
25
+ `model.update(params)` **before** calling `loss(model, ...)` -- so by the time our
26
+ loss runs, `model` already carries that step's (gradient-updated, on steps after the
27
+ first) parameters. A loss closure that snapshots a trainable head's weight array once
28
+ at construction would keep computing against that first snapshot forever -- the
29
+ weight *attribute* on `model` gets reassigned by `model.update`, but a frozen
30
+ dataclass field holding the old array object does not follow the reassignment (this
31
+ was confirmed empirically against the installed mlx before relying on it). `loss_fn`
32
+ therefore calls `split_model(model)` fresh on every invocation rather than once at
33
+ `make_loss_fn` construction time.
34
+
35
+ - **Model structure** (mlx_lm/models/llama.py, mlx_lm/models/qwen3.py -- identical
36
+ shape in both families): `Model.__call__` computes hidden states via `self.model(x)`
37
+ (the inner `LlamaModel` / `Qwen3Model`) and projects with `self.lm_head`
38
+ (`nn.Linear` or `nn.QuantizedLinear`) unless `args.tie_word_embeddings`, in which case
39
+ it projects via `self.model.embed_tokens.as_linear(hidden)`
40
+ (`nn.Embedding` / `nn.QuantizedEmbedding`). Both quantized module types expose
41
+ `weight, scales, biases, group_size, bits, mode` directly; `biases` is `None` and
42
+ `mode` is not `"affine"` for the non-affine quantization modes (`mxfp4`/`mxfp8`/
43
+ `nvfp4`) that this project's kernel and chunked paths do not implement.
44
+
45
+ Only the Llama and Qwen3 model families are supported (matched by
46
+ `type(model).__module__`); anything else raises `AdapterError` naming the support list.
47
+ """
48
+ from collections.abc import Callable
49
+ from typing import Any, Literal
50
+
51
+ import mlx.core as mx
52
+ from mlx import nn
53
+
54
+ from mlx_train_perf.core.loss import (
55
+ DenseHead,
56
+ HeadRef,
57
+ QuantizedHead,
58
+ linear_cross_entropy,
59
+ resolve_impl,
60
+ tied_head,
61
+ )
62
+ from mlx_train_perf.errors import AdapterError, MissingDependencyError
63
+
64
+ # Keyed by the exact `type(model).__module__` mlx-lm uses for each family (verified
65
+ # against the installed mlx_lm.models.llama / mlx_lm.models.qwen3 above).
66
+ _SUPPORTED_FAMILIES: dict[str, str] = {
67
+ "llama": "mlx_lm.models.llama",
68
+ "qwen3": "mlx_lm.models.qwen3",
69
+ }
70
+
71
+
72
+ def _require_mlx_lm() -> None:
73
+ try:
74
+ import mlx_lm # noqa: F401, PLC0415
75
+ except ImportError as exc:
76
+ raise MissingDependencyError(
77
+ "mlx-lm is required for mlx_train_perf.adapters.mlx_lm; install the optional "
78
+ "'mlx-lm' extra (pip install 'mlx-train-perf[mlx-lm]')"
79
+ ) from exc
80
+
81
+
82
+ def _quantized_head(module: nn.Module) -> QuantizedHead:
83
+ # `nn` resolves to `Any` here (see the pyproject.toml mypy override for `mlx.nn`),
84
+ # so these attribute reads are unchecked -- correctness is pinned by
85
+ # test_split_quantized_head_reads_affine_fields / test_split_rejects_non_affine_*.
86
+ mode = module.mode
87
+ if mode != "affine":
88
+ raise AdapterError(
89
+ f"quantized head/embedding mode {mode!r} is not supported; only "
90
+ "mode='affine' quantized heads are supported (the only mode this "
91
+ "project's kernel and chunked paths implement)"
92
+ )
93
+ biases = module.biases
94
+ if biases is None:
95
+ raise AdapterError(
96
+ "quantized head/embedding has no biases (mode without biases); "
97
+ "mlx_train_perf's QuantizedHead requires biases"
98
+ )
99
+ return QuantizedHead(
100
+ w_q=module.weight,
101
+ scales=module.scales,
102
+ biases=biases,
103
+ group_size=module.group_size,
104
+ bits=module.bits,
105
+ )
106
+
107
+
108
+ def _head_from_module(module: nn.Module) -> HeadRef:
109
+ """`model.lm_head` case: a dedicated (untied) output projection."""
110
+ if isinstance(module, nn.QuantizedLinear):
111
+ return _quantized_head(module)
112
+ if isinstance(module, nn.Linear):
113
+ # nn.Module.trainable_parameters is itself untyped in mlx's source.
114
+ trainable = "weight" in module.trainable_parameters() # type: ignore[no-untyped-call]
115
+ return DenseHead(weight=module.weight, trainable=trainable)
116
+ raise AdapterError(
117
+ f"unsupported head module type {type(module).__name__!r}; expected "
118
+ "nn.Linear or nn.QuantizedLinear"
119
+ )
120
+
121
+
122
+ def _tied_head_from_embedding(embedding: nn.Module) -> HeadRef:
123
+ """`model.model.embed_tokens` case, used as the head when weights are tied."""
124
+ if isinstance(embedding, nn.QuantizedEmbedding):
125
+ return _quantized_head(embedding)
126
+ if isinstance(embedding, nn.Embedding):
127
+ # nn.Module.trainable_parameters is itself untyped in mlx's source.
128
+ trainable = "weight" in embedding.trainable_parameters() # type: ignore[no-untyped-call]
129
+ return tied_head(embedding.weight, trainable=trainable)
130
+ raise AdapterError(
131
+ f"unsupported embedding module type {type(embedding).__name__!r}; expected "
132
+ "nn.Embedding or nn.QuantizedEmbedding"
133
+ )
134
+
135
+
136
+ def split_model(model: Any) -> tuple[Callable[[mx.array], mx.array], HeadRef]:
137
+ """Split an mlx-lm `Model` into a hidden-state trunk and a `HeadRef`.
138
+
139
+ `model` is typed `Any` deliberately: importing mlx-lm's model classes here (just
140
+ for a type annotation) would defeat the point of `mlx_lm` being an optional, lazily
141
+ imported dependency. Support is instead verified structurally, at call time.
142
+ """
143
+ _require_mlx_lm()
144
+ module_name = type(model).__module__
145
+ if module_name not in _SUPPORTED_FAMILIES.values():
146
+ supported = ", ".join(sorted(_SUPPORTED_FAMILIES))
147
+ raise AdapterError(
148
+ f"unsupported model architecture (module {module_name!r}); "
149
+ f"mlx_train_perf's mlx-lm adapter supports: {supported}"
150
+ )
151
+ inner = model.model # the inner LlamaModel / Qwen3Model -- yields hidden states
152
+
153
+ def trunk(x: mx.array) -> mx.array:
154
+ return inner(x) # type: ignore[no-any-return]
155
+
156
+ if model.args.tie_word_embeddings:
157
+ head = _tied_head_from_embedding(inner.embed_tokens)
158
+ else:
159
+ head = _head_from_module(model.lm_head)
160
+ return trunk, head
161
+
162
+
163
+ def make_loss_fn(
164
+ model: Any,
165
+ *,
166
+ impl: Literal["auto", "kernel", "chunked", "naive"] = "auto",
167
+ allow_unverified_mlx: bool = False,
168
+ ) -> Callable[[Any, mx.array, mx.array], tuple[mx.array, mx.array]]:
169
+ """Build a loss callable matching mlx-lm's trainer contract:
170
+ `loss(model, batch, lengths) -> (loss, ntoks)` (see the module docstring for the
171
+ exact, version-cited contract this reproduces).
172
+
173
+ Fails fast: an unsupported architecture (`AdapterError`) or a missing `mlx-lm`
174
+ install (`MissingDependencyError`) is raised immediately, before any training step
175
+ runs, rather than on the first call.
176
+ """
177
+ # Fail fast only -- the (trunk, head) pair itself is discarded. `loss_fn` below
178
+ # re-derives both from the live `model` argument on every call (see the module
179
+ # docstring for why a construction-time snapshot would go stale).
180
+ split_model(model)
181
+
182
+ # The kernel/chunked/naive decision depends on the hidden dtype and row count `n`,
183
+ # neither of which is known until a real batch has flowed through the trunk -- so
184
+ # it is resolved on the first call to `loss_fn` and cached here for every later
185
+ # step (the decision itself does not vary across steps of the same training run).
186
+ resolved_impl: Literal["kernel", "chunked", "naive"] | None = None
187
+
188
+ def loss_fn(
189
+ model_arg: Any, batch: mx.array, lengths: mx.array
190
+ ) -> tuple[mx.array, mx.array]:
191
+ nonlocal resolved_impl
192
+ trunk, head = split_model(model_arg)
193
+ inputs = batch[:, :-1]
194
+ targets = batch[:, 1:]
195
+ hidden = trunk(inputs)
196
+ if resolved_impl is None:
197
+ n = hidden.shape[0] * hidden.shape[1]
198
+ resolved_impl = resolve_impl(
199
+ head=head, dtype=hidden.dtype, n=n, impl=impl,
200
+ allow_unverified_mlx=allow_unverified_mlx,
201
+ ).impl
202
+ # validate_targets=False: mlx_lm's trainer wraps this step in mx.compile, which
203
+ # forbids the range check's host sync; the trainer feeds in-range tokenizer ids, so
204
+ # the check is both unusable here and unnecessary. This is what lets `ours` run
205
+ # through the real compiled train() step, on equal footing with stock.
206
+ nll = linear_cross_entropy(hidden, head, targets, impl=resolved_impl,
207
+ reduction="none", validate_targets=False)
208
+ steps = mx.arange(1, targets.shape[1] + 1)
209
+ mask = (steps >= lengths[:, 0:1]) & (steps <= lengths[:, 1:])
210
+ ntoks = mask.sum()
211
+ loss = (nll * mask).astype(mx.float32).sum() / ntoks
212
+ return loss, ntoks
213
+
214
+ return loss_fn
File without changes
@@ -0,0 +1,153 @@
1
+ """Bench artifact conventions: identity-keyed JSON results with resume integrity.
2
+
3
+ Same shape as the committed `scripts/bench_quant_thresholds.py` house pattern: a run's
4
+ identity captures everything a stale-vs-fresh decision depends on, results are written
5
+ atomically (`.tmp` + rename), and
6
+ freshness requires an EXACT identity match plus `status == "ok"`.
7
+
8
+ Identity carries a `code_sha` -- the same SHA-256-over-name+bytes pattern
9
+ `scripts/bench_quant_thresholds.py._code_sha` uses -- over a declared list of files the
10
+ measured path actually depends on (the bench worker plus the loss/kernel/guard modules
11
+ it calls into). This is the field `result_is_fresh` actually needs: an installed
12
+ package's `importlib.metadata.version` is hatch-vcs tag-derived and only refreshes on an
13
+ explicit reinstall, NOT on every `uv sync` -- a plain checkout can therefore have edited
14
+ measured-path files while `version("mlx-train-perf")` still reports last release's
15
+ string, which would let a stale artifact be served as fresh (the exact failure this
16
+ harness exists to prevent). `package_version` is kept in the identity too, but purely as
17
+ informational provenance, not as the staleness signal.
18
+ """
19
+ import hashlib
20
+ import json
21
+ import platform
22
+ import uuid
23
+ from importlib.metadata import PackageNotFoundError, version
24
+ from pathlib import Path
25
+
26
+ from mlx_train_perf._compat import _installed_mlx_version
27
+ from mlx_train_perf.errors import BenchInputError
28
+
29
+ SCHEMA_VERSION = 1
30
+
31
+ _PACKAGE_ROOT = Path(__file__).resolve().parent.parent # .../src/mlx_train_perf
32
+
33
+ # Every file whose bytes change what a `loss_layer` (or later `train_step`) condition
34
+ # actually measures -- the bench worker itself, plus everything on the loss-computation
35
+ # path it calls into. Deliberately explicit rather than "every .py in the repo": a
36
+ # docs-only or CLI-only edit must NOT invalidate every bench artifact.
37
+ CODE_SHA_DEPS: tuple[Path, ...] = tuple(
38
+ _PACKAGE_ROOT / rel for rel in (
39
+ "bench/worker.py",
40
+ "core/loss.py",
41
+ "core/chunked.py",
42
+ "core/naive.py",
43
+ "core/kernel/launch.py",
44
+ "core/kernel/dispatch.py",
45
+ "core/kernel/source.py",
46
+ "core/guards.py",
47
+ "adapters/mlx_lm.py",
48
+ )
49
+ )
50
+
51
+ _RESERVED_PARAM_KEYS = ("kind", "session_id")
52
+
53
+
54
+ def _code_sha(deps: tuple[Path, ...]) -> str:
55
+ """SHA-256 over name+bytes of each dep file, in the given order -- any edit to a
56
+ measured-path file changes this, which is what `result_is_fresh` uses to invalidate
57
+ a prior artifact (same convention as `scripts/bench_quant_thresholds.py._code_sha`).
58
+ Recomputed fresh on every call (no caching): identity is meant to reflect the
59
+ ON-DISK state of these files at the moment it's built, not a snapshot from import
60
+ time."""
61
+ h = hashlib.sha256()
62
+ for p in deps:
63
+ h.update(p.name.encode())
64
+ h.update(p.read_bytes())
65
+ return h.hexdigest()[:16]
66
+
67
+
68
+ def _installed_mlx_lm_version() -> str | None:
69
+ """`mlx-lm` is an optional extra -- absent installs must not fail identity
70
+ construction, they just record `None` (still a stable, comparable identity value)."""
71
+ try:
72
+ return version("mlx-lm")
73
+ except PackageNotFoundError:
74
+ return None
75
+
76
+
77
+ def new_session_id() -> str:
78
+ """A fresh identity token for one bench invocation -- distinguishes artifacts from
79
+ different runs so `report`'s ratio logic can refuse to compare across them."""
80
+ return uuid.uuid4().hex
81
+
82
+
83
+ def run_identity(**kw: object) -> dict[str, object]:
84
+ """Everything a result's freshness depends on: mlx/mlx-lm versions, machine, macos,
85
+ a `code_sha` over the measured-path dependency files (see `CODE_SHA_DEPS` -- the
86
+ staleness signal `result_is_fresh` actually relies on), this package's own installed
87
+ version (informational only), plus whatever identity-relevant kwargs the caller
88
+ supplies (condition kind, grid point, dtype, impl, tile/chunk, session_id, ...). Two
89
+ calls returning equal dicts describe the SAME run in every way that matters for
90
+ reuse; any difference is what `result_is_fresh` treats as staleness.
91
+
92
+ A caller kwarg that happens to reuse one of THIS function's own internal field names
93
+ (e.g. a condition param literally named `code_sha`) would otherwise silently hijack
94
+ that field via `{**internal, **kw}` -- caught here by checking for overlap against
95
+ `internal`'s own keys directly, so the guard can never drift out of sync with the
96
+ field list (unlike a separately maintained reserved-name constant)."""
97
+ internal: dict[str, object] = {
98
+ "schema_version": SCHEMA_VERSION,
99
+ "mlx_version": _installed_mlx_version(),
100
+ "mlx_lm_version": _installed_mlx_lm_version(),
101
+ "machine": platform.machine(),
102
+ "macos": platform.mac_ver()[0],
103
+ "code_sha": _code_sha(CODE_SHA_DEPS),
104
+ "package_version": version("mlx-train-perf"),
105
+ }
106
+ collision = set(kw) & internal.keys()
107
+ if collision:
108
+ raise BenchInputError(
109
+ f"condition params must not use reserved identity key(s) {sorted(collision)} "
110
+ "-- computed internally by run_identity"
111
+ )
112
+ return {**internal, **kw}
113
+
114
+
115
+ def condition_identity(
116
+ *, kind: str, session_id: str, params: dict[str, object],
117
+ ) -> dict[str, object]:
118
+ """The single call site both `runner.run_conditions` and `worker.main` use to build
119
+ one condition's identity. `kind`/`session_id` are supplied separately by the
120
+ caller -- a `params` dict that happens to reuse either name would otherwise reach
121
+ `run_identity(kind=kind, session_id=session_id, **params)` and fail with a raw
122
+ `TypeError: got multiple values for keyword argument`; this raises a clean, named
123
+ error instead."""
124
+ for key in _RESERVED_PARAM_KEYS:
125
+ if key in params:
126
+ raise BenchInputError(
127
+ f"condition params must not use the reserved key {key!r} -- it is "
128
+ "supplied separately by the bench runner/worker"
129
+ )
130
+ return run_identity(kind=kind, session_id=session_id, **params)
131
+
132
+
133
+ def write_result(path: Path, identity: dict[str, object], status: str, **fields: object) -> None:
134
+ """Atomic write (tmp + rename) -- an interrupted worker leaves either the PRIOR
135
+ artifact or nothing at `path`, never a half-written JSON `result_is_fresh` could
136
+ misparse as fresh."""
137
+ path.parent.mkdir(parents=True, exist_ok=True)
138
+ tmp = path.with_suffix(".tmp")
139
+ tmp.write_text(json.dumps({"identity": identity, "status": status, **fields}, indent=2))
140
+ tmp.rename(path)
141
+
142
+
143
+ def result_is_fresh(path: Path, identity: dict[str, object]) -> bool:
144
+ """Fresh = parses, `status == "ok"`, identity matches EXACTLY. A missing file, a
145
+ corrupt one, a recorded error/refusal, or ANY identity field drift (including one the
146
+ caller didn't think to check) all trigger a recompute rather than a stale reuse."""
147
+ if not path.exists():
148
+ return False
149
+ try:
150
+ data = json.loads(path.read_text())
151
+ except (json.JSONDecodeError, OSError):
152
+ return False
153
+ return bool(data.get("status") == "ok" and data.get("identity") == identity)
@@ -0,0 +1,161 @@
1
+ """Subprocess-per-condition bench runner: resume-safe, same-session ratio reporting.
2
+
3
+ Each `Condition` gets its OWN Python process (`python -m mlx_train_perf.bench.worker`) --
4
+ the spike-proven isolation pattern (MLX's lazy allocator otherwise holds buffers across
5
+ runs within one process) -- and its own artifact, written the instant it finishes. A
6
+ condition whose artifact is already fresh is skipped entirely, never spawned; a worker
7
+ that exits nonzero (crashes) gets its failure recorded as a `status="error"` result here,
8
+ on the CALLER's side, so one bad condition never aborts the rest of the sweep.
9
+ """
10
+ import json
11
+ import subprocess
12
+ import sys
13
+ import tempfile
14
+ from collections import defaultdict
15
+ from dataclasses import dataclass
16
+ from pathlib import Path
17
+ from typing import cast
18
+
19
+ from mlx_train_perf.bench.artifacts import condition_identity, result_is_fresh, write_result
20
+
21
+ _STDERR_TAIL_CHARS = 4000 # enough to see the failing assertion/traceback, not a full dump
22
+
23
+
24
+ @dataclass(frozen=True, slots=True, kw_only=True)
25
+ class Condition:
26
+ name: str
27
+ kind: str
28
+ params: dict[str, object]
29
+
30
+
31
+ def _spawn_worker(config_path: Path) -> subprocess.CompletedProcess[str]:
32
+ return subprocess.run(
33
+ [sys.executable, "-m", "mlx_train_perf.bench.worker", "--config", str(config_path)],
34
+ capture_output=True, text=True, check=False,
35
+ )
36
+
37
+
38
+ def run_conditions(
39
+ conditions: list[Condition], out_dir: Path, *, session_id: str,
40
+ ) -> list[Path]:
41
+ out_dir.mkdir(parents=True, exist_ok=True)
42
+ paths: list[Path] = []
43
+ for condition in conditions:
44
+ out_path = out_dir / f"{condition.name}.json"
45
+ ident = condition_identity(
46
+ kind=condition.kind, session_id=session_id, params=condition.params,
47
+ )
48
+ paths.append(out_path)
49
+ if result_is_fresh(out_path, ident):
50
+ continue
51
+ # By definition stale (missing, corrupt, an old error/refusal, or an identity
52
+ # mismatch) -- remove it BEFORE spawning so `out_path.exists()` after the worker
53
+ # returns means exactly "THIS worker wrote it". Without this, a worker that fails
54
+ # silently (exits 0, writes nothing) would leave the stale artifact in place and
55
+ # the sweep would report someone else's old "ok" result as this run's truth.
56
+ out_path.unlink(missing_ok=True)
57
+
58
+ config = {
59
+ "kind": condition.kind, "params": condition.params, "session_id": session_id,
60
+ "out": str(out_path),
61
+ }
62
+ # The config lives in the SYSTEM temp dir, deliberately never inside `out_dir` --
63
+ # an interrupted run must not leave a stray `.json` there for a later glob over
64
+ # the artifact directory to misread as a result.
65
+ with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
66
+ json.dump(config, f)
67
+ config_path = Path(f.name)
68
+ try:
69
+ proc = _spawn_worker(config_path)
70
+ if proc.returncode != 0:
71
+ # The worker itself wrote nothing (it crashed before/without reaching its
72
+ # own `write_result`) -- this IS the sweep-level failure envelope, keyed
73
+ # by the SAME identity the worker would have used, so a later resume run
74
+ # (after the underlying bug is fixed) still sees it as stale and retries.
75
+ stderr_tail = (proc.stderr or proc.stdout or "")[-_STDERR_TAIL_CHARS:]
76
+ write_result(
77
+ out_path, ident, "error", error_type="WorkerCrashed",
78
+ error_msg=stderr_tail, returncode=proc.returncode,
79
+ )
80
+ elif not out_path.exists():
81
+ # A clean exit that wrote nothing is still a sweep-level failure (e.g. a
82
+ # worker that swallowed its own crash) -- recorded the same way, so a
83
+ # later resume still sees this condition as stale and retries it.
84
+ write_result(
85
+ out_path, ident, "error", error_type="WorkerExitedWithoutArtifact",
86
+ error_msg="worker exited 0 without writing an artifact", returncode=0,
87
+ )
88
+ finally:
89
+ config_path.unlink(missing_ok=True)
90
+ return paths
91
+
92
+
93
+ def _identity_of(entry: dict[str, object]) -> dict[str, object]:
94
+ return cast(dict[str, object], entry["identity"])
95
+
96
+
97
+ def _group_key(identity: dict[str, object]) -> tuple[tuple[str, object], ...]:
98
+ """The "same experimental grid point" key: every identity field EXCEPT `impl` (the
99
+ dimension being compared) and `session_id` (checked separately, as a gate on whether
100
+ a ratio may be emitted at all -- see `report`)."""
101
+ return tuple(sorted((k, v) for k, v in identity.items() if k not in ("impl", "session_id")))
102
+
103
+
104
+ def _ratio_label_and_value(
105
+ impl_a: object, impl_b: object, wall_a: object, wall_b: object,
106
+ ) -> tuple[str, float] | None:
107
+ """Direction is by MEASURED speed (`f"{slower_impl}/{faster_impl}"`, value = slower
108
+ `wall_s` / faster `wall_s`) -- never by the impls' alphabetical name order, which is
109
+ an accident of spelling and not a proxy for which one actually ran slower. `None`
110
+ when either wall time is missing or non-numeric (can't rank), or one is zero (can't
111
+ divide by it)."""
112
+ if not (isinstance(wall_a, int | float) and isinstance(wall_b, int | float)):
113
+ return None
114
+ if not wall_a or not wall_b:
115
+ return None
116
+ if wall_a >= wall_b:
117
+ slower_impl, slower_wall, faster_impl, faster_wall = impl_a, wall_a, impl_b, wall_b
118
+ else:
119
+ slower_impl, slower_wall, faster_impl, faster_wall = impl_b, wall_b, impl_a, wall_a
120
+ return f"{slower_impl}/{faster_impl}", slower_wall / faster_wall
121
+
122
+
123
+ def report(paths: list[Path]) -> dict[str, object]:
124
+ """Aggregate a set of artifacts into pairwise `impl` ratios at each shared grid
125
+ point (direction/magnitude: see `_ratio_label_and_value`). A pair that shares a grid
126
+ point but NOT a `session_id` is a cross-machine/cross-run comparison -- refused, and
127
+ named in `cross_session_excluded`, rather than silently blended into `ratios`."""
128
+ entries: list[dict[str, object]] = []
129
+ for p in paths:
130
+ try:
131
+ data = json.loads(p.read_text())
132
+ except (json.JSONDecodeError, OSError):
133
+ continue
134
+ if data.get("status") != "ok":
135
+ continue
136
+ entries.append(data)
137
+
138
+ groups: dict[tuple[tuple[str, object], ...], list[dict[str, object]]] = defaultdict(list)
139
+ for data in entries:
140
+ groups[_group_key(_identity_of(data))].append(data)
141
+
142
+ ratios: dict[str, float] = {}
143
+ cross_session_excluded: list[dict[str, object]] = []
144
+ for members in groups.values():
145
+ for i, a in enumerate(members):
146
+ for b in members[i + 1:]:
147
+ ident_a, ident_b = _identity_of(a), _identity_of(b)
148
+ impl_a, impl_b = ident_a.get("impl"), ident_b.get("impl")
149
+ if impl_a == impl_b:
150
+ continue # a genuine re-run of the same impl, not a comparison
151
+ if ident_a.get("session_id") != ident_b.get("session_id"):
152
+ cross_session_excluded.append({
153
+ "impl_a": impl_a, "session_a": ident_a.get("session_id"),
154
+ "impl_b": impl_b, "session_b": ident_b.get("session_id"),
155
+ })
156
+ continue
157
+ result = _ratio_label_and_value(impl_a, impl_b, a.get("wall_s"), b.get("wall_s"))
158
+ if result is not None:
159
+ label, value = result
160
+ ratios[label] = value
161
+ return {"ratios": ratios, "cross_session_excluded": cross_session_excluded}