Back to Posts

Attention is recurrent

This blog is written to set some context about the inverse operation that comes up in many models like K3, Qwen, GLM-flash, ... and for our 2026 NeurIPS paper Fast and Stable Triangular Inversion for Delta-Rule Linear Transformers.

Most introductions to attention, define it as

$$O=\sigma (QK^T+M)V.$$

But this is a bit backwards, it's only the parallel form, used for prefill and training. And it makes the "causal mask" seem like a fundamental thing. But if you rather look from the recursion perspective, the casualness comes automatically, because when you auto-regressively generate a sequence, you can simply not use future tokens, as they do not yet exist. The parallel form is hugely important though, and it is what makes modern inference systems reasonable. Enabling LLMs to re-use shared prefixes efficiently, allowing systems to have huge system prompts and making back and forth conversation efficient. To get an understanding of how different they are the parallel form and decoding mode are, you can see 1000+ of tokens per second in prefill on macbooks with the M-chip, while the decoding part is stuck at 30 TPS, being dependent on the HBM/VRAM speed.1

When we generate tokens we have a prompt $x_{\leq n}=(x_0, x_1, ..., x_n)$ and we wish to compute $x_{n+1}$ and so on. I.e $p\left(x_{n+1} | x_{\leq n}\right)$.

We have regular attention defined as:

$$q_n, k_n, v_n = Proj(x_n)$$
$$o_n=\sigma \left( q_n K_{\leq n}^T\right)V_{\leq n}$$
$$h_n=f(o_n)$$
$$p\left(x_{n+1} | x_{\leq n}\right)=Proj(h_n).$$

And we have the generation process from prompt $x=(x_0,)$: $$x_0\longrightarrow o_0, h_0 \longrightarrow x_1\longrightarrow o_1,h_1 \longrightarrow x_2 \longrightarrow \cdots$$

With these definitions it turns out that parallel form, i.e how to compute $o_n$ directly for a prompt $x_{\leq n}=(x_0, x_1, ..., x_n)$ is embarrassingly parallel and we can compute all $o_0, o_1, ..., o_n$ simultaniously ( no need to first compute $o_0, o_1, ...$ in a sequential matter), allowing you to simple change $q$ for $Q$, and we suddenly go from GEMV to GEMMs and we see the matrix units go to work.

For recursive decoding, to generate $o_n$ they did not have to recalculate $k_{\leq n-1}$ and $v_{\leq n-1}$ as they were already calculated when $o_{n-1}$ was generated, i.e using the identities:2

$$K_{\leq n}=\begin{bmatrix}K_{\leq n-1} \\ k_n \end{bmatrix}, V_{\leq n}=\begin{bmatrix}V_{\leq n-1} \\ v_n \end{bmatrix}.$$

Linear attention

Why? You can deepen your Transformers without any big compute increase by interleaving linear layers. It can be used as positional encoding, so the full attention layers can get rid of RoPE. You can scale the context lengths without hitting the quadratic time complexity wall.

Now let's remove the softmax, we can collapse the history into a single time independent state $S\in \mathbb{R}^{d \times d}$.3 Since now we are no longer just keeping the information from all previous tokens, but compressing them into a fixed-sized state, there are more engineering tuning/hacks to improve the compression, just look at the complexity of LSTMs and GRUs...

Starting from the auto-regressive form

$$o_n = \sigma(q_nK_{\leq n}^T)V_{\leq n},$$

we get

$$o_n = q_nK_{\leq n}^TV_{\leq n}$$
$$o_n = \sum_{i=0}^n (q_nk_i^T)v_i$$
$$o_n = q_n\left(\sum_{i=0}^n k_i^Tv_i\right).$$

This is the key part, as if we define the summation as

$$S_n=\sum_{i=0}^n k_i^Tv_i,\qquad S_n\in\mathbb{R}^{d\times d},$$

we can be smart when calculating the attention for the $n$ th token, $o_n$, we don't have to recalculate the whole sum if we keep track of the last state $S_{n-1}$:

$$S_n=S_{n-1}+k_n^Tv_n$$
$$o_n=q_nS_n.$$

So when we want to calculate $o_n$ in linear attention, instead of keeping $(k_0,v_0),(k_1,v_1),\ldots,(k_n,v_n)$ i.e $(K_{\leq n}, V_{\leq n})$ (the KV-cache), we only need to keep track of the previous state $S_{n-1}$. We can say that for linear transformer $o_n=o_n(S_{n-1}, q_n, k_n, v_n)$, while regular attention $o_n=o_n(K_{\leq n}, V_{\leq n})$. The former arguments are all fixed-size, while the latter arguments grow linearly with $n$.

So generation (and at arrow between $i-1$ and $i$, the previous state $S_{i-1}$ is cached) now looks like

$$x_0\longrightarrow S_0,o_0,h_0\longrightarrow x_1\longrightarrow S_1,o_1,h_1\longrightarrow x_2\longrightarrow \cdots,$$

with

$$p(x_{n+1}\mid x_{\leq n})=Proj(h_n).$$

Prefill

Now suppose our known prompt is

$$x=(x_0,x_1,x_2,x_3).$$

Sequentially we would compute

$$S_0=k_0^Tv_0$$
$$S_1=S_0+k_1^Tv_1$$
$$S_2=S_1+k_2^Tv_2$$
$$S_3=S_2+k_3^Tv_3.$$

Then

$$o_i=q_iS_i.$$

The last output

$$o_3$$

is what eventually gives us

$$p(x_4\mid x_0,x_1,x_2,x_3).$$

Again, $x_4$ is not part of the prefill. We prefill $x_0,\ldots,x_3$, and the representation at the last known position predicts $x_4$.

Unlike regular attention, there is now an actual recurrence

$$S_0\to S_1\to S_2\to S_3.$$

But luckily the recurrence is just a cumulative sum:

$$S_n=\sum_{i\leq n} k_i^Tv_i.$$

And addition is associative, s o all the prefix states $S_0,S_1,S_2,S_3$ can be computed with a parallel prefix scan. This is the first place where the difference to regular attention becomes interesting.

For regular attention there wasn't really a state recurrence at all. Once $Q,K,V$ were known, every $o_i$ could be computed independently, so prefill was basically just stacking the autoregressive equations into a GEMM.

For linear attention there is a recurrent state, but the recurrence happens to be one of the easiest possible ones:

$$S_n=S_{n-1}+B_n,\qquad B_n=k_n^Tv_n.$$

So plain linear attention is just the simplest associative scan (a cumsum).

$$S_n=S_{n-1}+k_n^Tv_n,\qquad o_n=q_nS_n.$$

Note on non-causal simple linear attention

if we just remove the softmax from non-causal attention we have

$$O=(QK^T)V,$$

which we can regroup as $O=Q(K^TV)$ which complexity is linear in $n$. (https://en.wikipedia.org/wiki/Matrix_chain_multiplication) But when we have a causal pre-fill sequence, we must now use multiplicative mask $L$ which has entries in $\{0, 1\}$ rather than an additive $M$ with $\{-\infty, 0\}$ entries get zerod by the softmax,

$$O=\left(QK^T \odot L\right)V.$$

And we can sadly not do this re-arrangment as $\odot$ is elementwise multiplication.4

Gated linear attention

Now we can make the state forget with parameter $\alpha_n$, so instead of $S_n=S_{n-1}+k_n^Tv_n$ we have

$$S_n=\alpha_nS_{n-1}+k_n^Tv_n.$$

Autoregressively this is still trivial:

$$S_{n-1}\longrightarrow \alpha_nS_{n-1}+k_n^Tv_n\longrightarrow S_n.$$

If we unroll it,

$$S_n=k_n^Tv_n+\alpha_nk_{n-1}^Tv_{n-1}+\alpha_n\alpha_{n-1}k_{n-2}^Tv_{n-2}+\cdots.$$

So an old write at position $i$ survives until position $n$ with weight

$$\prod_{j=i+1}^n\alpha_j.$$

This is no longer the simplest assocative scan, as it's now a weighted scan.

Delta attention

Instead of blindly adding the new value, we first ask: what does the current memory already predict for this key? The current prediction is

$$\hat v_n=k_nS_{n-1}.$$

So rather than writing $v_n$, we write only the error

$$v_n-\hat v_n.$$

With a write strength $\beta_n$,

$$S_n=S_{n-1}+\beta_n k_n^T\left(v_n-k_nS_{n-1}\right).$$

This is still a simple autoregressive update. Given $S_{n-1}$ and the new token $x_n$, compute

$$q_n,k_n,v_n,\beta_n$$

then

$$u_n=\beta_n\left(v_n-k_nS_{n-1}\right)$$
$$S_n=S_{n-1}+k_n^Tu_n$$

and finally

$$o_n=q_nS_n.$$

Here $u_n$ is the actual corrected value we write into memory. And now something important has changed. For ordinary linear attention,

$$\text{write}_n=k_n^Tv_n$$

only depends on token $n$. But for delta attention,

$$\text{write}_n=k_n^Tu_n$$

and

$$u_n=\beta_n(v_n-k_nS_{n-1})$$

depends on the entire previous state. So the writes themselves are now causally coupled. This is why prefill stops being just a cumsum. What happens during prefill?

Prefill

Take again

$$x_0,x_1,x_2,x_3.$$

Assume $S_{-1}=0$.

Then

$$u_0=\beta_0v_0$$

and

$$S_0=k_0^Tu_0.$$

For token $1$,

$$u_1=\beta_1\left(v_1-k_1k_0^Tu_0\right).$$

For token $2$,

$$u_2=\beta_2\left(v_2-k_2k_0^Tu_0-k_2k_1^Tu_1\right).$$

And for token $3$,

$$u_3=\beta_3\left(v_3-k_3k_0^Tu_0-k_3k_1^Tu_1-k_3k_2^Tu_2\right).$$

So unlike normal attention, we cannot just say "compute every row independently".

And unlike vanilla linear attention, we cannot just cumsum independent writes either. The write $u_3$ depends on $u_2$, which depends on $u_1$, which depends on $u_0$. At first this looks inherently sequential. But notice that these dependencies are all linear. Move the previous writes to the left:

$$u_0=\beta_0v_0$$
$$u_1+\beta_1(k_1k_0^T)u_0=\beta_1v_1$$
$$u_2+\beta_2(k_2k_0^T)u_0+\beta_2(k_2k_1^T)u_1=\beta_2v_2$$

and so on.

So for the whole prompt we get one lower-triangular system

$$AU=BV,$$

where

$$U=\begin{bmatrix}u_0 \\ u_1 \\ u_2\\ u_3 \end{bmatrix}.$$

Therefore

$$U=A^{-1}BV.$$

And this is where the inverse in the parallel/chunkwise GDN/KDA formulas comes from.

It is just what happens when we take a causal recurrence

$$u_n=\beta_n ( v_n-\sum_{i < n} c_{ni} u_i )$$

and ask:

instead of solving $u_0$, then $u_1$, then $u_2$, can we solve all of them together?

The answer is a triangular solve.

So now the little hierarchy becomes

$$\text{linear attention}\Rightarrow\text{cumsum}$$
$$\text{gated linear attention}\Rightarrow\text{weighted associative scan}$$
$$\text{delta attention}\Rightarrow\text{triangular solve (inverse) / more complicated associative scan}.$$