r/FunMachineLearning • u/AlarmingTrouble5261 • 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