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.
- mlx_train_perf/__init__.py +12 -0
- mlx_train_perf/_compat.py +21 -0
- mlx_train_perf/adapters/__init__.py +0 -0
- mlx_train_perf/adapters/mlx_lm.py +214 -0
- mlx_train_perf/bench/__init__.py +0 -0
- mlx_train_perf/bench/artifacts.py +153 -0
- mlx_train_perf/bench/runner.py +161 -0
- mlx_train_perf/bench/worker.py +450 -0
- mlx_train_perf/cli.py +265 -0
- mlx_train_perf/core/__init__.py +0 -0
- mlx_train_perf/core/chunked.py +183 -0
- mlx_train_perf/core/guards.py +38 -0
- mlx_train_perf/core/kernel/__init__.py +0 -0
- mlx_train_perf/core/kernel/dispatch.py +25 -0
- mlx_train_perf/core/kernel/launch.py +522 -0
- mlx_train_perf/core/kernel/source.py +702 -0
- mlx_train_perf/core/loss.py +392 -0
- mlx_train_perf/core/naive.py +14 -0
- mlx_train_perf/devtools/__init__.py +0 -0
- mlx_train_perf/devtools/regpressure.py +279 -0
- mlx_train_perf/errors.py +49 -0
- mlx_train_perf/plan/__init__.py +0 -0
- mlx_train_perf/plan/calibration.py +69 -0
- mlx_train_perf/plan/calibration_data.json +16 -0
- mlx_train_perf/plan/estimate.py +405 -0
- mlx_train_perf/py.typed +0 -0
- mlx_train_perf-0.1.0.dist-info/METADATA +123 -0
- mlx_train_perf-0.1.0.dist-info/RECORD +31 -0
- mlx_train_perf-0.1.0.dist-info/WHEEL +4 -0
- mlx_train_perf-0.1.0.dist-info/entry_points.txt +2 -0
- mlx_train_perf-0.1.0.dist-info/licenses/LICENSE +21 -0
|
@@ -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}
|