fractus-cte / tests /test_attention_equivalence.py
thefinalboss's picture
opt: cumsum/chunked attention kernels, memory-flat CE, block checkpointing, v2 trainer (proven equivalent, 46 tests)
8d9996a verified
Raw History Blame
15.3 kB
"""Equivalence proofs: cumsum attention vs einsum reference.
Open-heart rule: any production-path change must be PROVEN mathematically
equivalent to the reference before it may wrap a live checkpoint.
Covered here (opt 1, 2026-08-22):
- forward outputs match (with and without state carry)
- gradients wrt q, k, v match
- carried final states (S_final, z_final) match
- both match the scalar looped reference (_linear_attention_causal_one_head)
"""
import sys
from pathlib import Path
import pytest
import torch
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from fractus.nn.attention import FractalLinearAttention
from fractus.nn.stats import elu_plus_one
ATOL = 1e-4 # float32 accumulation-order tolerance
RTOL = 1e-4
def _make(attn, G, L=32, D=64, seed=0):
torch.manual_seed(seed)
q = torch.rand(G, L, D) + 0.5 # positive like elu+1 features
k = torch.rand(G, L, D) + 0.5
v = torch.randn(G, L, D) * 0.5
S0 = torch.randn(G, D, D) * 0.05
z0 = torch.rand(G, D) + 0.5 # z from positive features
return q, k, v, S0, z0
def test_forward_matches_einsum_no_carry():
attn = FractalLinearAttention(d_model=128, n_heads=2, d_head=64, n_levels=2)
q, k, v, _, _ = _make(attn, G=6)
y_ref = attn._linear_attention_causal_einsum(q, k, v)
y_new = attn._linear_attention_causal_cumsum(q, k, v)
assert torch.allclose(y_ref, y_new, atol=ATOL, rtol=RTOL), \
f"max diff {(y_ref - y_new).abs().max().item():.3e}"
def test_forward_matches_einsum_with_carry():
attn = FractalLinearAttention(d_model=128, n_heads=2, d_head=64, n_levels=2)
q, k, v, S0, z0 = _make(attn, G=6)
y_ref, (Sr, zr) = attn._linear_attention_causal_einsum(q, k, v, carry=(S0, z0))
y_new, (Sn, zn) = attn._linear_attention_causal_cumsum(q, k, v, carry=(S0, z0))
assert torch.allclose(y_ref, y_new, atol=ATOL, rtol=RTOL)
assert torch.allclose(Sr, Sn, atol=ATOL * 10, rtol=RTOL) # S sums are larger
assert torch.allclose(zr, zn, atol=ATOL, rtol=RTOL)
def test_gradients_match_with_carry():
attn = FractalLinearAttention(d_model=128, n_heads=2, d_head=64, n_levels=2)
q, k, v, S0, z0 = _make(attn, G=4)
grads = {}
for name, fn in (("einsum", attn._linear_attention_causal_einsum),
("cumsum", attn._linear_attention_causal_cumsum)):
qg = q.clone().requires_grad_(True)
kg = k.clone().requires_grad_(True)
vg = v.clone().requires_grad_(True)
y, _ = fn(qg, kg, vg, carry=(S0, z0))
# weighted loss so every element matters (avoid symmetric cancellation)
w = torch.linspace(0.1, 1.0, y.numel()).view(y.shape)
(y * w).sum().backward()
grads[name] = (qg.grad.clone(), kg.grad.clone(), vg.grad.clone())
for i, dim in enumerate(("q", "k", "v")):
gr, gn = grads["einsum"][i], grads["cumsum"][i]
assert torch.allclose(gr, gn, atol=1e-3, rtol=1e-3), \
f"grad[{dim}] max diff {(gr - gn).abs().max().item():.3e}"
def test_both_match_looped_reference():
"""The original scalar loop is the ultimate ground truth."""
torch.manual_seed(3)
attn = FractalLinearAttention(d_model=64, n_heads=1, d_head=64, n_levels=1)
G, L, D = 2, 16, 64
q = torch.rand(G, L, D) + 0.5
k = torch.rand(G, L, D) + 0.5
v = torch.randn(G, L, D) * 0.5
y_loop = attn._linear_attention_causal_one_head(q, k, v)
y_einsum = attn._linear_attention_causal_einsum(q, k, v)
y_cumsum = attn._linear_attention_causal_cumsum(q, k, v)
assert torch.allclose(y_loop, y_einsum, atol=1e-4, rtol=1e-4)
assert torch.allclose(y_loop, y_cumsum, atol=1e-4, rtol=1e-4)
def test_chunked_matches_einsum_no_carry():
attn = FractalLinearAttention(d_model=128, n_heads=2, d_head=64, n_levels=2)
for L in (32, 128): # includes multi-block case
q, k, v, _, _ = _make(attn, G=6, L=L, seed=L)
y_ref = attn._linear_attention_causal_einsum(q, k, v)
y_ch = attn._linear_attention_causal_chunked(q, k, v, block=16)
assert torch.allclose(y_ref, y_ch, atol=ATOL * 2, rtol=RTOL), \
f"L={L}: max diff {(y_ref - y_ch).abs().max().item():.3e}"
def test_chunked_matches_einsum_with_carry():
attn = FractalLinearAttention(d_model=128, n_heads=2, d_head=64, n_levels=2)
q, k, v, S0, z0 = _make(attn, G=6, L=64)
y_ref, (Sr, zr) = attn._linear_attention_causal_einsum(q, k, v, carry=(S0, z0))
y_ch, (Sc, zc) = attn._linear_attention_causal_chunked(q, k, v, carry=(S0, z0), block=16)
assert torch.allclose(y_ref, y_ch, atol=ATOL * 4, rtol=RTOL), \
f"max diff {(y_ref - y_ch).abs().max().item():.3e}"
assert torch.allclose(Sr, Sc, atol=1e-2, rtol=1e-3)
assert torch.allclose(zr, zc, atol=ATOL, rtol=RTOL)
def test_chunked_gradients_match():
attn = FractalLinearAttention(d_model=128, n_heads=2, d_head=64, n_levels=2)
q, k, v, S0, z0 = _make(attn, G=4, L=32)
grads = {}
for name, fn in (("einsum", attn._linear_attention_causal_einsum),
("chunked", lambda *a, **kw: attn._linear_attention_causal_chunked(*a, block=16, **kw))):
qg = q.clone().requires_grad_(True)
kg = k.clone().requires_grad_(True)
vg = v.clone().requires_grad_(True)
y, _ = fn(qg, kg, vg, carry=(S0, z0))
w = torch.linspace(0.1, 1.0, y.numel()).view(y.shape)
(y * w).sum().backward()
grads[name] = (qg.grad.clone(), kg.grad.clone(), vg.grad.clone())
for i, dim in enumerate(("q", "k", "v")):
gr, gn = grads["einsum"][i], grads["chunked"][i]
assert torch.allclose(gr, gn, atol=5e-3, rtol=5e-3), \
f"grad[{dim}] max diff {(gr - gn).abs().max().item():.3e}"
def test_chunked_ragged_length_falls_back():
"""Non-multiple lengths must still be exact via the cumsum fallback."""
attn = FractalLinearAttention(d_model=128, n_heads=2, d_head=64, n_levels=2)
q, k, v, _, _ = _make(attn, G=6, L=50)
y_ref = attn._linear_attention_causal_einsum(q, k, v)
y_ch = attn._linear_attention_causal_chunked(q, k, v, block=16)
assert torch.allclose(y_ref, y_ch, atol=ATOL * 2, rtol=RTOL)
def test_impl_switch_affects_dispatch():
from fractus.nn import attention as A
attn = FractalLinearAttention(d_model=128, n_heads=2, d_head=64, n_levels=2)
q, k, v, S0, z0 = _make(attn, G=6, L=64)
prev = A._ACTIVE_IMPL
try:
A.set_attention_impl("chunked")
y_ch, _ = attn._linear_attention_causal_vectorized(q, k, v, carry=(S0, z0))
A.set_attention_impl("cumsum")
y_cs, _ = attn._linear_attention_causal_vectorized(q, k, v, carry=(S0, z0))
y_ref, _ = attn._linear_attention_causal_einsum(q, k, v, carry=(S0, z0))
assert torch.allclose(y_ref, y_ch, atol=ATOL * 4, rtol=RTOL)
assert torch.allclose(y_ref, y_cs, atol=ATOL, rtol=RTOL)
finally:
A.set_attention_impl(prev)
# ---------------------------------------------------------------------------
# Chunked cross-entropy equivalence (opt 3)
# ---------------------------------------------------------------------------
def _ce_case(N=1024, d=128, V=50257, seed=7):
torch.manual_seed(seed)
h = torch.randn(N, d) * 0.5
w = torch.randn(V, d) * 0.05
t = torch.randint(0, V, (N,))
return h, w, t
def test_ce_loss_matches_dense():
from fractus.nn.ce import chunked_cross_entropy
import torch.nn.functional as F
h, w, t = _ce_case()
logits = F.linear(h, w)
ref = F.cross_entropy(logits.float(), t)
got = chunked_cross_entropy(h, w, t, ce_chunk=333) # ragged chunks
assert torch.allclose(ref, got, atol=1e-4, rtol=1e-5), \
f"ref={ref.item():.6f} got={got.item():.6f}"
def test_ce_gradients_match_dense():
from fractus.nn.ce import chunked_cross_entropy
import torch.nn.functional as F
h, w, t = _ce_case(N=512, d=64, V=2000)
hg = h.clone().requires_grad_(True)
wg = w.clone().requires_grad_(True)
loss_dense = F.cross_entropy(F.linear(hg, wg).float(), t)
loss_dense.backward()
hc = h.clone().requires_grad_(True)
wc = w.clone().requires_grad_(True)
loss_chunk = chunked_cross_entropy(hc, wc, t, ce_chunk=128)
loss_chunk.backward()
assert torch.allclose(loss_dense, loss_chunk, atol=1e-5, rtol=1e-5)
assert torch.allclose(hg.grad, hc.grad, atol=1e-4, rtol=1e-3), \
f"h grad max diff {(hg.grad - hc.grad).abs().max().item():.3e}"
# weight grad: rows never touched by targets are zero in BOTH paths;
# compare only touched rows to avoid dense-vs-chunk zero-row noise
touched = torch.zeros(V := 2000, dtype=torch.bool).scatter_(0, t, True)
assert torch.allclose(wg.grad[touched], wc.grad[touched], atol=1e-5, rtol=1e-3)
def test_engine_tick_chunk_train_ce_matches_train():
"""End-to-end: tick_chunk_train_ce loss == CE(tick_chunk_train logits).
Weights must be CLONED via load_state_dict — two consecutive constructions
under one seed draw DIFFERENT parameters (the original bug behind
dense=63.7 vs chunked=61.2). With identical weights and the same
production attention kernel the forward is deterministic (zero-init
thought state), so any residual gap is CE reduction order only (~1e-6).
"""
import torch.nn.functional as F
from fractus.continuous_engine import ContinuousThoughtEngine
cfg = dict(vocab_size=1000, d_model=64, n_heads=1, d_head=64,
n_levels=2, n_oscillators=8, coupling_rank=4,
n_experts=4, top_k=2, expert_d_ff=64, siren_rank=16,
n_layers=2)
torch.manual_seed(5)
eng_d = ContinuousThoughtEngine(**cfg)
eng_c = ContinuousThoughtEngine(**cfg)
eng_c.load_state_dict(eng_d.state_dict())
torch.manual_seed(6)
toks = torch.randint(0, 1000, (2, 33))
chunk, target = toks[:, :-1], toks[:, 1:]
losses = []
for eng, mode in ((eng_d, "dense"), (eng_c, "chunked")):
eng.reset_thought(batch_size=2)
if mode == "dense":
logits, lb = eng.tick_chunk_train(chunk)
ce = F.cross_entropy(logits.reshape(-1, logits.size(-1)),
target.reshape(-1))
else:
ce, lb = eng.tick_chunk_train_ce(chunk, target, ce_chunk=17)
losses.append((ce.detach(), lb.detach()))
assert torch.allclose(losses[0][0], losses[1][0], atol=1e-4, rtol=1e-4), \
f"dense={losses[0][0].item():.5f} chunked={losses[1][0].item():.5f}"
assert torch.allclose(losses[0][1], losses[1][1])
def test_dispatcher_uses_production_path():
"""The engine calls _linear_attention_causal_vectorized — it must route to
the cumsum implementation and stay equivalent to the reference."""
attn = FractalLinearAttention(d_model=128, n_heads=2, d_head=64, n_levels=2)
q, k, v, S0, z0 = _make(attn, G=6)
y_ref, st_ref = attn._linear_attention_causal_einsum(q, k, v, carry=(S0, z0))
y_disp, st_disp = attn._linear_attention_causal_vectorized(q, k, v, carry=(S0, z0))
assert torch.allclose(y_ref, y_disp, atol=ATOL, rtol=RTOL)
assert torch.allclose(st_ref[0], st_disp[0], atol=ATOL * 10, rtol=RTOL)
def test_engine_end_to_end_chunk_equivalence():
"""Two-level proof around CTEBlock.tick_chunk_core:
Level 1 (strict): the attention stage — the code we actually replaced —
must match the einsum reference tightly on the exact flattened shapes
the engine feeds it, carry included.
Level 2 (documented): downstream of attention, the von Mises gate + topk
are DISCRETE. A float32-rounding difference near a gate tie can flip
which expert a token routes to, producing a localized output jump far
above rounding scale. This is not an implementation error — it is the
measure-zero boundary behavior of argmax under two equally-valid
summation orders. We therefore assert the mismatch is RARE and bounded,
not zero, and print it for the log.
"""
from fractus.continuous_engine import CTEBlock
cfg = dict(d_model=128, n_heads=2, d_head=64, n_levels=2,
n_oscillators=8, coupling_rank=4, n_experts=8,
top_k=2, expert_d_ff=128, siren_rank=32)
torch.manual_seed(11)
blk_r = CTEBlock(**cfg)
blk_n = CTEBlock(**cfg)
# same trap as the engine CE test: one shared seed does NOT give two
# constructions identical weights — clone explicitly.
blk_n.load_state_dict(blk_r.state_dict())
# ---- Level 1: attention stage, engine shapes -------------------------
attn = blk_r.attn
nH, dH, nL = attn.n_heads, attn.d_head, attn.n_levels
B, C = 3, 64
torch.manual_seed(12)
h_normed = torch.randn(B, C, cfg["d_model"])
q_all = h_normed.view(B, C, nH, dH)
k_all = h_normed.view(B, C, nH, dH)
v_all = h_normed.view(B, C, nH, dH)
offsets = attn.level_offsets
q_feat = elu_plus_one(q_all.unsqueeze(1) + offsets.view(nL, 1, 1, 1), alpha=1.0)
k_feat = elu_plus_one(k_all.unsqueeze(1) + offsets.view(nL, 1, 1, 1), alpha=1.0)
v_lev = v_all.unsqueeze(1).expand(B, nL, C, nH, dH)
qf = q_feat.permute(0, 1, 3, 2, 4).reshape(B * nL * nH, C, dH)
kf = k_feat.permute(0, 1, 3, 2, 4).reshape(B * nL * nH, C, dH)
vf = v_lev.permute(0, 1, 3, 2, 4).reshape(B * nL * nH, C, dH)
S0 = torch.rand(B * nL * nH, dH, dH) * 0.01
z0 = torch.rand(B * nL * nH, dH) + 0.5
y_ref, (Sr, zr) = attn._linear_attention_causal_einsum(qf, kf, vf, carry=(S0, z0))
y_new, (Sn, zn) = attn._linear_attention_causal_vectorized(qf, kf, vf, carry=(S0, z0))
assert torch.allclose(y_ref, y_new, atol=ATOL * 4, rtol=RTOL), \
f"attention y max diff {(y_ref - y_new).abs().max().item():.3e}"
assert torch.allclose(Sr, Sn, atol=1e-2, rtol=1e-3)
assert torch.allclose(zr, zn, atol=ATOL * 10, rtol=RTOL)
# ---- Level 2: full block with flip tolerance -------------------------
torch.manual_seed(13)
h_in = torch.randn(B, C, cfg["d_model"])
results = []
for blk, mode in ((blk_r, "einsum"), (blk_n, "production")):
orig = blk.attn._linear_attention_causal_vectorized
if mode == "einsum":
blk.attn._linear_attention_causal_vectorized = \
lambda q, k, v, carry=None, _ref=blk.attn._linear_attention_causal_einsum: \
_ref(q, k, v, carry=carry)
try:
out, lb = blk.tick_chunk_core(h_in.clone())
results.append((out.detach(), lb.detach()))
finally:
blk.attn._linear_attention_causal_vectorized = orig
out_r, lb_r = results[0]
out_n, lb_n = results[1]
diff = (out_r - out_n).abs()
frac_mismatch = (diff > 1e-3).float().mean().item()
# routing flips, when they occur, touch a tiny fraction of positions;
# everything else matches at rounding scale.
assert diff.median() < 1e-4, f"median diff {diff.median().item():.3e} too large"
assert frac_mismatch < 0.05, \
f"{frac_mismatch:.1%} of elements differ >1e-3 — too many for tie flips"
assert abs(float(lb_r) - float(lb_n)) < 0.5
print(f"\n[end-to-end] median diff {diff.median().item():.2e}, "
f"max {diff.max().item():.2e}, mismatch>1e-3: {frac_mismatch:.2%}")
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v"]))