"""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"]))