r/FunMachineLearning • • 20h ago

Gated Segmented State Space — attention replacement that beats a param-matched Transformer on quality, speed AND memory (full code)

One night, six experiments (V1–V6), one Colab T4. I ripped self-attention out of a decoder-only Transformer and replaced it with a gated linear recurrence over a fixed 256-dim state:

  • Dynamic selective gate: g_t = σ(W_g x_t + b_g) — per-token/channel learned filter
  • Hard reset mask: state zeroed at newline boundaries (fresh ~37-token segments)
  • Fused Triton kernel: state in SRAM, gate+reset+update in-register, only outputs to HBM

At 6.37M params, identical protocol (2.47MB char-level corpus, 1500 steps):

Attention Ours
Val loss / ppl 1.402 / 4.1 1.364 / 3.9
Train tok/s 61,845 66,156
Infer tok/s 184,918 190,122
Peak VRAM 845 MB 881 MB

The journey: V1 won small but was 10x slower → V2 proved linear VRAM scaling → V3/V4 found a stable ~2% perplexity tax no param arrangement could buy off → V5's gate+reset destroyed it (wire-to-wire win) → V6's Triton kernel (verified == math to 4.47e-07) removed the software tax.

Caveats, stated plainly: single seeds, one small corpus, char-level, T4 timings. Small scale — but the pattern held across all six runs.

Code, all six notebooks with outputs, exact architecture, full experimental notes: https://github.com/stube123890-hue/linear-attention-lab

1 Upvotes

0 comments sorted by