thefinalboss commited on
Commit
8d9996a
·
verified ·
1 Parent(s): 938dd12

opt: cumsum/chunked attention kernels, memory-flat CE, block checkpointing, v2 trainer (proven equivalent, 46 tests)

Browse files
Files changed (1) hide show
  1. tests/test_attention_equivalence.py +349 -0
tests/test_attention_equivalence.py ADDED
@@ -0,0 +1,349 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Equivalence proofs: cumsum attention vs einsum reference.
2
+
3
+ Open-heart rule: any production-path change must be PROVEN mathematically
4
+ equivalent to the reference before it may wrap a live checkpoint.
5
+
6
+ Covered here (opt 1, 2026-08-22):
7
+ - forward outputs match (with and without state carry)
8
+ - gradients wrt q, k, v match
9
+ - carried final states (S_final, z_final) match
10
+ - both match the scalar looped reference (_linear_attention_causal_one_head)
11
+ """
12
+ import sys
13
+ from pathlib import Path
14
+
15
+ import pytest
16
+ import torch
17
+
18
+ sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
19
+
20
+ from fractus.nn.attention import FractalLinearAttention
21
+ from fractus.nn.stats import elu_plus_one
22
+
23
+
24
+ ATOL = 1e-4 # float32 accumulation-order tolerance
25
+ RTOL = 1e-4
26
+
27
+
28
+ def _make(attn, G, L=32, D=64, seed=0):
29
+ torch.manual_seed(seed)
30
+ q = torch.rand(G, L, D) + 0.5 # positive like elu+1 features
31
+ k = torch.rand(G, L, D) + 0.5
32
+ v = torch.randn(G, L, D) * 0.5
33
+ S0 = torch.randn(G, D, D) * 0.05
34
+ z0 = torch.rand(G, D) + 0.5 # z from positive features
35
+ return q, k, v, S0, z0
36
+
37
+
38
+ def test_forward_matches_einsum_no_carry():
39
+ attn = FractalLinearAttention(d_model=128, n_heads=2, d_head=64, n_levels=2)
40
+ q, k, v, _, _ = _make(attn, G=6)
41
+ y_ref = attn._linear_attention_causal_einsum(q, k, v)
42
+ y_new = attn._linear_attention_causal_cumsum(q, k, v)
43
+ assert torch.allclose(y_ref, y_new, atol=ATOL, rtol=RTOL), \
44
+ f"max diff {(y_ref - y_new).abs().max().item():.3e}"
45
+
46
+
47
+ def test_forward_matches_einsum_with_carry():
48
+ attn = FractalLinearAttention(d_model=128, n_heads=2, d_head=64, n_levels=2)
49
+ q, k, v, S0, z0 = _make(attn, G=6)
50
+ y_ref, (Sr, zr) = attn._linear_attention_causal_einsum(q, k, v, carry=(S0, z0))
51
+ y_new, (Sn, zn) = attn._linear_attention_causal_cumsum(q, k, v, carry=(S0, z0))
52
+ assert torch.allclose(y_ref, y_new, atol=ATOL, rtol=RTOL)
53
+ assert torch.allclose(Sr, Sn, atol=ATOL * 10, rtol=RTOL) # S sums are larger
54
+ assert torch.allclose(zr, zn, atol=ATOL, rtol=RTOL)
55
+
56
+
57
+ def test_gradients_match_with_carry():
58
+ attn = FractalLinearAttention(d_model=128, n_heads=2, d_head=64, n_levels=2)
59
+ q, k, v, S0, z0 = _make(attn, G=4)
60
+
61
+ grads = {}
62
+ for name, fn in (("einsum", attn._linear_attention_causal_einsum),
63
+ ("cumsum", attn._linear_attention_causal_cumsum)):
64
+ qg = q.clone().requires_grad_(True)
65
+ kg = k.clone().requires_grad_(True)
66
+ vg = v.clone().requires_grad_(True)
67
+ y, _ = fn(qg, kg, vg, carry=(S0, z0))
68
+ # weighted loss so every element matters (avoid symmetric cancellation)
69
+ w = torch.linspace(0.1, 1.0, y.numel()).view(y.shape)
70
+ (y * w).sum().backward()
71
+ grads[name] = (qg.grad.clone(), kg.grad.clone(), vg.grad.clone())
72
+
73
+ for i, dim in enumerate(("q", "k", "v")):
74
+ gr, gn = grads["einsum"][i], grads["cumsum"][i]
75
+ assert torch.allclose(gr, gn, atol=1e-3, rtol=1e-3), \
76
+ f"grad[{dim}] max diff {(gr - gn).abs().max().item():.3e}"
77
+
78
+
79
+ def test_both_match_looped_reference():
80
+ """The original scalar loop is the ultimate ground truth."""
81
+ torch.manual_seed(3)
82
+ attn = FractalLinearAttention(d_model=64, n_heads=1, d_head=64, n_levels=1)
83
+ G, L, D = 2, 16, 64
84
+ q = torch.rand(G, L, D) + 0.5
85
+ k = torch.rand(G, L, D) + 0.5
86
+ v = torch.randn(G, L, D) * 0.5
87
+
88
+ y_loop = attn._linear_attention_causal_one_head(q, k, v)
89
+ y_einsum = attn._linear_attention_causal_einsum(q, k, v)
90
+ y_cumsum = attn._linear_attention_causal_cumsum(q, k, v)
91
+
92
+ assert torch.allclose(y_loop, y_einsum, atol=1e-4, rtol=1e-4)
93
+ assert torch.allclose(y_loop, y_cumsum, atol=1e-4, rtol=1e-4)
94
+
95
+
96
+ def test_chunked_matches_einsum_no_carry():
97
+ attn = FractalLinearAttention(d_model=128, n_heads=2, d_head=64, n_levels=2)
98
+ for L in (32, 128): # includes multi-block case
99
+ q, k, v, _, _ = _make(attn, G=6, L=L, seed=L)
100
+ y_ref = attn._linear_attention_causal_einsum(q, k, v)
101
+ y_ch = attn._linear_attention_causal_chunked(q, k, v, block=16)
102
+ assert torch.allclose(y_ref, y_ch, atol=ATOL * 2, rtol=RTOL), \
103
+ f"L={L}: max diff {(y_ref - y_ch).abs().max().item():.3e}"
104
+
105
+
106
+ def test_chunked_matches_einsum_with_carry():
107
+ attn = FractalLinearAttention(d_model=128, n_heads=2, d_head=64, n_levels=2)
108
+ q, k, v, S0, z0 = _make(attn, G=6, L=64)
109
+ y_ref, (Sr, zr) = attn._linear_attention_causal_einsum(q, k, v, carry=(S0, z0))
110
+ y_ch, (Sc, zc) = attn._linear_attention_causal_chunked(q, k, v, carry=(S0, z0), block=16)
111
+ assert torch.allclose(y_ref, y_ch, atol=ATOL * 4, rtol=RTOL), \
112
+ f"max diff {(y_ref - y_ch).abs().max().item():.3e}"
113
+ assert torch.allclose(Sr, Sc, atol=1e-2, rtol=1e-3)
114
+ assert torch.allclose(zr, zc, atol=ATOL, rtol=RTOL)
115
+
116
+
117
+ def test_chunked_gradients_match():
118
+ attn = FractalLinearAttention(d_model=128, n_heads=2, d_head=64, n_levels=2)
119
+ q, k, v, S0, z0 = _make(attn, G=4, L=32)
120
+
121
+ grads = {}
122
+ for name, fn in (("einsum", attn._linear_attention_causal_einsum),
123
+ ("chunked", lambda *a, **kw: attn._linear_attention_causal_chunked(*a, block=16, **kw))):
124
+ qg = q.clone().requires_grad_(True)
125
+ kg = k.clone().requires_grad_(True)
126
+ vg = v.clone().requires_grad_(True)
127
+ y, _ = fn(qg, kg, vg, carry=(S0, z0))
128
+ w = torch.linspace(0.1, 1.0, y.numel()).view(y.shape)
129
+ (y * w).sum().backward()
130
+ grads[name] = (qg.grad.clone(), kg.grad.clone(), vg.grad.clone())
131
+
132
+ for i, dim in enumerate(("q", "k", "v")):
133
+ gr, gn = grads["einsum"][i], grads["chunked"][i]
134
+ assert torch.allclose(gr, gn, atol=5e-3, rtol=5e-3), \
135
+ f"grad[{dim}] max diff {(gr - gn).abs().max().item():.3e}"
136
+
137
+
138
+ def test_chunked_ragged_length_falls_back():
139
+ """Non-multiple lengths must still be exact via the cumsum fallback."""
140
+ attn = FractalLinearAttention(d_model=128, n_heads=2, d_head=64, n_levels=2)
141
+ q, k, v, _, _ = _make(attn, G=6, L=50)
142
+ y_ref = attn._linear_attention_causal_einsum(q, k, v)
143
+ y_ch = attn._linear_attention_causal_chunked(q, k, v, block=16)
144
+ assert torch.allclose(y_ref, y_ch, atol=ATOL * 2, rtol=RTOL)
145
+
146
+
147
+ def test_impl_switch_affects_dispatch():
148
+ from fractus.nn import attention as A
149
+ attn = FractalLinearAttention(d_model=128, n_heads=2, d_head=64, n_levels=2)
150
+ q, k, v, S0, z0 = _make(attn, G=6, L=64)
151
+
152
+ prev = A._ACTIVE_IMPL
153
+ try:
154
+ A.set_attention_impl("chunked")
155
+ y_ch, _ = attn._linear_attention_causal_vectorized(q, k, v, carry=(S0, z0))
156
+ A.set_attention_impl("cumsum")
157
+ y_cs, _ = attn._linear_attention_causal_vectorized(q, k, v, carry=(S0, z0))
158
+ y_ref, _ = attn._linear_attention_causal_einsum(q, k, v, carry=(S0, z0))
159
+ assert torch.allclose(y_ref, y_ch, atol=ATOL * 4, rtol=RTOL)
160
+ assert torch.allclose(y_ref, y_cs, atol=ATOL, rtol=RTOL)
161
+ finally:
162
+ A.set_attention_impl(prev)
163
+
164
+
165
+ # ---------------------------------------------------------------------------
166
+ # Chunked cross-entropy equivalence (opt 3)
167
+ # ---------------------------------------------------------------------------
168
+
169
+ def _ce_case(N=1024, d=128, V=50257, seed=7):
170
+ torch.manual_seed(seed)
171
+ h = torch.randn(N, d) * 0.5
172
+ w = torch.randn(V, d) * 0.05
173
+ t = torch.randint(0, V, (N,))
174
+ return h, w, t
175
+
176
+
177
+ def test_ce_loss_matches_dense():
178
+ from fractus.nn.ce import chunked_cross_entropy
179
+ import torch.nn.functional as F
180
+
181
+ h, w, t = _ce_case()
182
+ logits = F.linear(h, w)
183
+ ref = F.cross_entropy(logits.float(), t)
184
+ got = chunked_cross_entropy(h, w, t, ce_chunk=333) # ragged chunks
185
+ assert torch.allclose(ref, got, atol=1e-4, rtol=1e-5), \
186
+ f"ref={ref.item():.6f} got={got.item():.6f}"
187
+
188
+
189
+ def test_ce_gradients_match_dense():
190
+ from fractus.nn.ce import chunked_cross_entropy
191
+ import torch.nn.functional as F
192
+
193
+ h, w, t = _ce_case(N=512, d=64, V=2000)
194
+
195
+ hg = h.clone().requires_grad_(True)
196
+ wg = w.clone().requires_grad_(True)
197
+ loss_dense = F.cross_entropy(F.linear(hg, wg).float(), t)
198
+ loss_dense.backward()
199
+
200
+ hc = h.clone().requires_grad_(True)
201
+ wc = w.clone().requires_grad_(True)
202
+ loss_chunk = chunked_cross_entropy(hc, wc, t, ce_chunk=128)
203
+ loss_chunk.backward()
204
+
205
+ assert torch.allclose(loss_dense, loss_chunk, atol=1e-5, rtol=1e-5)
206
+ assert torch.allclose(hg.grad, hc.grad, atol=1e-4, rtol=1e-3), \
207
+ f"h grad max diff {(hg.grad - hc.grad).abs().max().item():.3e}"
208
+ # weight grad: rows never touched by targets are zero in BOTH paths;
209
+ # compare only touched rows to avoid dense-vs-chunk zero-row noise
210
+ touched = torch.zeros(V := 2000, dtype=torch.bool).scatter_(0, t, True)
211
+ assert torch.allclose(wg.grad[touched], wc.grad[touched], atol=1e-5, rtol=1e-3)
212
+
213
+
214
+ def test_engine_tick_chunk_train_ce_matches_train():
215
+ """End-to-end: tick_chunk_train_ce loss == CE(tick_chunk_train logits).
216
+
217
+ Weights must be CLONED via load_state_dict — two consecutive constructions
218
+ under one seed draw DIFFERENT parameters (the original bug behind
219
+ dense=63.7 vs chunked=61.2). With identical weights and the same
220
+ production attention kernel the forward is deterministic (zero-init
221
+ thought state), so any residual gap is CE reduction order only (~1e-6).
222
+ """
223
+ import torch.nn.functional as F
224
+ from fractus.continuous_engine import ContinuousThoughtEngine
225
+
226
+ cfg = dict(vocab_size=1000, d_model=64, n_heads=1, d_head=64,
227
+ n_levels=2, n_oscillators=8, coupling_rank=4,
228
+ n_experts=4, top_k=2, expert_d_ff=64, siren_rank=16,
229
+ n_layers=2)
230
+ torch.manual_seed(5)
231
+ eng_d = ContinuousThoughtEngine(**cfg)
232
+ eng_c = ContinuousThoughtEngine(**cfg)
233
+ eng_c.load_state_dict(eng_d.state_dict())
234
+ torch.manual_seed(6)
235
+ toks = torch.randint(0, 1000, (2, 33))
236
+ chunk, target = toks[:, :-1], toks[:, 1:]
237
+
238
+ losses = []
239
+ for eng, mode in ((eng_d, "dense"), (eng_c, "chunked")):
240
+ eng.reset_thought(batch_size=2)
241
+ if mode == "dense":
242
+ logits, lb = eng.tick_chunk_train(chunk)
243
+ ce = F.cross_entropy(logits.reshape(-1, logits.size(-1)),
244
+ target.reshape(-1))
245
+ else:
246
+ ce, lb = eng.tick_chunk_train_ce(chunk, target, ce_chunk=17)
247
+ losses.append((ce.detach(), lb.detach()))
248
+
249
+ assert torch.allclose(losses[0][0], losses[1][0], atol=1e-4, rtol=1e-4), \
250
+ f"dense={losses[0][0].item():.5f} chunked={losses[1][0].item():.5f}"
251
+ assert torch.allclose(losses[0][1], losses[1][1])
252
+
253
+
254
+ def test_dispatcher_uses_production_path():
255
+ """The engine calls _linear_attention_causal_vectorized — it must route to
256
+ the cumsum implementation and stay equivalent to the reference."""
257
+ attn = FractalLinearAttention(d_model=128, n_heads=2, d_head=64, n_levels=2)
258
+ q, k, v, S0, z0 = _make(attn, G=6)
259
+ y_ref, st_ref = attn._linear_attention_causal_einsum(q, k, v, carry=(S0, z0))
260
+ y_disp, st_disp = attn._linear_attention_causal_vectorized(q, k, v, carry=(S0, z0))
261
+ assert torch.allclose(y_ref, y_disp, atol=ATOL, rtol=RTOL)
262
+ assert torch.allclose(st_ref[0], st_disp[0], atol=ATOL * 10, rtol=RTOL)
263
+
264
+
265
+ def test_engine_end_to_end_chunk_equivalence():
266
+ """Two-level proof around CTEBlock.tick_chunk_core:
267
+
268
+ Level 1 (strict): the attention stage — the code we actually replaced —
269
+ must match the einsum reference tightly on the exact flattened shapes
270
+ the engine feeds it, carry included.
271
+
272
+ Level 2 (documented): downstream of attention, the von Mises gate + topk
273
+ are DISCRETE. A float32-rounding difference near a gate tie can flip
274
+ which expert a token routes to, producing a localized output jump far
275
+ above rounding scale. This is not an implementation error — it is the
276
+ measure-zero boundary behavior of argmax under two equally-valid
277
+ summation orders. We therefore assert the mismatch is RARE and bounded,
278
+ not zero, and print it for the log.
279
+ """
280
+ from fractus.continuous_engine import CTEBlock
281
+
282
+ cfg = dict(d_model=128, n_heads=2, d_head=64, n_levels=2,
283
+ n_oscillators=8, coupling_rank=4, n_experts=8,
284
+ top_k=2, expert_d_ff=128, siren_rank=32)
285
+ torch.manual_seed(11)
286
+ blk_r = CTEBlock(**cfg)
287
+ blk_n = CTEBlock(**cfg)
288
+ # same trap as the engine CE test: one shared seed does NOT give two
289
+ # constructions identical weights — clone explicitly.
290
+ blk_n.load_state_dict(blk_r.state_dict())
291
+
292
+ # ---- Level 1: attention stage, engine shapes -------------------------
293
+ attn = blk_r.attn
294
+ nH, dH, nL = attn.n_heads, attn.d_head, attn.n_levels
295
+ B, C = 3, 64
296
+ torch.manual_seed(12)
297
+ h_normed = torch.randn(B, C, cfg["d_model"])
298
+ q_all = h_normed.view(B, C, nH, dH)
299
+ k_all = h_normed.view(B, C, nH, dH)
300
+ v_all = h_normed.view(B, C, nH, dH)
301
+ offsets = attn.level_offsets
302
+ q_feat = elu_plus_one(q_all.unsqueeze(1) + offsets.view(nL, 1, 1, 1), alpha=1.0)
303
+ k_feat = elu_plus_one(k_all.unsqueeze(1) + offsets.view(nL, 1, 1, 1), alpha=1.0)
304
+ v_lev = v_all.unsqueeze(1).expand(B, nL, C, nH, dH)
305
+ qf = q_feat.permute(0, 1, 3, 2, 4).reshape(B * nL * nH, C, dH)
306
+ kf = k_feat.permute(0, 1, 3, 2, 4).reshape(B * nL * nH, C, dH)
307
+ vf = v_lev.permute(0, 1, 3, 2, 4).reshape(B * nL * nH, C, dH)
308
+ S0 = torch.rand(B * nL * nH, dH, dH) * 0.01
309
+ z0 = torch.rand(B * nL * nH, dH) + 0.5
310
+
311
+ y_ref, (Sr, zr) = attn._linear_attention_causal_einsum(qf, kf, vf, carry=(S0, z0))
312
+ y_new, (Sn, zn) = attn._linear_attention_causal_vectorized(qf, kf, vf, carry=(S0, z0))
313
+ assert torch.allclose(y_ref, y_new, atol=ATOL * 4, rtol=RTOL), \
314
+ f"attention y max diff {(y_ref - y_new).abs().max().item():.3e}"
315
+ assert torch.allclose(Sr, Sn, atol=1e-2, rtol=1e-3)
316
+ assert torch.allclose(zr, zn, atol=ATOL * 10, rtol=RTOL)
317
+
318
+ # ---- Level 2: full block with flip tolerance -------------------------
319
+ torch.manual_seed(13)
320
+ h_in = torch.randn(B, C, cfg["d_model"])
321
+ results = []
322
+ for blk, mode in ((blk_r, "einsum"), (blk_n, "production")):
323
+ orig = blk.attn._linear_attention_causal_vectorized
324
+ if mode == "einsum":
325
+ blk.attn._linear_attention_causal_vectorized = \
326
+ lambda q, k, v, carry=None, _ref=blk.attn._linear_attention_causal_einsum: \
327
+ _ref(q, k, v, carry=carry)
328
+ try:
329
+ out, lb = blk.tick_chunk_core(h_in.clone())
330
+ results.append((out.detach(), lb.detach()))
331
+ finally:
332
+ blk.attn._linear_attention_causal_vectorized = orig
333
+
334
+ out_r, lb_r = results[0]
335
+ out_n, lb_n = results[1]
336
+ diff = (out_r - out_n).abs()
337
+ frac_mismatch = (diff > 1e-3).float().mean().item()
338
+ # routing flips, when they occur, touch a tiny fraction of positions;
339
+ # everything else matches at rounding scale.
340
+ assert diff.median() < 1e-4, f"median diff {diff.median().item():.3e} too large"
341
+ assert frac_mismatch < 0.05, \
342
+ f"{frac_mismatch:.1%} of elements differ >1e-3 — too many for tie flips"
343
+ assert abs(float(lb_r) - float(lb_n)) < 0.5
344
+ print(f"\n[end-to-end] median diff {diff.median().item():.2e}, "
345
+ f"max {diff.max().item():.2e}, mismatch>1e-3: {frac_mismatch:.2%}")
346
+
347
+
348
+ if __name__ == "__main__":
349
+ sys.exit(pytest.main([__file__, "-v"]))