[Paper Review] Transformers are RNNs: Fast Autoregressive Transformers with Linear Attention

Jumyung Song·2026년 9월 29일

Paper review

목록 보기
4/6
post-thumbnail

Summary

  • If the similarity in attention is expressed as a kernel ϕ(Qi)Tϕ(Kj)\phi(Q_i)^T\phi(K_j) using a feature map ϕ\phi, ϕ(Qi)\phi(Q_i) can be separated out of the sum over jj. Then ∑jϕ(Kj)VjT\sum_j \phi(K_j)V_j^T and ∑jϕ(Kj)\sum_j \phi(K_j) can be computed only once and reused by every query, so the computation time and memory become O(N)O(N).

  • With causal masking, it is enough to carry ∑j≤iϕ(Kj)VjT\sum_{j \le i}\phi(K_j)V_j^T as a cumulative sum (state), which means that a Transformer layer = an RNN with a fixed-size hidden state.

  • As a result, autoregressive generation runs inference with constant time and constant memory per step without a KV cache, and is about 4,462 times faster than softmax on CIFAR-10 image generation.


1. Problem & Motivation

Transformers have shown strong performance on a variety of tasks including natural language, and pre-training with autoregressive / masked LM objectives yields powerful representations even from unlabeled data. However, since self-attention computes the similarity of every key-query pair, for an input of length NN it has time complexity O(N2)O(N^2) and memory O(N2)O(N^2), which limits the context length. In particular, in autoregressive inference, attention is applied over all tokens so far every time a single token is generated, so the per-step cost keeps growing as generation gets longer.

To reduce the existing time complexity and memory, this paper introduces the Linear transformer architecture.

Linear Transformer computes similarity with a kernel-based formulation so that both memory and computation become O(N)O(N).

2.1. Efficient Transformers

Several methods have previously been proposed to speed up training and inference. Weight pruning, factorization, quantization, etc. speed up training and inference, but the time complexity of the attention operation is still O(N2)O(N^2).

Attempts to increase context by reducing complexity are as follows. Both methods*reduce complexity by having each query look at only a subset of keys.

  • Sparse Transformer (Child et al.): factorizes the attention matrix sparsely so that each query only looks at positions in a fixed pattern, with complexity O(NN)O(N\sqrt{N}).
  • Reformer: puts similar vectors into the same bucket and computes attention only within a bucket. Hashing requires the constraint Q=KQ = K, and the complexity is O(Nlog⁡N)O(N\log N).

Context is the maximum range of the sequence that can be used when computing self-attention. In O(N2)O(N^2) attention, as the sequence gets longer, the amount of computation and the memory for storing the attention matrix grow quadratically, which limits the context length. Reducing complexity allows longer sequences to be processed with the same resources, so reducing complexity can ultimately be seen as a way to handle long sequences.

Unlike previous methods, Linear Transformer scales linearly with sequence length without any constraints on queries and keys.

2.2. Understanding Self-Attention

From a kernel perspective, attention can be thought of as applying to the input a kernel smoother, which takes an average giving higher weights to more similar data. Here, the kernel value is the similarity between inputs.

This paper extends the above perspective and applies the idea that any kernel giving positive similarity scores can be used as attention. Additionally, it shows that a self-attention layer trained with an autoregressive objective can be viewed as an RNN.

3. Linear Transformers

3.1. Transformers

The input x∈RN×Fx \in \mathbb{R}^{N \times F} is a sequence of NN feature vectors of dimension FF. A Transformer layer TlT_l consists of self-attention AlA_l and a feedforward flf_l.

Tl(x)=fl(Al(x)+x),Al(x)=softmax(QKTD)VT_l(x) = f_l(A_l(x) + x), \qquad A_l(x) = \text{softmax}\left(\frac{QK^T}{\sqrt{D}}\right)V

Here Q=xWQQ = xW_Q, K=xWKK = xW_K, V=xWVV = xW_V. The ii-th output of softmax attention can be generalized using a similarity function sim(⋅,⋅)\text{sim}(\cdot,\cdot).

Vi′=∑j=1Nsim(Qi,Kj) Vj∑j=1Nsim(Qi,Kj)V_i' = \frac{\sum_{j=1}^{N} \text{sim}(Q_i, K_j)\, V_j}{\sum_{j=1}^{N} \text{sim}(Q_i, K_j)}

If sim(q,k)=exp⁡(qTkD)\text{sim}(q, k) = \exp\left(\frac{q^Tk}{\sqrt{D}}\right), it becomes the standard softmax attention. That is, as long as sim\text{sim} is non-negative, any function can be used as attention.

3.2. Linearized Attention

If the kernel k(x,y):R2×F→R+k(x, y): \mathbb{R}^{2 \times F} \to \mathbb{R}_+ is expressed with a feature map ϕ\phi so that sim(Qi,Kj)=ϕ(Qi)Tϕ(Kj)\text{sim}(Q_i, K_j) = \phi(Q_i)^T\phi(K_j),

Vi′=∑j=1Nϕ(Qi)Tϕ(Kj) Vj∑j=1Nϕ(Qi)Tϕ(Kj)=ϕ(Qi)T∑j=1Nϕ(Kj)VjTϕ(Qi)T∑j=1Nϕ(Kj)V_i' = \frac{\sum_{j=1}^{N} \phi(Q_i)^T \phi(K_j)\, V_j}{\sum_{j=1}^{N} \phi(Q_i)^T \phi(K_j)} = \frac{\phi(Q_i)^T \sum_{j=1}^{N} \phi(K_j) V_j^T}{\phi(Q_i)^T \sum_{j=1}^{N} \phi(K_j)}

Since ϕ(Qi)\phi(Q_i) is independent of jj, it can be pulled out of the sum. Because QiQ_i can be separated, the computation for QiQ_i can be done only once.

(ϕ(Q)ϕ(K)T)V=ϕ(Q)(ϕ(K)TV)\left(\phi(Q)\phi(K)^T\right)V = \phi(Q)\left(\phi(K)^TV\right)
  • The original approach builds the N×NN \times N matrix first → O(N2)O(N^2)
  • The approach with QIQ_I pulled out builds ϕ(K)TV\phi(K)^TV (a D×MD \times M matrix) first → O(N)O(N)

Since ∑jϕ(Kj)VjT\sum_{j}\phi(K_j)V_j^T and ∑jϕ(Kj)\sum_{j}\phi(K_j) are computed only once and reused by every query, time and memory become linear in NN.

3.3. Feature Maps and Computational Cost

This paper chose the following feature map.

ϕ(x)=elu(x)+1,elu(x)={x(x>0)α(ex−1)(x≤0)\phi(x) = \text{elu}(x) + 1, \qquad \text{elu}(x) = \begin{cases} x & (x > 0) \\ \alpha(e^x - 1) & (x \le 0) \end{cases}

The function ϕ(x)=elu(x)+1\phi(x)=elu(x)+1 has the following properties

  • Thanks to the +1+1, ϕ(x)>0\phi(x) > 0, so the similarity is always positive
  • It prevents the gradient from becoming 0 when xx is negative.

3.4. Causal Masking

In an autoregressive model that generates one token at a time, the ii-th position must only be influenced by positions with j≤ij \le i. To apply causal masking to the linearized attention expressed as in 3.2.3.2., only the range of the sum needs to change from NN to ii.

Vi′=∑j=1isim(Qi,Kj) Vj∑j=1isim(Qi,Kj)=ϕ(Qi)T∑j=1iϕ(Kj)VjTϕ(Qi)T∑j=1iϕ(Kj)V_i' = \frac{\sum_{j=1}^{i} \text{sim}(Q_i, K_j)\, V_j}{\sum_{j=1}^{i} \text{sim}(Q_i, K_j)} = \frac{\phi(Q_i)^T \sum_{j=1}^{i} \phi(K_j) V_j^T}{\phi(Q_i)^T \sum_{j=1}^{i} \phi(K_j)}

Here, the terms grouped by ∑\sum are sums from the first to the ii-th position, growing by one term as ii increases. That is, instead of computing from i=1i=1 every time, it can be thought of as a cumulative sum where one value is added for each new position ii.

If the cumulative sums are defined as follows

Si=∑j=1iϕ(Kj)VjT,Zi=∑j=1iϕ(Kj)S_i = \sum_{j=1}^{i} \phi(K_j) V_j^T, \qquad Z_i = \sum_{j=1}^{i} \phi(K_j)
Vi′=ϕ(Qi)TSiϕ(Qi)TZiV_i' = \frac{\phi(Q_i)^T S_i}{\phi(Q_i)^T Z_i}

and they are updated from the previous values in constant time as Si=Si−1+ϕ(Ki)ViTS_i = S_{i-1} + \phi(K_i)V_i^T, Zi=Zi−1+ϕ(Ki)Z_i = Z_{i-1} + \phi(K_i). Therefore, the entire causal attention is linear in sequence length.

3.5. Gradient Computation

If all SiS_i were stored and differentiated with autograd, memory would become O(NDM)O(NDM), possibly using even more memory than softmax. Therefore, the paper solves this by expressing the gradients themselves as cumulative sums as well.

∇ϕ(Qi)L=∇VˉiL(∑j=1iϕ(Kj)VjT)T\nabla_{\phi(Q_i)}\mathcal{L} = \nabla_{\bar{V}_i}\mathcal{L} \left(\sum_{j=1}^{i} \phi(K_j) V_j^T\right)^T
∇ϕ(Ki)L=(∑j=iNϕ(Qj)(∇VˉjL)T)Vi\nabla_{\phi(K_i)}\mathcal{L} = \left(\sum_{j=i}^{N} \phi(Q_j) \left(\nabla_{\bar{V}_j}\mathcal{L}\right)^T\right) V_i
∇ViL=(∑j=iNϕ(Qj)(∇VˉjL)T)Tϕ(Ki)\nabla_{V_i}\mathcal{L} = \left(\sum_{j=i}^{N} \phi(Q_j) \left(\nabla_{\bar{V}_j}\mathcal{L}\right)^T\right)^T \phi(K_i)

In this way, both forward and backward can be processed in linear time and memory.

3.6. Transformers and RNNs

The key idea of interpreting linear attention as an RNN is as follows.

Any Transformer layer with causal masking can be written as a model that takes an input, updates an internal state, and produces an output, i.e., an RNN.

s0=0,z0=0si=si−1+ϕ(xiWK)(xiWV)Tzi=zi−1+ϕ(xiWK)yi=fl(ϕ(xiWQ)Tsiϕ(xiWQ)Tzi+xi)\begin{aligned} s_0 &= 0, \quad z_0 = 0 \\ s_i &= s_{i-1} + \phi(x_i W_K)(x_i W_V)^T \\ z_i &= z_{i-1} + \phi(x_i W_K) \\ y_i &= f_l\left(\frac{\phi(x_i W_Q)^T s_i}{\phi(x_i W_Q)^T z_i} + x_i\right) \end{aligned}
  • Hidden state: si∈RD×Ms_i \in \mathbb{R}^{D \times M} (attention memory), zi∈RDz_i \in \mathbb{R}^{D} (normalizer memory)
  • When a new token xix_i comes in, the state is updated, and the output is produced by reading the state with the current query.
  • During training it can be used like a Transformer that processes all positions in parallel, and during inference like an RNN that carries only the state.

A softmax Transformer must store all past tokens' K and V in the KV cache, so the cost per generation step grows as O(N)O(N), whereas a Linear Transformer only needs to maintain a single fixed-size state, so the per-step cost and memory are constant.

4. Experiments

4.1. Setup

The three compared methods are as follows.

  • Softmax: standard full attention Transformer
  • LSH-X: Reformer. X is the number of hashing rounds
  • Linear: Linear Transformer using ϕ(x)=elu(x)+1\phi(x) = \text{elu}(x) + 1

The experiments cover a synthetic task, image generation (MNIST, CIFAR-10), and more.

4.2. Synthetic Task: Convergence

Convergence behavior was compared on a copy task (copying a sequence of symbols as is), similar to sequence duplication.

  • Linear converges stably (smoothly) and reaches a lower loss than LSH-4.
  • It eventually reaches a loss similar to softmax. However, in the middle range (about 4,000–7,000 steps), softmax shows a lower loss.

4.3. Synthetic Task: Memory & Computation

GPU peak memory and forward/backward time were measured while varying the sequence length from 292^9 to 2162^{16}.

  • Softmax grows quadratically with length and cannot be measured beyond 2122^{12} (out of GPU memory).
  • Both LSH and Linear grow linearly, but Linear is the fastest and uses the least memory at every length.

4.4. Image Generation: MNIST & CIFAR-10

Images are generated autoregressively pixel by pixel. The evaluation metrics are bits/dim (lower is better) and the number of images generated per second.

MNIST

CIFAR-10

  • Linear achieves quality similar to softmax while generating 317 times faster on MNIST and 4,462 times faster on CIFAR-10. That is, while softmax generates a single CIFAR-10 image, Linear generates about 4,460 images.
  • The speed gap becomes much larger on CIFAR-10, which has longer sequences. This is because softmax's cost grows every step, while Linear's per-step cost stays constant like an RNN.

5. Conclusion

Linear Transformer replaces attention with an inner product of kernel feature maps and uses the associativity of matrix multiplication to reduce time and memory complexity from O(N2)O(N^2) to O(N)O(N). In addition, by interpreting causal attention as an RNN, it enables autoregressive inference with constant time and constant memory. In experiments, it maintained performance close to softmax while achieving speedups of hundreds to thousands of times in autoregressive generation.

6. My Take

1. An optimization in a different direction from MQA → GQA → MLA
MQA, GQA, and MLA, which I have read so far, were all answers to "how small can the KV cache be stored?" Whether heads are shared (MQA, GQA) or compressed into a latent (MLA), the cache still grows in proportion to the number of tokens NN. Linear Transformer discards per-token K and V of past tokens altogether and summarizes them into a single state of size D×MD \times M. It is a fundamentally different approach in that the KV cache size becomes independent of NN.

2. Compressing the entire past into a fixed-size state
Softmax attention keeps all past tokens and can retrieve exactly what it needs. In contrast, Linear attention's SiS_i simply adds up all ϕ(Kj)VjT\phi(K_j)V_j^T, so as the length grows, information gets mixed and it becomes hard to accurately recover a specific token. I think the much lower loss of softmax in the middle range of the copy task and the slightly worse MNIST bits/dim can also be seen in connection with this limitation.

References

  • Katharopoulos, A., Vyas, A., Pappas, N., & Fleuret, F. (2020). Transformers are RNNs: Fast Autoregressive Transformers with Linear Attention. ICML 2020. arXiv:2006.16236
  • Vaswani, A. et al. (2017). Attention Is All You Need. arXiv:1706.03762
  • Child, R. et al. (2019). Generating Long Sequences with Sparse Transformers. arXiv:1904.10509

This post was translated from the original Korean version with the help of AI, so some errors may remain.

0개의 댓글