1The problem attention solves
A sequence model reads vectors , one per token. At the input these are token embeddings; deeper in the network they are the hidden states of the previous layer. At position the model has to produce an output that may depend on anything seen so far. Which earlier tokens matter is not fixed by position. In "the cat that the dog chased ran away", the verb "ran" needs "cat", five tokens back, and in another sentence the relevant word will be somewhere else. The model has to find earlier information by its content.
The simplest content-based lookup uses the dot product. For two vectors, is large when they point in similar directions, zero when they are orthogonal, and negative when they point apart. So position could score every earlier token by and return an average of the weighted by these scores.
This asks one vector to do three different jobs. As the thing doing the searching, describes what position is looking for. As the thing being searched, describes how position can be found. As the result, is also what gets returned. Using a single vector for all three causes concrete failures. A token scores highest against tokens that resemble itself, so it mostly retrieves near-copies of itself. The score is symmetric, so if "ran" finds "cat", "cat" finds "ran" equally well, although what each needs from the other is different. And what is returned is the raw token, with no choice about which of its features to pass on.
The repair is to give each job its own learned linear map:
The query is what position is looking for. The key is how position advertises itself to queries. The value is what position hands over when it is found. Queries and keys are compared by a dot product, so they must live in the same space, . Values are never compared with anything, so their dimension is a separate choice. The three matrices are learned by gradient descent along with the rest of the network; this post takes them as given and asks what happens once the vectors exist.
2Softmax attention
The scores can be any real numbers. To average values we want nonnegative weights that sum to one. Exponentiating makes every score positive, and dividing by the total makes the weights sum to one:
The exponential does more than make weights positive. A score higher by gets times the weight, so if one key matches the query clearly better than the others, its value dominates the average and the lookup behaves almost like picking a single entry. The sum runs over so that position never uses tokens that come after it, which is what makes the model usable for generation. This operation is attention (Vaswani et al., 2017). We leave out the scaling and the use of several heads in parallel; neither changes anything below.
Now look at where appears. It sits inside every exponential and inside the denominator, which couples all the keys together. Before the query arrives there is nothing useful to compute in advance, so every pair has to be kept until it is needed. This store is the key-value cache, and it grows with . Each new token must be compared with all earlier ones, so producing tokens costs time proportional to .
Suppose instead we are only allowed a fixed number of numbers, a matrix , and we must write each pair into it as it arrives and discard the pair afterwards. Two questions follow. How should a pair be written? And how many pairs can such a matrix hold before reads start to go wrong?
3Removing the softmax
Replace by and drop the denominator (Katharopoulos et al., 2020). Each term is a vector times a number, and that number is linear in . Writing and then multiplying by is the same as first forming the matrix and then applying it to . Since every term has this form, the query can be pulled out of the sum:
The matrix is the outer product of and , with entries . Pulling out of the sum is the entire derivation. Everything the past contributes to any future read is collected into before the query is known. The pairs can be thrown away, and is updated by adding one term per step:
A model that carries a fixed-size state from one step to the next and updates it with each input is a recurrent network, and is its state. Storing a pair means adding its outer product to the state. Retrieving means multiplying the state by a query.
The same matrix appeared long before transformers, as the correlation matrix memory (Kohonen, 1972; Anderson, 1972), and Schlag et al. (2021) show that linear attention is a fast weight programmer, a network whose weights are rewritten at every step. To understand what this memory can and cannot do, we need to understand the object being added.
4One stored pair
Apply to an arbitrary input :
Read the right-hand side from the inside out. The scalar depends only on the component of along . That number is then used as a coefficient on . So the map ignores every direction of the input perpendicular to , and it can only ever output multiples of . The set of inputs sent to zero, the null space, is the set of vectors perpendicular to . The set of possible outputs, the image, is the line through . A matrix whose image is a single line has rank 1, and every rank-1 matrix is an outer product of this kind.
Querying with the key itself gives . The stored value comes back scaled by the squared length of the key, so for exact recall we want keys of length one. We assume from here on.
In the figures , so keys, queries and values are arrows in a plane and everything can be drawn. Nothing in the argument depends on the dimension being 2. The left panel of each figure is the space of keys and queries, the right panel is the space of values, and the matrix in between is the map from one to the other.
Inputs (drag )
Outputs
5Two pairs, and the Gram matrix
Store a second pair. The memory is , and a read is the sum of what each outer product does on its own:
The query is dropped onto each key separately, and each drop weights its own value. Querying with gives , where is the angle between the keys. The second term is interference: part of the other value comes back as well, in proportion to how much the keys overlap.
Keys and query (drag )
Values and output
The same computation works for any number of pairs. Stack the keys as the columns of and the values as the columns of . The sum of outer products is , and querying with every stored key at once gives
The matrix of all dot products between keys is the Gram matrix. Column of is the read for , namely , so column of lists how much of each value comes back. The diagonal entries are , the wanted value at full strength. The off-diagonal entries are the cosines between keys, and each one is the weight of a wrong value leaking into a read.
Every stored value is read back exactly if and only if , which means the keys are orthonormal: unit length and mutually perpendicular. The values can point anywhere. In at most vectors are mutually perpendicular, so at most pairs can be stored without interference, and in the plane a third key always overlaps the first two. The capacity depends on the key dimension alone. The values are only written down; the keys are what has to be kept apart. Schlag et al. (2021) discuss this limit for linear transformers.
This explains why softmax attention does not run out of room. Expand the exponential as a power series, . Each power is itself a dot product, between the vectors of all -fold products of coordinates of and of . Collecting these vectors for every gives a map with , where has infinitely many coordinates. Apart from the denominator, softmax attention is the memory of this post with keys in an infinite-dimensional space, where the dimension bound never applies. The price is that cannot be stored, so the original pairs are stored instead. Linear attention with a finite map of dimension applied to queries and keys has capacity . Tsai et al. (2019) develop this kernel view of attention, and Choromanski et al. (2021) approximate with finitely many random features.
6The best a fixed matrix can do
Adding outer products is one way to fill . Is it the best? Suppose all pairs are available and we choose the matrix whose reads are closest to the stored values, measured by the total squared error:
The subscript denotes the Frobenius norm, the square root of the sum of squared entries; the two expressions are the same sum. This is a least squares problem. When the keys are linearly independent, is invertible and the solution is
Every value is recalled exactly, even for keys that overlap. To see what does, write , where is column of . These vectors satisfy
So has dot product one with its own key and zero with every other key. The vectors are called the dual basis of the keys. The least squares memory stores each value under the dual of its key, and since the dual of is perpendicular to , querying with picks up nothing of .
There is a second way to read the same formula. A query in the span of the keys can be written in exactly one way as , and taking the dot product with shows . Hence : the least squares memory is the linear map that sends each key to its value, applied to the coordinates of the query in the basis of keys. The additive memory weights by , a perpendicular drop onto each key separately; the least squares memory weights it by , which is found by walking to the query along the key directions. The two agree exactly when the keys are orthonormal.
Keys and query (drag )
Values and output
Exact recall of overlapping keys has a cost, and Figure 3 shows it. For two unit keys at angle , the dual vectors have length . As the keys approach each other the dual vectors grow without bound, and stretches the direction in which the two keys differ by a large factor. The largest stretch of a matrix, its largest singular value , is the long semi-axis of the ellipse. A query that is slightly off then gives a read that is far from . Least squares separates nearby keys by amplifying the small difference between them, and it amplifies small errors in the query by the same factor.
When the keys are linearly dependent, which is unavoidable once , is not invertible. The minimizer is then , where is the pseudoinverse, and is the best compromise in squared error rather than an exact memory. This optimal linear associative memory was described by Kohonen and Ruohonen (1973). Linear attention cannot use it. Forming requires the dot product between every pair of keys, and a recurrent memory discards each key after writing it.
7How much can a matrix remember?
Section 5 answered this for keys we choose. A trained network does not choose its keys to be perpendicular; they are for whatever tokens arrive, and they overlap. To see what happens then, take the simplest model of overlapping keys: keys drawn independently and uniformly from the unit sphere in , and values drawn independently with mean zero and . We measure the average squared read error , which is on the same scale as a stored value.
Additive memory. The error in reading is the interference term, . Different values are uncorrelated, so the cross terms average to zero and
The expectation can be found without integrating. The distribution of keys does not change under rotations, so we may rotate until is the first coordinate axis. Then is just the first coordinate of . The squared coordinates of a unit vector add up to 1, and by symmetry none of the coordinates is special, so each squared coordinate has mean . With interfering pairs,
The error grows in a straight line from the first extra pair onward. There is no threshold: the interference reaches the size of the signal at , and keeping the error below a fraction allows only about pairs.
Least squares memory. For , random keys are linearly independent with probability one, so . For the reads are , and is the orthogonal projection of onto a -dimensional subspace (the row space of ). Each row of is a vector in holding one coordinate of all values. If the values are Gaussian, this row points in a uniformly random direction independent of the keys, and projecting it onto a -dimensional subspace of keeps, on average, a fraction of its squared length. The lost fraction is the error:
This memory is exact up to exactly pairs and then degrades, reaching half the signal at .
This answers the question in the title for the two memories we have. A matrix in can hold pairs exactly if it is fitted with all the keys in hand. Filled by adding outer products, as linear attention does, it recalls a single pair exactly and holds about pairs at relative error . The value dimension appears in neither answer. Everything a linear associative memory can store is limited by how many directions its keys can occupy.
References
- Anderson, J. A. (1972). A simple neural network generating an interactive memory. Mathematical Biosciences, 14(3–4), 197–220.
- Choromanski, K., Likhosherstov, V., Dohan, D., Song, X., Gane, A., Sarlos, T., et al. (2021). Rethinking attention with Performers. International Conference on Learning Representations.
- Katharopoulos, A., Vyas, A., Pappas, N., & Fleuret, F. (2020). Transformers are RNNs: Fast autoregressive transformers with linear attention. International Conference on Machine Learning.
- Kohonen, T. (1972). Correlation matrix memories. IEEE Transactions on Computers, C-21(4), 353–359.
- Kohonen, T., & Ruohonen, M. (1973). Representation of associated data by matrix operators. IEEE Transactions on Computers, C-22(7), 701–702.
- Schlag, I., Irie, K., & Schmidhuber, J. (2021). Linear transformers are secretly fast weight programmers. International Conference on Machine Learning.
- Tsai, Y.-H. H., Bai, S., Yamada, M., Morency, L.-P., & Salakhutdinov, R. (2019). Transformer dissection: A unified understanding of transformer's attention via the lens of kernel. Empirical Methods in Natural Language Processing.
- Vaswani, A., Shazeer, N., Parmar, N., Uszkoreit, J., Jones, L., Gomez, A. N., Kaiser, Ł., & Polosukhin, I. (2017). Attention is all you need. Advances in Neural Information Processing Systems.