
If the similarity in attention is expressed as a kernel using a feature map , can be separated out of the sum over . Then and can be computed only once and reused by every query, so the computation time and memory become .
With causal masking, it is enough to carry 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.
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 it has time complexity and memory , 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 .
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 .
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.
Context is the maximum range of the sequence that can be used when computing self-attention. In 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.
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.
The input is a sequence of feature vectors of dimension . A Transformer layer consists of self-attention and a feedforward .
Here , , . The -th output of softmax attention can be generalized using a similarity function .
If , it becomes the standard softmax attention. That is, as long as is non-negative, any function can be used as attention.
If the kernel is expressed with a feature map so that ,
Since is independent of , it can be pulled out of the sum. Because can be separated, the computation for can be done only once.
Since and are computed only once and reused by every query, time and memory become linear in .
This paper chose the following feature map.
The function has the following properties
In an autoregressive model that generates one token at a time, the -th position must only be influenced by positions with . To apply causal masking to the linearized attention expressed as in , only the range of the sum needs to change from to .
Here, the terms grouped by are sums from the first to the -th position, growing by one term as increases. That is, instead of computing from every time, it can be thought of as a cumulative sum where one value is added for each new position .
If the cumulative sums are defined as follows
and they are updated from the previous values in constant time as , . Therefore, the entire causal attention is linear in sequence length.
If all were stored and differentiated with autograd, memory would become , possibly using even more memory than softmax. Therefore, the paper solves this by expressing the gradients themselves as cumulative sums as well.
In this way, both forward and backward can be processed in linear time and memory.
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.
A softmax Transformer must store all past tokens' K and V in the KV cache, so the cost per generation step grows as , whereas a Linear Transformer only needs to maintain a single fixed-size state, so the per-step cost and memory are constant.
The three compared methods are as follows.
The experiments cover a synthetic task, image generation (MNIST, CIFAR-10), and more.
Convergence behavior was compared on a copy task (copying a sequence of symbols as is), similar to sequence duplication.
GPU peak memory and forward/backward time were measured while varying the sequence length from to .
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 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 to . 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.
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 . Linear Transformer discards per-token K and V of past tokens altogether and summarizes them into a single state of size . It is a fundamentally different approach in that the KV cache size becomes independent of .
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 simply adds up all , 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.
This post was translated from the original Korean version with the help of AI, so some errors may remain.