Fast Attention and Long Context

Linear attention

Linear attention drops softmax so the maths can be rearranged, turning the transformer into a recurrent network with a fixed-size memory and constant cost per token.

On this page 10
  1. The short answer
  2. The shopkeeper's two systems
  3. What is given up, stated plainly
  4. Why it is faster
  5. How the rearrangement works
  6. The interesting equivalence
  7. What actually gets shipped
  8. Where you have already seen this
  9. Remember this
  10. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

The short answer

Linear attention keeps one fixed-size summary of everything read so far, instead of keeping every word and searching through them.

The shopkeeper's two systems

A shopkeeper can run the day in two ways.

Keep every bill in a box. At closing time you can answer any question — what did the man in the blue shirt buy at eleven? — by finding that bill. The box grows all day, and searching takes longer as it fills.

Or keep a running total on one slip. Total sales, total items, cash in hand. One slip, same size at closing as at opening, updated in a second per customer.

Ordinary attention is the box of bills. Linear attention is the slip.

What is given up, stated plainly

The slip cannot answer "what did the man in the blue shirt buy". That information went into the total and cannot be pulled back out.

This is not a detail. It is the central trade of this whole family of methods. Every honest paper about them asks how much fits on the slip before it stops being enough.

Why it is faster

With the box, writing each new word means reading every bill again. Work grows with the square of the day's length.

With the slip, each word costs the same as the last. Nothing grows. A model can read a million words at a steady speed and a steady memory.

   ordinary attention           linear attention
   ------------------           ----------------
   store every word             store one summary
   cost per word: grows         cost per word: constant
   memory: grows                memory: constant
   can recall exactly           can recall approximately

How the rearrangement works

Ordinary attention compares every question with every earlier word, then combines. The comparison step is what costs so much.

Take away one step in the middle, a percentage calculation called softmax. The remaining arithmetic can then be regrouped. Instead of comparing questions to words, you fold all the words into one small table first. Each question then reads that table.

The regrouping is ordinary school algebra. Multiplying three things left-to-right or right-to-left gives the same answer, and one order is far cheaper here.

Taking away that middle step is the price. Softmax is what makes attention sharp enough to pick out one word.

The interesting equivalence

Each word costs the same, and one fixed summary is carried forward. That is a recurrent network: a model reading one item at a time, carrying a memory forward.

That is the family transformers replaced. Linear attention arrives back there from the other direction. A well-known paper on this is titled, accurately, Transformers are RNNs.

What actually gets shipped

Almost nobody uses linear attention alone. Quality falls short on tasks that need exact recall.

The pattern that works is a mixture. Three cheap linear layers for every one full-attention layer. The cheap layers carry the bulk; the occasional full layer restores exact lookup.

At least one large lab tried a fast attention variant, then removed it in the next model. Their reason: nothing they tested reliably matched full attention on reasoning, coding and agent work. Another ships the mixture happily. The question is not settled.

Where you have already seen this

  • Models advertising a very long context that stays fast at the far end.
  • On-device assistants that run at a fixed memory budget.
  • Anything descended from the RWKV or Mamba families.

Remember this

  • Linear attention drops softmax so the arithmetic can be regrouped and made cheap.
  • It becomes a recurrent network with a fixed-size memory that never grows.
  • Constant cost per word, at the price of exact recall. Real models mix it with full attention.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Written against PyTorch 2.5.1, Python 3.10. CPU only, instant.

The regrouping, and the recurrence it implies

linear_attention.py
import torch

torch.manual_seed(0)
N, d = 8, 4
Q, K, V = torch.randn(N, d), torch.randn(N, d), torch.randn(N, d)
phi = torch.nn.functional.elu                      # feature map; elu(x)+1 keeps scores positive
q, k = phi(Q) + 1, phi(K) + 1

# 1. Two ways to bracket the same product. Non-causal, for clarity.
left  = (q @ k.T) @ V                              # N x N appears here
right = q @ (k.T @ V)                              # d x d appears here instead
print("bracketing is associative:", torch.allclose(left, right, atol=1e-5))
print("  (q k^T) is", tuple((q @ k.T).shape), " but  (k^T V) is", tuple((k.T @ V).shape))

# 2. Causal linear attention as a recurrence: one fixed-size state, one pass.
S = torch.zeros(d, d)                              # the whole memory of the past
z = torch.zeros(d)
out_rnn = torch.zeros(N, d)
for t in range(N):
    S = S + torch.outer(k[t], V[t])                # no dependence on t in the SIZE of S
    z = z + k[t]
    out_rnn[t] = (q[t] @ S) / (q[t] @ z + 1e-6)

# 3. The same thing written as masked quadratic attention, to check it.
scores = q @ k.T
mask = torch.tril(torch.ones(N, N))
scores = scores * mask
out_quad = (scores @ V) / (scores.sum(-1, keepdim=True) + 1e-6)
print("recurrent form matches quadratic form:", torch.allclose(out_rnn, out_quad, atol=1e-4))
print("max difference:", (out_rnn - out_quad).abs().max().item())

print(f"\nstate carried between steps: {S.numel()} numbers, and it never grows")
print("a softmax KV cache at step t holds 2 * t * d =", 2 * N * d, "numbers at t =", N)

# 4. Where the crossover sits.
print(f"\n{'seq len':>9}{'softmax N*N*d':>16}{'linear N*d*d':>15}{'ratio':>8}")
for n in (64, 512, 4096, 131072):
    a, b = n * n * d, n * d * d
    print(f"{n:>9}{a:>16,}{b:>15,}{a/b:>8.1f}")
Output
bracketing is associative: True
  (q k^T) is (8, 8)  but  (k^T V) is (4, 4)
recurrent form matches quadratic form: True
max difference: 4.76837158203125e-07

state carried between steps: 16 numbers, and it never grows
a softmax KV cache at step t holds 2 * t * d = 64 numbers at t = 8

  seq len   softmax N*N*d   linear N*d*d   ratio
       64          16,384          1,024    16.0
      512       1,048,576          8,192   128.0
     4096      67,108,864         65,536  1024.0
   131072  68,719,476,736      2,097,152 32768.0

Reading the output

(q k^T) is 8x8 and (k^T V) is 4x4. At N = 8, d = 4 the difference is unimpressive. At N = 131072, d = 128 it is a 131072x131072 matrix against a 128x128 one. Same product, different bracketing.

The recurrent form matches the quadratic form to 5e-07. This is the load-bearing check. The loop over t carries a d x d state and produces bit-comparable output to the masked matrix version. Linear attention is a recurrent network; the equality is not an approximation.

The state is 16 numbers and does not depend on t. That line is the property everything else follows from. Constant memory, constant time per token, no KV cache, no growth.

The last table is why anyone tolerates the quality loss. At 131k tokens the quadratic term is 32,768 times the linear one.

What the pieces are for

phi(x) = elu(x) + 1 is the feature map. Softmax guarantees non-negative weights; without it, scores can go negative and the normaliser can approach zero. The map from Transformers are RNNs keeps everything positive with cheap arithmetic. Other choices (random Fourier features, ReLU, learned maps) trade accuracy against cost.

S = S + outer(k[t], V[t]) is the memory update. Every token adds a rank-one term. Nothing is ever removed, which is the weakness later variants attack.

z is the running normaliser, playing the role of softmax's denominator.

The 1e-6 in the divisor is not decoration. With a poor feature map q[t] @ z can get very small, and the output explodes. Numerical fragility is a real reason linear attention is harder to train than it looks.

What modern variants change

The plain recurrence above never forgets. Write enough tokens and S saturates into noise. Every serious variant adds a way to forget or overwrite:

VariantAdded mechanism
Gated linear attentiona learned decay factor multiplying S each step
DeltaNeta delta rule that replaces rather than adds, removing the old value for a key
Gated DeltaNetboth: decay plus replacement
RWKVtime-decay weights with a token-shift mixing scheme

Qwen3-Next-80B-A3B uses Gated DeltaNet. Its config lists layer_types as three linear_attention layers followed by one full_attention, repeated across 48 layers, with linear_key_head_dim = 128 and linear_num_value_heads = 32. That 3:1 hybrid is the shape most current designs converge on.

Common mistakes

Expecting a speed-up at short context. At 512 tokens the quadratic path is faster in practice, because it maps to a dense matmul and linear attention's chunked scan does not. The crossover is thousands of tokens, not hundreds.

Writing the naive for t in range(N) loop for training. It is correct and unusably slow on a GPU. Real implementations use a chunked parallel scan: quadratic within a chunk, recurrent across chunks. Use flash-linear-attention or the kernels shipped with the model rather than writing your own.

Assuming constant memory means unlimited memory. The state is a fixed d x d. Information beyond its capacity is lost, and no amount of context length changes that.

Comparing on perplexity alone. Linear attention's weakness is exact recall, and perplexity averages over mostly-local predictions where it does fine. Test with retrieval tasks, as in advertised context vs usable context.

Treating "linear attention" as one method. Plain kernelised attention, gated variants, delta rules and state-space models have different capacities and very different results. The label covers a decade of divergent work.

Try it yourself

Add decay to the recurrence: S = 0.9 * S + torch.outer(k[t], V[t]). Confirm it no longer matches the quadratic form, then work out which weighted quadratic form it does match. That exercise is the whole content of the gated-linear-attention literature.

What to learn next

Researcher — Mathematics and papers.

Derivation

Write attention with a general similarity function $\operatorname{sim}$:

$$ o_i = \frac{\sum_{j \le i} \operatorname{sim}(q_i, k_j)\, v_j}{\sum_{j \le i} \operatorname{sim}(q_i, k_j)} $$

Softmax attention is $\operatorname{sim}(q,k) = \exp(q^{\top}k/\sqrt{d})$, which does not factorise. Katharopoulos, Vyas, Pappas and Fleuret (2020), Transformers are RNNs: Fast Autoregressive Transformers with Linear Attention (arXiv:2006.16236, ICML 2020) take $\operatorname{sim}(q,k) = \phi(q)^{\top}\phi(k)$ for a feature map $\phi: \mathbb{R}^{d} \to \mathbb{R}^{m}$ with non-negative outputs. Then

$$ o_i = \frac{\phi(q_i)^{\top} \sum_{j \le i} \phi(k_j) v_j^{\top}}{\phi(q_i)^{\top} \sum_{j \le i} \phi(k_j)} $$

Both sums are prefix sums independent of $i$, giving the recurrence

$$ S_i = S_{i-1} + \phi(k_i) v_i^{\top} \in \mathbb{R}^{m \times d}, \qquad z_i = z_{i-1} + \phi(k_i) \in \mathbb{R}^{m} $$

$$ o_i = \frac{S_i^{\top} \phi(q_i)}{z_i^{\top}\phi(q_i)} $$

Complexity is $\Theta(N m d)$ time and $\Theta(md)$ state, against $\Theta(N^2 d)$ and $\Theta(Nd)$. The paper's chosen map is $\phi(x) = \operatorname{elu}(x) + 1$, cheap and positive.

Feature-map families

Choromanski et al. (2020), Rethinking Attention with Performers (arXiv:2009.14794) construct FAVOR+, positive orthogonal random features whose inner products give an unbiased estimate of the softmax kernel. This is the one variant that approximates softmax rather than replacing it, and it recovers much of the quality at the cost of variance that grows with $m$.

Simpler maps (ReLU, identity-plus-shift, learned MLPs) abandon the softmax connection entirely and are trained from scratch. They are what current models use, because at scale the kernel-approximation framing stopped predicting quality.

Why plain linear attention underperforms

Three arguments, all defensible.

Capacity. The state is $m \times d$ numbers regardless of $N$. Softmax attention's implicit state is $2Nd$. Information-theoretically, exact recall of $N$ arbitrary key-value pairs from a fixed $md$ state is impossible for $N > m$. Associative-recall benchmarks show precisely this failure.

No forgetting. Additive updates make $S$ accumulate every token forever. The retrieval signal-to-noise ratio degrades as $O(1/\sqrt{N})$ under a random-vector model.

No sharpness. $\exp$ produces near-one-hot distributions when one key dominates. A polynomial-degree-one kernel cannot. Induction heads, the mechanism behind in-context learning, need that sharpness.

The modern fixes map onto these directly. Gating $S_i = \gamma_i \odot S_{i-1} + \phi(k_i)v_i^{\top}$ addresses forgetting; the delta rule $S_i = S_{i-1}(I - \beta_i k_i k_i^{\top}) + \beta_i k_i v_i^{\top}$ addresses overwriting; larger head dimensions address capacity.

Yang, Kautz and Hatamizadeh (2024), Gated Delta Networks: Improving Mamba2 with Delta Rule (arXiv:2412.06464, ICLR 2025) combine both — gating for fast erasure, the delta rule for targeted overwrite — and that construction is what Qwen3-Next ships as Gated DeltaNet.

The hybrid consensus

Pure linear models have not displaced transformers. Hybrids have gained real ground.

Qwen3-Next-80B-A3B interleaves three Gated DeltaNet layers per full-attention layer across 48 layers. Kimi Linear reports the same 3:1 ratio. The intuition is that exact retrieval is needed at a few points in the depth of the network, not at every layer, so a quarter of the layers can carry the associative-recall load while three quarters run in constant memory.

Cache implications are worth stating: in a 3:1 hybrid, only a quarter of layers hold a growing KV cache. Combine that with grouped-query attention on those layers and long-context serving cost falls by close to an order of magnitude.

The negative result

MiniMax shipped Lightning Attention in MiniMax-01 and reverted to full attention in MiniMax-M2, reporting that no efficient-attention variant they evaluated reliably matched full attention across reasoning, coding and agentic tasks in production.

That is a strong claim from a team with direct experience of both. Set against Qwen and Moonshot shipping hybrids, the reasonable reading is that hybrid ratios, gating design and training recipe matter more than the linear-versus-softmax label, and that the margin is narrow enough for teams to reach opposite conclusions in good faith.

Implementation reality

Never write the sequential recurrence for training. The standard approach is chunkwise parallel: split into chunks of 64 or 128, compute quadratic attention within a chunk, and propagate $S$ across chunks with a scan. This recovers matmul throughput while keeping linear complexity, and it is what every production kernel does.

Numerical stability needs care. The normaliser $z^{\top}\phi(q)$ can approach zero; most modern variants drop the normaliser entirely and rely on a normalisation layer after the attention output instead, which is more robust.

What to learn next