opt: cumsum/chunked attention kernels, memory-flat CE, block checkpointing, v2 trainer (proven equivalent, 46 tests)
8d9996a verified Download tests/test_attention_equivalence.py from thefinalboss/fractus-cte: direct link, hf CLI and curl.
- Browser
- Download file 15.3 kB
-
https://huggingface.co/thefinalboss/fractus-cte/resolve/de10ad18c4cded81ef93b7317e3b618b33b5e9d6/tests/test_attention_equivalence.py
- Command line
-
hf download hf://thefinalboss/fractus-cte@de10ad18c4cded81ef93b7317e3b618b33b5e9d6/tests/test_attention_equivalence.py
-
curl -L -o test_attention_equivalence.py https://huggingface.co/thefinalboss/fractus-cte/resolve/de10ad18c4cded81ef93b7317e3b618b33b5e9d6/tests/test_attention_equivalence.py
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"])) | |