File size: 6,636 Bytes
ecc81b3 eebb8d5 ecc81b3 eebb8d5 ecc81b3 eebb8d5 ecc81b3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 | """Compare two `pretrain.py` runs — one machine against another.
python benchmarks/compare.py "MPS bench" "CUDA bench" --out COMPARISON.md
Both runs start from bit-identical weights on bit-identical data, so every
number below is a difference the *arithmetic* produced. What that means
depends on the row:
- **the portable models** (``*_portable_*``, RNNs, conv, attention) run the
same pure-torch code on both machines, so a difference is float
non-associativity — different reduction orders on different hardware.
Expect small and growing-with-steps, not zero.
- **the vendored models** run the authors' *fused* kernels on CUDA and their
reference path on MPS. Those rows are fused-vs-reference, and they are the
ones worth reading: this is the check PLAN.md fixes as the rule for any
fast path.
A loss trajectory is chaotic, so late-step agreement is not the right test —
two runs of the same model can separate simply because gradient descent
amplifies. The honest summary is the *early* divergence and whether both runs
landed in the same place, which is why both are reported.
"""
from __future__ import annotations
import argparse
import json
import math
from pathlib import Path
import torch
def load(path: Path) -> dict:
manifest = json.loads((path / "manifest.json").read_text())
manifest["_dir"] = path
return manifest
def weights_of(run: Path, name: str) -> dict | None:
blob = run / name / "weights.pt"
if not blob.exists():
return None
return torch.load(blob, map_location="cpu", weights_only=True)
def weight_delta(a: dict, b: dict) -> tuple[float, float]:
"""Relative and absolute distance between two trained parameter sets."""
num = 0.0
den = 0.0
worst = 0.0
for key, va in a["model"].items():
vb = b["model"][key]
d = (va.double() - vb.double()).abs()
num += float((d**2).sum())
den += float((va.double() ** 2).sum())
worst = max(worst, float(d.max()))
return (math.sqrt(num / max(den, 1e-30)), worst)
def first_divergence(la: list[float], lb: list[float], tol: float = 1e-6) -> int | None:
"""The first step whose losses differ by more than `tol` relatively.
Says when the two machines stopped agreeing, which is far more informative
than how far apart they ended up: the end of a chaotic trajectory is not a
measurement of anything.
"""
for i, (x, y) in enumerate(zip(la, lb, strict=False)):
scale = max(abs(x), abs(y), 1e-12)
if abs(x - y) / scale > tol:
return i
return None
def _weights_provenance(left: dict, right: dict) -> list[str]:
"""Say — and check — where each run's starting weights came from.
The comparison used to assert that both machines began from bit-identical
weights. For S4 and S4D that was false: `torch.linalg.eigh` fixes
eigenvectors only up to a phase, so their `B` and `P` differ between macOS
and Linux under the same seed, and the comparison reported a 2.6e-01
difference that was two different models rather than two devices.
A claim the instrument cannot check is worth less than one it can, so this
reads what the runs recorded instead of assuming.
"""
def sources(run: dict) -> set[str]:
return {m.get("weights_from", "seed") for m in run["models"] if "error" not in m}
both = sources(left) | sources(right)
if both <= {"written", "loaded"}:
return [
"Both runs load one shared set of starting weights, so the initial",
"conditions are bit-identical and every difference below is arithmetic.",
"",
]
return [
"> **The two runs did not share starting weights** "
f"({', '.join(sorted(both))}). The seed alone does not give identical",
"> S4/S4D weights across platforms — `eigh`'s eigenvectors are fixed only",
"> up to a phase — so differences below may be initialisation rather than",
"> arithmetic. Re-run both sides with `--init <dir>` pointing at the same",
"> directory. See benchmarks/init_weights.py.",
"",
]
def main() -> int:
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument("left")
ap.add_argument("right")
ap.add_argument("--out", default=None, help="write a markdown table here")
args = ap.parse_args()
left, right = load(Path(args.left)), load(Path(args.right))
by_name = {
"left": {m["name"]: m for m in left["models"] if "error" not in m},
"right": {m["name"]: m for m in right["models"] if "error" not in m},
}
shared = [n for n in by_name["left"] if n in by_name["right"]]
lines = [
"# Device comparison",
"",
f"- **{args.left}** — {left['device_name']} · torch {left['torch']}"
f" · torch-dimensions {left['torch_dimensions']}",
f"- **{args.right}** — {right['device_name']} · torch {right['torch']}"
f" · torch-dimensions {right['torch_dimensions']}",
"",
f"Both runs: seed {left['seed']}, {left['steps']} steps, batch {left['batch']}.",
"Data is drawn on CPU from a seeded generator, so both machines see the",
"same batches in the same order.",
"",
*_weights_provenance(left, right),
"| model | loss (left) | loss (right) | Δ loss | first divergence "
"| rel Δw | max Δw | speed |",
"|---|---|---|---|---|---|---|---|",
]
for name in shared:
a, b = by_name["left"][name], by_name["right"][name]
wa, wb = weights_of(left["_dir"], name), weights_of(right["_dir"], name)
if wa and wb:
rel_w, max_w = weight_delta(wa, wb)
wcol, mcol = f"{rel_w:.2e}", f"{max_w:.2e}"
else:
wcol = mcol = "—"
step = first_divergence(a["losses"], b["losses"])
dcol = "identical" if step is None else f"step {step}"
speed = b["steps_per_second"] / max(a["steps_per_second"], 1e-9)
lines.append(
f"| `{name}` | {a['loss_final']:.5f} | {b['loss_final']:.5f} | "
f"{abs(a['loss_final'] - b['loss_final']):.2e} | {dcol} | {wcol} | {mcol} | "
f"{speed:.2f}× |"
)
missing = sorted(set(by_name["left"]) ^ set(by_name["right"]))
if missing:
lines += ["", f"Not in both runs: {', '.join('`' + m + '`' for m in missing)}."]
text = "\n".join(lines) + "\n"
print(text)
if args.out:
Path(args.out).write_text(text)
print(f"wrote {args.out}")
return 0
if __name__ == "__main__":
raise SystemExit(main())
|