← All Posts
Deep Learning · Transformers· Recurrent and hybrid models

Linear Attention: From Pairwise Reads to a Recurrent State

Linear attention changes what is remembered. Instead of keeping every previous key/value pair, it accumulates their contributions into a fixed-size state. This can make each decoding update independent of context length, at the cost of a different retrieval mechanism.

The reassociation that works—and the one that does not

Without softmax, $(QK^\top)V=Q(K^\top V)$ by associativity. The left side suggests an $n\times n$ intermediate; the right side first forms a $d_k\times d_v$ matrix. Ordinary softmax attention cannot be rearranged this way because row-wise normalization and exponentiation sit between the two multiplications.

Define a different similarity kernel $\operatorname{sim}(q,k)=\phi(q)^\top\phi(k)$, with feature map $\phi:\mathbb R^{d_k}\to\mathbb R^r$. A normalized causal read is

$$o_t=\frac{\sum_{i\le t}\phi(k_i)^\top\phi(q_t)\,v_i}{\sum_{i\le t}\phi(k_i)^\top\phi(q_t)}.$$

Positive features are one way to make the denominator and weights nonnegative. For the examples here, assume the denominator is strictly positive. An epsilon clamp is a numerical convention that changes the operation near zero; it should not silently replace a missing mathematical assumption.

Store the sufficient sums

Use column vectors and define the state with key-feature rows and value columns:

$$S_t=S_{t-1}+\phi(k_t)v_t^\top\in\mathbb R^{r\times d_v},\qquad z_t=z_{t-1}+\phi(k_t)\in\mathbb R^r,$$
$$o_t=\frac{S_t^\top\phi(q_t)}{z_t^\top\phi(q_t)}.$$

The matrix stores value-weighted key features; the vector stores the corresponding normalization features. Both are independent of sequence length in shape. The query is applied only when reading. The state changes with the input sequence but is not a learned model parameter.

An exact two-token read

Let features already be $k_1=(1,0)$, $k_2=(1,1)$, scalar values be $v_1=2$, $v_2=4$, and query be $q=(1,1)$. After the first update, $S_1=(2,0)^\top$, $z_1=(1,0)$. After the second, $S_2=(6,4)^\top$, $z_2=(2,1)$.

The recurrent read is $(6+4)/(2+1)=10/3$. Direct pairwise similarity gives weights proportional to $(q^\top k_1,q^\top k_2)=(1,2)$, hence $(1\cdot2+2\cdot4)/3=10/3$. This equality is exact for this kernel. Applying softmax to scores $(1,2)$ would instead give a different result.

Try the mechanism Write associations and compare two ways to read

This equals the direct kernel-weighted average: the two keys receive unnormalized weights 1 and 2. Enable JavaScript to change the example inputs; the complete calculation remains in the article.

Linear in which variable?

Each update and read costs $O(rd_v)$ and the persistent state occupies $rd_v+r$ scalars per head, ignoring feature-map computation. Across $n$ tokens, recurrence costs $O(nrd_v)$. It is linear in sequence length at fixed feature and value dimensions, not linear in every model dimension.

For $r=d_v=128$, the state has 16,512 scalars per head. At short contexts, a conventional K/V cache may be smaller; at sufficiently long contexts it grows beyond this fixed state. Training can still require activations, chunk boundary states, and gradients. Fixed decode state is not a statement that an entire training run has constant memory.

A recurrence need not force token-serial training

The additive states are prefix sums. Parallel scans or chunked matrix operations can compute them across a known training sequence. A chunk implementation separates a within-chunk contribution from the state arriving from previous chunks, letting much of the work use matrix multiplication.

Which formulation is fastest depends on hardware and dimensions. A Python for-loop demonstrates the recurrence but is not a credible GPU throughput benchmark. Transformers Are RNNs develops the linear-attention connection between parallel and recurrent views.

What compression makes difficult

The state superposes contributions from many tokens. Different histories can map to the same state, particularly when keys overlap. It cannot generally provide independently addressable storage for an unbounded collection of arbitrary key/value pairs at fixed state size and precision.

That observation does not imply recurrence is useless for long context. Many tasks admit useful compressed statistics. It does explain why exact retrieval, copying, and overwriting associations are useful diagnostic tasks, and why some models retain occasional full-attention layers.

Exact for a chosen kernel is different from approximating softmax

An explicit feature map such as $\operatorname{ELU}(x)+1$ defines its own kernel. Random-feature methods can instead approximate a softmax-like kernel, with approximation error depending on the feature construction and dimension. Performer is a primary reference for that route.

Neither approach licenses writing “linear attention computes exact softmax in linear time.” The exactness claim must name the kernel and the normalization being computed. FlashAttention preserves softmax while addressing memory traffic, and keeps quadratic pairwise arithmetic for dense attention.

From adding associations to editing them

The normalized accumulator here includes $z_t$. Some modern linear recurrent modules instead use an unnormalized matrix state followed by learned normalization and gating; they are not this exact weighted-average formula. DeltaNet begins that branch by replacing additive writes with error-correcting writes.

Try it: If two histories produce the same $S_t$ and $z_t$, can this read distinguish them using the same query?

No. Its output depends only on those states and the query. Any information distinguishing those histories has been discarded by this representation.

Run the parallel-versus-recurrent identity in the numerical lab before changing the state update rule.