thefinalboss commited on
Commit
8b9e686
·
verified ·
1 Parent(s): deacc41

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

Browse files
Files changed (1) hide show
  1. fractus/nn/attention.py +155 -7
fractus/nn/attention.py CHANGED
@@ -12,12 +12,30 @@ Math (Katharopoulos 2020, normalized causal form):
12
  """
13
 
14
  import math
 
15
  import torch
16
  import torch.nn as nn
17
 
18
  from .stats import elu_plus_one, stable_softmax
19
 
20
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
21
  def _mandelbrot_offsets(n_levels: int) -> torch.Tensor:
22
  """Offsets ω_level = (φ2)^{-level} for level = 0..n_levels-1.
23
 
@@ -126,17 +144,17 @@ class FractalLinearAttention(nn.Module):
126
  outputs.append(y_t)
127
  return torch.stack(outputs, dim=1) # (B, L, D)
128
 
129
- def _linear_attention_causal_vectorized(
130
  self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
131
  carry: tuple = None,
132
  ) -> torch.Tensor:
133
- """Vectorized version of _linear_attention_causal_one_head.
134
 
135
- Same mathematics, but without a Python loop over L. The trick:
136
- precompute the cumulative sums S_t and z_t via a lower-triangular
137
- convolution, then compute all y_t in parallel.
138
-
139
- Equivalence guaranteed by test_attention_vectorized.py (atol 1e-5).
140
 
141
  L8 STATE-CARRY: if `carry = (S0, z0)` is provided (each (B, D, D) and
142
  (B, D)), the running state is INITIALIZED with (S0, z0) instead of
@@ -186,6 +204,136 @@ class FractalLinearAttention(nn.Module):
186
  return y, (S_final, z_final)
187
  return y # (B, L, D)
188
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
189
  def forward(self, x: torch.Tensor) -> torch.Tensor:
190
  """x: (B, L, d_model) → output (B, L, d_model).
191
 
 
12
  """
13
 
14
  import math
15
+ import os
16
  import torch
17
  import torch.nn as nn
18
 
19
  from .stats import elu_plus_one, stable_softmax
20
 
21
 
22
+ _ACTIVE_IMPL = os.environ.get("FRACTUS_ATTN_IMPL", "cumsum") # cumsum | chunked
23
+
24
+
25
+ def set_attention_impl(name: str) -> None:
26
+ """Select the production linear-attention kernel.
27
+
28
+ 'cumsum' — fastest arithmetic (default; CPU-friendly).
29
+ 'chunked' — memory-flat O(block²+D²) activations (GPU/VRAM-constrained).
30
+ Both are proven equal to the einsum reference by
31
+ tests/test_attention_equivalence.py.
32
+ """
33
+ global _ACTIVE_IMPL
34
+ if name not in ("cumsum", "chunked"):
35
+ raise ValueError(f"unknown attention impl: {name!r}")
36
+ _ACTIVE_IMPL = name
37
+
38
+
39
  def _mandelbrot_offsets(n_levels: int) -> torch.Tensor:
40
  """Offsets ω_level = (φ2)^{-level} for level = 0..n_levels-1.
41
 
 
144
  outputs.append(y_t)
145
  return torch.stack(outputs, dim=1) # (B, L, D)
146
 
147
+ def _linear_attention_causal_einsum(
148
  self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
149
  carry: tuple = None,
150
  ) -> torch.Tensor:
151
+ """REFERENCE implementation (einsum-with-mask form).
152
 
153
+ Kept as numerical ground truth for tests/test_attention_equivalence.py.
154
+ Computes S_t = Σ_{j<=t} outer[j] via a matmul against a triangular
155
+ mask: O(L^2·D^2) multiply-adds. The production path
156
+ (_linear_attention_causal_vectorized → cumsum form) computes the same
157
+ sums with O(L·D^2) adds.
158
 
159
  L8 STATE-CARRY: if `carry = (S0, z0)` is provided (each (B, D, D) and
160
  (B, D)), the running state is INITIALIZED with (S0, z0) instead of
 
204
  return y, (S_final, z_final)
205
  return y # (B, L, D)
206
 
207
+ def _linear_attention_causal_cumsum(
208
+ self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
209
+ carry: tuple = None,
210
+ ) -> torch.Tensor:
211
+ """PRODUCTION implementation: causal linear attention via cumsum.
212
+
213
+ Same mathematics as _linear_attention_causal_einsum — proven equal by
214
+ tests/test_attention_equivalence.py (forward AND gradients, with and
215
+ without carry). The difference is purely computational:
216
+
217
+ S_t = Σ_{j<=t} outer[j] IS a cumulative sum along the L axis.
218
+ The mask-matmul form computed that sum with O(L^2·D^2)
219
+ multiply-adds; cumsum needs O(L·D^2) adds — at L=128 that is
220
+ ~128x less arithmetic on the dominant term. Backward is a reverse
221
+ cumsum instead of another full masked matmul.
222
+
223
+ Numerical note: accumulation order differs from the matmul form, so
224
+ results agree within float32 rounding (~1e-4 relative), far below the
225
+ training signal scale. Checkpoint compatibility is exact: parameter
226
+ shapes and semantics untouched (open-heart surgery rule).
227
+ """
228
+ B, L, D = q.shape
229
+
230
+ outer = torch.einsum("btp,btq->btpq", k, v) # (B, L, D, D)
231
+ S = torch.cumsum(outer, dim=1) # S_t = Σ_{j<=t} k_j⊗v_j
232
+ z = torch.cumsum(k, dim=1) # (B, L, D)
233
+
234
+ if carry is not None:
235
+ S0, z0 = carry
236
+ S = S + S0.unsqueeze(1)
237
+ z = z + z0.unsqueeze(1)
238
+
239
+ num = torch.einsum("btp,btpq->btq", q, S) # (B, L, D)
240
+ denom = (q * z).sum(dim=-1, keepdim=True) # (B, L, 1)
241
+ safe = denom.abs() > 1e-10
242
+ y = torch.where(safe, num / (denom + 1e-20), torch.zeros_like(num))
243
+
244
+ if carry is not None:
245
+ return y, (S[:, -1], z[:, -1])
246
+ return y
247
+
248
+ def _linear_attention_causal_chunked(
249
+ self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
250
+ carry: tuple = None, block: int = 64,
251
+ ) -> torch.Tensor:
252
+ """MEMORY-FLAT implementation: blockwise causal linear attention.
253
+
254
+ Same inclusive-causal mathematics as the einsum/cumsum forms (proven by
255
+ tests/test_attention_equivalence.py), but S_t is never materialized
256
+ for every t. The sequence is processed in blocks of `block` positions:
257
+
258
+ intra-block : y_t += Σ_{j<=t, same block} (φq_t·φk_j) v_j
259
+ via an (block×block) masked score matrix
260
+ inter-block : y_t += φq_t · S_run / (φq_t·z_run contribution)
261
+ state update : S_run += Σ_block φk⊗v ; z_run += Σ_block φk
262
+
263
+ Peak activation per group drops from O(L·D²) (the outer/S tensors of
264
+ the cumsum form) to O(block² + D²). At L=128, D=64, block=64 that is
265
+ a ~32x reduction of the dominant activation — the lever that should
266
+ re-enable torch.compile and large batches on VRAM-constrained pods.
267
+
268
+ FLOP note: intra work is O(L·block·D) instead of O(L·D²); with
269
+ block=64 < D=64 it is on par. Choose this impl when MEMORY is the
270
+ binding constraint (GPU), cumsum when raw speed is (CPU).
271
+
272
+ `carry` semantics identical: (S0, z0) added to every position's sums;
273
+ returned state includes it.
274
+ """
275
+ B, L, D = q.shape
276
+ if L % block != 0:
277
+ # fall back to the exact cumsum path for ragged lengths
278
+ return self._linear_attention_causal_cumsum(q, k, v, carry=carry)
279
+
280
+ dev, dt = q.device, q.dtype
281
+ if carry is not None:
282
+ S_run, z_run = carry[0].clone(), carry[1].clone()
283
+ else:
284
+ S_run = torch.zeros(B, D, D, dtype=dt, device=dev)
285
+ z_run = torch.zeros(B, D, dtype=dt, device=dev)
286
+
287
+ causal = torch.tril(
288
+ torch.ones(block, block, dtype=torch.bool, device=dev))
289
+
290
+ outs = []
291
+ for a in range(0, L, block):
292
+ qb = q[:, a : a + block] # (G, b, D)
293
+ kb = k[:, a : a + block]
294
+ vb = v[:, a : a + block]
295
+
296
+ scores = torch.bmm(qb, kb.transpose(1, 2)) # (G, b, b) φq·φk
297
+
298
+ num_intra = torch.bmm(scores.masked_fill(~causal, 0.0), vb)
299
+ den_intra = scores.masked_fill(~causal, 0.0).sum(dim=-1) # (G, b)
300
+
301
+ num_inter = torch.bmm(qb, S_run) # (G, b, D)
302
+ den_inter = torch.bmm(qb, z_run.unsqueeze(-1)).squeeze(-1)
303
+
304
+ num = num_intra + num_inter
305
+ den = den_intra + den_inter
306
+ safe = den.abs() > 1e-10
307
+ outs.append(torch.where(
308
+ safe.unsqueeze(-1), num / (den.unsqueeze(-1) + 1e-20),
309
+ torch.zeros_like(num)))
310
+
311
+ # inclusive update AFTER readout: current token joins state for
312
+ # the NEXT query positions only (matches S_t including token t).
313
+ S_run = S_run + torch.einsum("btp,btq->bpq", kb, vb)
314
+ z_run = z_run + kb.sum(dim=1)
315
+
316
+ y = torch.cat(outs, dim=1)
317
+ if carry is not None:
318
+ return y, (S_run, z_run)
319
+ return y
320
+
321
+ def _linear_attention_causal_vectorized(
322
+ self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
323
+ carry: tuple = None,
324
+ ) -> torch.Tensor:
325
+ """Dispatches to the active production kernel (default: cumsum).
326
+
327
+ Signature and return contract are unchanged from the original
328
+ vectorized implementation, so all call sites (CTEBlock.tick_chunk_core,
329
+ FractalLinearAttention.forward) work without modification. Switch
330
+ kernels with set_attention_impl('chunked'|'cumsum') or the
331
+ FRACTUS_ATTN_IMPL environment variable.
332
+ """
333
+ if _ACTIVE_IMPL == "chunked":
334
+ return self._linear_attention_causal_chunked(q, k, v, carry=carry)
335
+ return self._linear_attention_causal_cumsum(q, k, v, carry=carry)
336
+
337
  def forward(self, x: torch.Tensor) -> torch.Tensor:
338
  """x: (B, L, d_model) → output (B, L, d_model).
339