File size: 6,112 Bytes
eebb8d5 | 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 | """The starting weights a cross-device comparison actually shares.
`pretrain.py` and `agreement.py` both said, in a docstring, that building on
CPU under a fixed seed gives bit-identical weights on any machine, and every
number they produced rested on it. For most models it is true — an LSTM built
that way hashes identically on macOS and on Linux.
For S4 and S4D it is false, and not because of a bug. `hippo.nplr` diagonalises
the HiPPO matrix with `torch.linalg.eigh`. Eigen*values* are unique and matched
across the two platforms to every digit printed, which is why `A_imag` looked
fine. Eigen*vectors* are fixed only up to a phase, and macOS Accelerate and
Linux LAPACK are each free to return a different one. `B` and `P` are
projections through those vectors, so they inherit it: measured across a Mac
Studio and an RTX 5090 box, `B` differed by a relative 1.5 and `P` by 0.53
while `A_imag` was identical.
The comparison was therefore reporting a 2.6e-01 output difference for the
vendored S4D — unchanged in float64, which reads exactly like a different
kernel — when the two machines had simply built two different models. With the
weights shared, the same pair agrees at 4e-07.
A test on one machine cannot observe a cross-platform difference. What it can
do is pin the mechanism that now carries the assumption: that weights written
by one run are what a later run gets, exactly, in place of whatever the seed
would have produced.
"""
from __future__ import annotations
import importlib.util
import sys
from pathlib import Path
import pytest
import torch
import torch.nn as nn
import torch_dimensions as td
ROOT = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(ROOT / "benchmarks"))
spec = importlib.util.spec_from_file_location(
"_td_init_weights", ROOT / "benchmarks" / "init_weights.py"
)
init_weights = importlib.util.module_from_spec(spec)
sys.modules["_td_init_weights"] = init_weights
spec.loader.exec_module(init_weights)
LAT = td.Lattice(shape=(4, 5), names=("h", "w"), time=True)
def test_the_first_run_writes_and_the_second_loads(tmp_path):
a, b = nn.Linear(8, 8), nn.Linear(8, 8)
assert init_weights.sync(a, tmp_path, "m") == "written"
assert init_weights.sync(b, tmp_path, "m") == "loaded"
for pa, pb in zip(a.parameters(), b.parameters(), strict=True):
assert torch.equal(pa, pb)
def test_loading_overrides_whatever_the_seed_produced(tmp_path):
"""The point of the mechanism: the loaded values win over construction.
If `sync` returned "loaded" while leaving the model on its own weights,
every comparison would still be measuring initialisation drift and would
still report a plausible-looking number.
"""
torch.manual_seed(0)
reference = nn.Linear(8, 8)
init_weights.sync(reference, tmp_path, "m")
torch.manual_seed(999) # deliberately a different draw
other = nn.Linear(8, 8)
assert not torch.equal(other.weight, reference.weight)
assert init_weights.sync(other, tmp_path, "m") == "loaded"
assert torch.equal(other.weight, reference.weight)
assert torch.equal(other.bias, reference.bias)
def test_no_store_means_no_change_and_says_so(tmp_path):
"""`--init` is opt-in; without it the behaviour is exactly what it was."""
torch.manual_seed(3)
model = nn.Linear(8, 8)
before = model.weight.detach().clone()
assert init_weights.sync(model, None, "m") == "seed"
assert torch.equal(model.weight, before)
def test_every_parameter_and_buffer_round_trips_for_a_real_model(tmp_path):
"""A `state_dict` is not just parameters. S4's kernel keeps buffers, and a
mechanism that restored parameters while leaving buffers to the seed would
reintroduce the bug in the exact place it came from."""
pytest.importorskip("einops", reason="the vendored S4 needs the [upstream] extra")
pytest.importorskip("hydra", reason="the s4 pipeline needs hydra-core")
torch.manual_seed(0)
first = td.S4(32, 2, LAT, d_input=1, d_state=16)
init_weights.sync(first, tmp_path, "s4")
torch.manual_seed(1)
second = td.S4(32, 2, LAT, d_input=1, d_state=16)
assert init_weights.sync(second, tmp_path, "s4") == "loaded"
sd_a, sd_b = first.state_dict(), second.state_dict()
assert set(sd_a) == set(sd_b)
for key in sd_a:
assert torch.equal(sd_a[key], sd_b[key]), f"{key} did not round trip"
def test_shared_weights_make_two_builds_agree_where_the_seed_would_not(tmp_path):
"""End to end, in the shape the benchmark uses it: build, sync, and the two
models compute the same thing. On one machine the seed would also have
achieved this — the value of the test is that it fails loudly if `sync`
ever stops applying, which is the failure that hid for a whole run."""
pytest.importorskip("einops", reason="the vendored S4 needs the [upstream] extra")
pytest.importorskip("hydra", reason="the s4 pipeline needs hydra-core")
torch.manual_seed(0)
a = td.S4D(32, 2, LAT, d_input=1, d_state=16).eval()
init_weights.sync(a, tmp_path, "s4d")
torch.manual_seed(7)
b = td.S4D(32, 2, LAT, d_input=1, d_state=16).eval()
init_weights.sync(b, tmp_path, "s4d")
x = torch.randn(2, 3, *LAT.shape, 1)
with torch.no_grad():
assert torch.equal(a(x), b(x))
def test_the_eigenvalues_are_the_reproducible_part(tmp_path):
"""Why the bug was invisible: the part of the decomposition that *is*
unique agrees, so `A_imag` matched across platforms to twelve decimals and
the initialisation looked sound. Pinned here so the diagnosis stays
attached to the code it explains."""
pytest.importorskip("einops", reason="the vendored S4 needs the [upstream] extra")
from torch_dimensions._vendor.s4.src.models.hippo.hippo import nplr
w1, p1, b1, v1 = nplr("legs", 32)
w2, p2, b2, v2 = nplr("legs", 32)
# Same machine, so everything repeats; the eigenvalues are the only part
# that also repeats across machines.
assert torch.allclose(w1, w2)
assert w1.is_complex(), "the eigenvalues are the spectrum of the HiPPO matrix"
|