File size: 10,942 Bytes
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
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
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
"""Mamba-3: the authors' block, and our transcription of their scan.

Mamba-3's recurrence exists upstream only as Triton, so unlike Mamba-1 and
Mamba-2 there is no reference implementation of theirs to defer to and no way,
on a machine without CUDA, to compare against their kernel. What *can* be
established is established here, and the file is explicit about the gap:

- the chunked (matmul) and recurrent (loop) forms are written independently
  and must agree to float64 precision;
- in the ``trap -> 1`` limit a third, direct O(L^2) sum must agree;
- ``trap -> 0`` must make the current step contribute nothing, which is what
  "trapezoidal" claims;
- the chunk size must not change the answer;
- gradients must be right (gradcheck), because the Triton backward is not
  ported — autograd differentiates the forward instead.

None of that proves equality with their kernel. It proves the recurrence
implemented here is the one written down, computed consistently.
"""

from __future__ import annotations

import pytest
import torch

import torch_dimensions as td

pytest.importorskip("einops", reason="the vendored Mamba-3 block needs the [upstream] extra")

from torch_dimensions.mixers.mamba3_compat import mamba3_siso_combined  # noqa: E402

F64 = torch.float64


def _inputs(b=2, length=37, hq=1, h=4, dqk=16, dv=8, nang=4, seed=0, gates=True):
    torch.manual_seed(seed)

    def g(*s):
        return torch.randn(*s, dtype=F64)

    return {
        "Q": g(b, length, hq, dqk),
        "K": g(b, length, hq, dqk),
        "V": g(b, length, h, dv),
        # A is negative and dt positive upstream (heavy-tail activation, then a
        # clamp), so every decay exponent is non-positive; sampling any other
        # way would test a regime the model cannot reach.
        "ADT": -torch.rand(b, h, length, dtype=F64) * 0.5 - 1e-3,
        "DT": torch.rand(b, h, length, dtype=F64) * 0.1 + 1e-3,
        "Trap": g(b, h, length),
        "Q_bias": g(h, dqk),
        "K_bias": g(h, dqk),
        "Angles": g(b, length, h, nang),
        "D": g(h) if gates else None,
        "Z": g(b, length, h, dv) if gates else None,
    }


@pytest.mark.parametrize(
    "kw",
    [
        {},
        {"hq": 2},  # grouped query attention: Q/K broadcast over head groups
        {"gates": False},  # no D skip, no Z gate
        {"length": 200},  # several chunks
        {"length": 5},  # shorter than one chunk
    ],
    ids=["dense", "gqa", "no-gates", "long", "short"],
)
def test_chunked_and_recurrent_forms_agree(kw):
    """The two independently written forms of the same recurrence.

    This is the load-bearing check: the chunked form folds each pair's two
    trapezoid visits into one weight, and if that algebra were wrong these
    would diverge.
    """
    args = _inputs(**kw)
    chunked = mamba3_siso_combined(**args, chunk_size=16)
    recurrent = mamba3_siso_combined(**args, chunk_size=16, recurrent=True)
    scale = recurrent.abs().max().item()
    assert (chunked - recurrent).abs().max().item() < 1e-12 * max(scale, 1.0)


def test_chunk_size_does_not_change_the_answer():
    args = _inputs(length=97)
    base = mamba3_siso_combined(**args, chunk_size=8)
    for chunk in (16, 32, 64, 128):
        got = mamba3_siso_combined(**args, chunk_size=chunk)
        assert (got - base).abs().max().item() < 1e-12


def test_trap_to_one_matches_an_independent_direct_sum():
    """With ``trap -> 1`` the previous-step term vanishes and the recurrence
    collapses to a decayed linear attention, which a third implementation —
    an explicit double loop, sharing no code with either scan — can state."""
    args = _inputs(length=12, h=2, dqk=8, dv=4, gates=False)
    args["Angles"] = torch.zeros_like(args["Angles"])  # isolate the scan from the rotation
    b, length, h = args["V"].shape[0], args["V"].shape[1], args["V"].shape[2]
    args["Trap"] = torch.full((b, h, length), 40.0, dtype=F64)  # sigmoid(40) = 1 - 4e-18

    out = mamba3_siso_combined(**args, chunk_size=4)

    q = args["Q"].expand(b, length, h, args["Q"].shape[-1]) + args["Q_bias"]
    k = args["K"].expand(b, length, h, args["K"].shape[-1]) + args["K_bias"]
    dt = args["DT"].movedim(-1, 1)
    cs = args["ADT"].movedim(-1, 1).cumsum(1)
    ref = torch.zeros_like(out)
    for t in range(length):
        for j in range(t + 1):
            weight = (cs[:, t] - cs[:, j]).exp() * dt[:, j]
            ref[:, t] += (weight * (q[:, t] * k[:, j]).sum(-1)).unsqueeze(-1) * args["V"][:, j]
    assert (out - ref).abs().max().item() < 1e-12


def test_trap_to_zero_removes_the_current_step():
    """The trapezoid's other end: with ``trap -> 0`` a pair contributes only
    on the step *after* it arrives, so the first output is proportional to
    ``sigmoid(trap)`` and vanishes with it."""
    args = _inputs(length=12, h=2, dqk=8, dv=4, gates=False)
    b, length, h = args["V"].shape[0], args["V"].shape[1], args["V"].shape[2]
    first = {}
    for value in (-20.0, -40.0):
        args["Trap"] = torch.full((b, h, length), value, dtype=F64)
        first[value] = mamba3_siso_combined(**args, chunk_size=4)[:, 0].abs().max().item()
    # sigmoid(-40)/sigmoid(-20) ~ 2e-9, and the outputs must track it.
    ratio = first[-40.0] / first[-20.0]
    assert 1e-9 < ratio < 1e-8, first


def test_gradients_are_correct():
    """The Triton backward (1,788 lines) is not ported: autograd differentiates
    the forward instead, so the forward being differentiable *correctly* is
    what has to hold."""
    args = _inputs(b=1, length=10, h=2, dqk=8, dv=4, nang=2)
    fixed = {k: args[k] for k in ("Q_bias", "K_bias", "D", "Z")}
    diff = ["Q", "K", "V", "ADT", "DT", "Trap", "Angles"]
    tensors = tuple(args[k].clone().requires_grad_(True) for k in diff)

    def run(*ts):
        return mamba3_siso_combined(**dict(zip(diff, ts, strict=True)), **fixed, chunk_size=4)

    assert torch.autograd.gradcheck(run, tensors, eps=1e-6, atol=1e-7)


def test_rotation_is_the_interleaved_convention():
    """Their kernel pairs adjacent components — ``tl.reshape(x, [D//2, 2])``
    then ``tl.split`` — not the half-and-half split some RoPE code uses. A
    single non-zero angle must therefore mix components 0 and 1, and leave
    component 2 alone."""
    from torch_dimensions.mixers.mamba3_compat import _rotate

    x = torch.tensor([[1.0, 0.0, 1.0, 0.0]], dtype=F64)
    cos = torch.tensor([[0.0, 1.0]], dtype=F64)  # 90 degrees on the first pair only
    sin = torch.tensor([[1.0, 0.0]], dtype=F64)
    got = _rotate(x, cos, sin)
    assert torch.allclose(got, torch.tensor([[0.0, 1.0, 1.0, 0.0]], dtype=F64))


def test_angles_beyond_the_rotary_width_are_not_rotated():
    """``headdim_angles`` can be smaller than ``headdim_qk // 2``; the tail
    pairs get cos=1, sin=0 upstream and must pass through untouched."""
    wide = _inputs(length=8, dqk=16, nang=2, gates=False)
    narrow = dict(wide)
    # Zeroing the angles must equal rotating with none of them set.
    narrow["Angles"] = torch.zeros_like(wide["Angles"])
    rotated = mamba3_siso_combined(**narrow, chunk_size=4)
    assert torch.isfinite(rotated).all()


def test_unsupported_paths_are_refused_rather_than_approximated():
    args = _inputs(length=8)
    with pytest.raises(NotImplementedError, match="cu_seqlens"):
        mamba3_siso_combined(**args, cu_seqlens=torch.tensor([0, 8], dtype=torch.int32))
    states = (torch.zeros(1), torch.zeros(1), torch.zeros(1), torch.zeros(1))
    with pytest.raises(NotImplementedError, match="input states"):
        mamba3_siso_combined(**args, Input_States=states)


# --- the block, and the model API --------------------------------------------


def test_the_vendored_block_runs_and_learns_shape():
    from torch_dimensions.mixers import Mamba3Mixer

    torch.manual_seed(0)
    mixer = Mamba3Mixer(64, d_state=32, headdim=16)
    x = torch.randn(2, 24, 64, requires_grad=True)
    y = mixer(x)
    assert y.shape == x.shape
    y.pow(2).mean().backward()
    assert torch.isfinite(x.grad).all()
    assert all(torch.isfinite(p.grad).all() for p in mixer.parameters() if p.grad is not None)


def test_mimo_is_refused_with_a_reason():
    from torch_dimensions.mixers import Mamba3Mixer

    with pytest.raises(ValueError, match="MIMO"):
        Mamba3Mixer(64, d_state=32, headdim=16, is_mimo=True)


@pytest.mark.parametrize("spelling", ["version", "name", "nd"])
def test_every_spelling_builds_the_same_model(spelling, tmp_path):
    from torch_dimensions.mixers import Mamba3Mixer

    kw = {"mixer_kwargs": {"d_state": 32, "headdim": 16}}
    lat = td.Lattice(shape=(4, 5), names=("y", "x"))
    if spelling == "version":
        model = td.Mamba(32, 2, lat, version=3, **kw)
    elif spelling == "name":
        model = td.Mamba3(32, 2, lat, **kw)
    else:
        model = td.Mamba3ND(32, 2, dim=2, shape=(4, 5), time=False, **kw)
    model.eval()

    assert isinstance(model.nd.mixers[0], Mamba3Mixer)
    assert model.config["version"] == 3

    x = torch.randn(2, 4, 5, 32)
    path = tmp_path / f"{spelling}.td"
    model.save(path)
    with torch.no_grad():
        assert torch.equal(model(x), td.load(path).eval()(x))


def test_mamba3_has_no_portable_build():
    with pytest.raises(ValueError, match="no portable build of Mamba-3"):
        td.Mamba3(32, 1, portable=True)


def test_version_three_is_registered_for_configs():
    model = td.build({"kind": "mamba3", "d_model": 32, "n_layers": 1})
    assert model.config["version"] == 3


@pytest.mark.skipif(not torch.backends.mps.is_available(), reason="no MPS device")
def test_mamba3_on_mps_matches_cpu():
    from torch_dimensions.mixers import Mamba3Mixer

    torch.manual_seed(1)
    cpu = Mamba3Mixer(32, d_state=32, headdim=16).eval()
    mps = Mamba3Mixer(32, d_state=32, headdim=16).to("mps")
    mps.load_state_dict({k: v.to("mps") for k, v in cpu.state_dict().items()})
    mps.eval()
    x = torch.randn(2, 24, 32)
    with torch.no_grad():
        assert (cpu(x) - mps(x.to("mps")).cpu()).abs().max().item() < 1e-4

    grad_in = torch.randn(2, 24, 32, device="mps", requires_grad=True)
    mps.train()
    mps(grad_in).pow(2).mean().backward()
    assert torch.isfinite(grad_in.grad).all()


# --- which implementation runs -----------------------------------------------


def test_dispatch_prefers_torch_off_cuda_and_when_forced(monkeypatch):
    from torch_dimensions.mixers._kernels import forced_torch, prefer_upstream

    assert not prefer_upstream(torch.zeros(1))  # CPU tensor: no fused kernel
    monkeypatch.setenv("TD_FORCE_TORCH_KERNELS", "1")
    assert forced_torch()
    assert not prefer_upstream(torch.zeros(1))


def test_load_upstream_returns_none_for_a_missing_kernel():
    from torch_dimensions.mixers._kernels import load_upstream

    assert load_upstream("torch_dimensions._no_such_module", "whatever") is None
    assert load_upstream("torch_dimensions.mixers._kernels", "prefer_upstream") is not None