File size: 9,041 Bytes
611aea1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
"""Parameter groups that honour what the upstream models ask for.

`A` controls how fast a state decays and `dt` its timescale, both inside an
exponential — so weight decay on them is not a mild regulariser but a change
to the dynamics, and a learning rate suited to a projection walks them out of
the stable region. Both upstream repos say so in the parameters themselves
(`_optim` in s4's kernel, `_no_weight_decay` on Mamba's `A_log`/`D`/`dt_bias`),
and those tags do nothing unless an optimizer reads them.

Which means the obvious line — `AdamW(model.parameters(), lr=...)` — trains an
S4 the way its authors avoid, silently. These tests pin the reading of the
tags, because the failure mode is a model that trains, converges to something,
and is misconfigured the whole way.
"""

from __future__ import annotations

import pytest
import torch

import torch_dimensions as td

pytest.importorskip("einops", reason="the tagged parameters live in the vendored code")
pytest.importorskip("hydra", reason="the s4 pipeline needs hydra-core")

LAT = td.Lattice(shape=(4, 5), names=("h", "w"), time=True)


def _all_params(groups):
    return [p for g in groups for p in g["params"]]


def test_every_parameter_lands_in_exactly_one_group():
    """A parameter in two groups is stepped twice; one in none is frozen
    without anyone saying so."""
    model = td.S4(32, 2, LAT, d_input=1, d_state=16)
    groups = td.param_groups(model, lr=1e-3)
    listed = _all_params(groups)
    ids = [id(p) for p in listed]
    assert len(ids) == len(set(ids)), "a parameter appears in more than one group"
    expected = {id(p) for p in model.parameters() if p.requires_grad}
    assert set(ids) == expected, "a parameter was dropped or invented"


def test_s4_kernel_parameters_keep_the_settings_upstream_gave_them():
    """s4 attaches `_optim` to its SSM kernel parameters. Those settings are
    the authors', not ours, and must survive verbatim."""
    model = td.S4(32, 2, LAT, d_input=1, d_state=16)
    tagged = [p for p in model.parameters() if getattr(p, "_optim", None)]
    assert tagged, "the vendored s4 kernel no longer tags its parameters"

    groups = td.param_groups(model, lr=1e-2, weight_decay=0.1)
    for param in tagged:
        group = next(g for g in groups if any(p is param for p in g["params"]))
        assert group["weight_decay"] == 0.0
        assert group["lr"] <= td.optim.SSM_MAX_LR


def test_mamba_no_weight_decay_parameters_get_none():
    """`A_log`, `D` and `dt_bias` carry `_no_weight_decay`. Decay on `A_log`
    pulls every state toward the same timescale, which is the opposite of what
    a bank of states is for."""
    model = td.Mamba(32, 2, LAT, d_input=1, d_state=8)
    tagged = [p for p in model.parameters() if getattr(p, "_no_weight_decay", False)]
    assert tagged, "the vendored mamba block no longer tags its parameters"

    groups = td.param_groups(model, lr=1e-3, weight_decay=0.1)
    for param in tagged:
        group = next(g for g in groups if any(p is param for p in g["params"]))
        assert group["weight_decay"] == 0.0


def test_the_ssm_rate_is_capped_not_raised():
    """Upstream fixes the SSM rate at 1e-3. A caller asking for 1e-2 gets it
    for ordinary weights and 1e-3 for the SSM; a caller asking for 1e-5 keeps
    1e-5 everywhere rather than having it raised to the cap."""
    model = td.S4(32, 1, LAT, d_input=1, d_state=16)

    high = td.param_groups(model, lr=1e-2)
    assert max(g["lr"] for g in high) == pytest.approx(1e-2)
    assert min(g["lr"] for g in high) == pytest.approx(1e-3)

    low = td.param_groups(model, lr=1e-5)
    assert all(g["lr"] == pytest.approx(1e-5) for g in low)


def test_norms_and_biases_are_excluded_from_decay_by_default():
    model = td.LSTM(32, 2, LAT, d_input=1)
    groups = td.param_groups(model, lr=1e-3, weight_decay=0.1)
    for group in groups:
        if group["weight_decay"] > 0:
            assert all(p.ndim > 1 for p in group["params"])

    including = td.param_groups(model, lr=1e-3, weight_decay=0.1, decay_1d=True)
    assert any(p.ndim <= 1 for g in including if g["weight_decay"] > 0 for p in g["params"])


def test_frozen_parameters_are_left_out():
    model = td.LSTM(16, 1, LAT, d_input=1)
    frozen = next(iter(model.parameters()))
    frozen.requires_grad_(False)
    listed = _all_params(td.param_groups(model, lr=1e-3))
    assert all(p is not frozen for p in listed)


def test_the_groups_drive_a_real_optimizer_step():
    """The point of the split is that the optimizer applies it — checked by
    stepping and confirming the tagged parameters moved less."""
    torch.manual_seed(0)
    model = td.S4(32, 1, LAT, d_input=1, d_state=16)
    tagged = [p for p in model.parameters() if getattr(p, "_optim", None)]
    before = [p.detach().clone() for p in tagged]

    opt = torch.optim.AdamW(td.param_groups(model, lr=1e-1), lr=1e-1)
    x = torch.randn(2, 3, 4, 5, 1)
    model(x).pow(2).mean().backward()
    opt.step()

    moved = max(float((p.detach() - b).abs().max()) for p, b in zip(tagged, before, strict=True))
    # Adam's step is ~lr in magnitude, so a 100x smaller rate must show.
    assert moved < 1e-2, f"SSM parameters moved by {moved:.3e} despite the capped rate"


# --- the schedule -------------------------------------------------------------


def test_warmup_reaches_the_peak_then_decays():
    model = td.LSTM(16, 1, LAT, d_input=1)
    opt = torch.optim.AdamW(td.param_groups(model, lr=1e-2), lr=1e-2)
    sched = td.warmup_cosine(opt, warmup=10, total=100)

    seen = []
    for _ in range(100):
        seen.append(opt.param_groups[0]["lr"])
        opt.step()
        sched.step()

    assert seen[0] < seen[5] < seen[9]  # ramping
    assert seen[9] == pytest.approx(1e-2, rel=1e-6)  # peak at the end of warmup
    assert seen[-1] < seen[50] < seen[10]  # decaying thereafter
    # `seen` records before each step, so the last entry is the rate *going
    # into* the final step rather than the one after it: near zero, not zero.
    assert seen[-1] < 1e-5
    opt.step()
    sched.step()
    assert opt.param_groups[0]["lr"] == pytest.approx(0.0, abs=1e-9)


def test_the_schedule_preserves_the_ratio_between_groups():
    """Scaling is multiplicative on purpose: a schedule that flattened every
    group to one rate would undo the separation that is the whole point."""
    model = td.S4(32, 1, LAT, d_input=1, d_state=16)
    opt = torch.optim.AdamW(td.param_groups(model, lr=1e-2), lr=1e-2)
    sched = td.warmup_cosine(opt, warmup=5, total=50)

    for _ in range(25):
        opt.step()
        sched.step()

    rates = [g["lr"] for g in opt.param_groups]
    assert max(rates) == pytest.approx(10 * min(rates), rel=1e-6)


def test_a_floor_keeps_the_rate_above_zero():
    model = td.LSTM(16, 1, LAT, d_input=1)
    opt = torch.optim.AdamW(td.param_groups(model, lr=1e-2), lr=1e-2)
    sched = td.warmup_cosine(opt, warmup=2, total=20, floor=0.1)
    for _ in range(20):
        opt.step()
        sched.step()
    assert opt.param_groups[0]["lr"] == pytest.approx(1e-3, rel=1e-6)


def test_an_impossible_schedule_is_refused():
    model = td.LSTM(16, 1, LAT, d_input=1)
    opt = torch.optim.AdamW(td.param_groups(model, lr=1e-3), lr=1e-3)
    with pytest.raises(ValueError, match="total > 0"):
        td.warmup_cosine(opt, warmup=5, total=0)


def test_the_recipe_trains_every_family_without_per_family_tuning():
    """The practical claim. Hand-picked per-family learning rates were needed
    only because the SSM parameters were being trained at the projection rate;
    with the groups and a warmup, one recipe holds for all of them."""
    lat = td.Lattice(shape=(4, 5), names=("h", "w"), time=True)
    builders = {
        "s4": lambda: td.S4(32, 2, lat, d_input=1, d_state=16),
        "mamba": lambda: td.Mamba(32, 2, lat, d_input=1, d_state=8),
        "transformer": lambda: td.Transformer(32, 2, lat, d_input=1),
        "lstm": lambda: td.LSTM(32, 2, lat, d_input=1),
    }
    steps = 60
    for name, build in builders.items():
        torch.manual_seed(0)
        model = build()
        head = torch.nn.Linear(32, 1)
        opt = torch.optim.AdamW(
            td.param_groups(model, lr=3e-3) + [{"params": head.parameters(), "lr": 3e-3}],
            lr=3e-3,
            betas=(0.9, 0.95),
        )
        sched = td.warmup_cosine(opt, warmup=steps // 10, total=steps)

        gen = torch.Generator().manual_seed(1)
        losses = []
        for _ in range(steps):
            x = torch.randn(3, 4, 4, 5, 1, generator=gen)
            loss = (head(model(x)) - x.cumsum(dim=3)).pow(2).mean()
            opt.zero_grad()
            loss.backward()
            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
            opt.step()
            sched.step()
            losses.append(float(loss))

        assert all(torch.isfinite(torch.tensor(losses))), f"{name} diverged to non-finite"
        assert losses[-1] < losses[0], (
            f"{name} did not improve: {losses[0]:.3f} -> {losses[-1]:.3f}"
        )