How attention worksLisa Dlima ↗colour keyX inputW weightsV valueQ queryK keyA attentionO output

How attention works

By Lisa Dlima
When I took the Georgia Tech masters in CSci course on deep learning we briefly covered LLMs and how they are built by attention blocks. We trained a model in PyTorch and needed to manipulate some special Q, K, and V matrices the right way to achieve convergence. But it wasn't that intuitive and I find myself constantly forgetting how it works. Here is my opportunity to learn it again and distill it in a way where you could rediscover it by intuition whenever you need to.I hope you have some base information on neural networks. If not that's okay you will just have to accept one premise: NNs almost all involve operations where some input data x is matrix multiplied by a weight matrix W to create a new projection of x.x' = WxFor example, imagine we have $x$ represented by an image of 128 flattened numbers. If we multiply it by $W$ which might be $128\times256$ then we project that $x$ into a new vector of length $256$.x0.81.00.30.50.70.3×W0.34-0.33-0.01-0.06-0.200.220.330.15-0.30-0.480.28-0.090.300.20-0.39-0.09-0.030.200.28-0.390.18-0.19-0.33-0.040.06-0.350.410.200.300.37-0.170.160.26-0.350.14-0.38=-0.40.0-0.10.3-0.4-0.3Sometimes you will see this represented by a set of perceptrons densely connected to each other.
0.81.00.30.50.70.3-0.40.0-0.10.3-0.4-0.3
Normally we would apply some activation function nonlinearity at the end, but we will ignore this. Attention is able to bake in its nonlinearity differently.If we want to compute more images simultaneously (likely on a GPU) we can feed more than one x into the weight matrix by just creating an X with more rows. Note that adding a new x2 does not change the projection results from x1. They are independent as they should be. You wouldn't want the processing of one image to affect the result of the other.But lets change the input data. Instead of X's rows being an unordered set of images, now they are an ordered set of tokens representing words.Xthe0.30.30.00.40.30.9cat0.81.00.30.50.70.3sat0.40.00.30.80.70.7on0.30.50.10.80.60.5the0.30.30.00.40.30.9mat0.50.70.70.90.50.4×W0.34-0.33-0.01-0.06-0.200.220.330.15-0.30-0.480.28-0.090.300.20-0.39-0.09-0.030.200.28-0.390.18-0.19-0.33-0.040.06-0.350.410.200.300.37-0.170.160.26-0.350.14-0.38=Vthe-0.30.30.30.3-0.3-0.3cat-0.40.0-0.10.3-0.4-0.3sat-0.10.30.30.4-0.30.2on-0.10.10.20.3-0.3-0.0the-0.30.30.30.3-0.3-0.3mat-0.2-0.1-0.20.3-0.5-0.2We will call this output matrix V, not because it is particularly special, but just because it is a fine variable name consistent with the attention paper.Now imagine we wanted to predict the next token as in a chatbot. Right now we process each token independently. If our model is to learn grammar, ordering, and complex nuance between words we need a way to intelligently mix these outputs.Here is one concrete example: We could take V and set each row to be some normalized combination of the rows of V. For example the last row could be a uniform combination of all rows in V. Keep in mind though we are building a chatbot so we don't want to leak information from a future token to a past token. So our rule will be that we can only mix rows with itself and earlier rows. One consequence of this is that the first row will always be unchanged because it is a mixture only with itself.Vthe-0.30.30.30.3-0.3-0.3cat-0.40.0-0.10.3-0.4-0.3sat-0.10.30.30.4-0.30.2on-0.10.10.20.3-0.3-0.0the-0.30.30.30.3-0.3-0.3mat-0.2-0.1-0.20.3-0.5-0.2mixedthe-0.30.30.30.3-0.3-0.3cat-0.30.20.10.3-0.4-0.3sat-0.30.20.20.3-0.3-0.1on-0.20.20.20.3-0.3-0.1the-0.20.20.20.3-0.3-0.1mat-0.20.10.10.3-0.4-0.1This uniform mixing is just one choice. Naturally, we will want the model to learn what the optimal distribution is. To do this we will need to learn an $n\times n$ matrix called $A$, with $n$ being the size of the sentence.A1.00000000.500.5000000.330.330.330000.250.250.250.25000.200.200.200.200.2000.170.170.170.170.170.17×Vthe-0.30.30.30.3-0.3-0.3cat-0.40.0-0.10.3-0.4-0.3sat-0.10.30.30.4-0.30.2on-0.10.10.20.3-0.3-0.0the-0.30.30.30.3-0.3-0.3mat-0.2-0.1-0.20.3-0.5-0.2=mixedthe-0.30.30.30.3-0.3-0.3cat-0.30.20.10.3-0.4-0.3sat-0.30.20.20.3-0.3-0.1on-0.20.20.20.3-0.3-0.1the-0.20.20.20.3-0.3-0.1mat-0.20.10.10.3-0.4-0.1This is how an attention module does it: Take $X$ and create two weight matrices $W_q$ and $W_k$. Compute $XW_q$ and $XW_k$ just like we did for $XW_v$. This will produce two new matrices $Q$ and $K$. Just like $V$, $Q$ and $K$ are just new, unique projections of $X$. Exactly the same as before. In the attention is all you need paper these are named for Key, Value, and Query. I found this more confusing than helpful. Just remember they are created the same way with the only difference being the different weights.To create an $n\times n$ matrix all we have to do is matrix multiply $Q$ by $K^\top$ (the transpose of $K$, so $n\times d$ times $d\times n$ gives an $n\times n$ result).
Xthe0.30.30.00.40.30.9cat0.81.00.30.50.70.3sat0.40.00.30.80.70.7on0.30.50.10.80.60.5the0.30.30.00.40.30.9mat0.50.70.70.90.50.4×Wq0.490.350.32-0.060.010.01-0.440.190.460.210.180.26-0.100.25-0.210.480.050.200.350.34-0.46-0.31-0.23-0.180.48-0.02-0.10-0.21-0.220.430.490.140.36-0.340.08-0.44=Qthe-0.30.20.10.60.20.5cat-0.50.30.10.80.50.3sat-0.0-0.4-0.00.40.50.6on-0.2-0.0-0.20.40.40.5the-0.30.20.10.60.20.5mat-0.4-0.10.10.80.60.8Xthe0.30.30.00.40.30.9cat0.81.00.30.50.70.3sat0.40.00.30.80.70.7on0.30.50.10.80.60.5the0.30.30.00.40.30.9mat0.50.70.70.90.50.4×Wk0.20-0.24-0.12-0.450.18-0.48-0.060.030.110.370.12-0.13-0.21-0.12-0.13-0.170.370.380.16-0.33-0.030.26-0.05-0.450.220.130.32-0.37-0.48-0.080.010.220.160.380.43-0.15=Kthe0.2-0.10.50.20.30.5cat0.1-0.40.9-0.40.40.0sat0.4-0.10.30.20.60.5on0.2-0.30.5-0.00.20.5the0.2-0.10.50.20.30.5mat0.1-0.20.6-0.20.30.2Qthe-0.30.20.10.60.20.5cat-0.50.30.10.80.50.3sat-0.0-0.4-0.00.40.50.6on-0.2-0.0-0.20.40.40.5the-0.30.20.10.60.20.5mat-0.4-0.10.10.80.60.8×Kᵀ
thecatsatonthemat
0.20.10.40.20.20.1-0.1-0.4-0.1-0.3-0.1-0.20.50.90.30.50.50.60.2-0.40.2-0.00.2-0.20.30.40.60.20.30.30.50.00.50.50.50.2
=A0.4-0.20.40.20.40.00.4-0.20.40.10.4-0.00.50.20.70.50.50.30.3-0.20.50.20.30.00.4-0.20.40.20.40.00.70.10.80.50.70.3
A matrix multiplication is just a lot of dot products, in this case dot products between every pair of $K$/$Q$ representations of our original tokens. The value of a dot product is higher the more similar a vector $k$ is to $q$. Therefore this $n\times n$ matrix is some measure of similarity. Our hope is that as our model learns this measure of similarity becomes valuable.But remember we only want information from future tokens to mix with past tokens, so first we mask; every pair where a token attends to a future token is set to $-\infty$, blanking the top-right of the matrix.A0.4−∞−∞−∞−∞−∞0.4-0.2−∞−∞−∞−∞0.50.20.7−∞−∞−∞0.3-0.20.50.2−∞−∞0.4-0.20.40.20.4−∞0.70.10.80.50.70.3Then we convert each row into probabilities with a row-wise softmax. The $-\infty$ entries become exactly zero, so each row sums to one.A1.00000000.650.3500000.350.250.410000.280.160.320.25000.230.130.230.180.2300.200.100.220.160.200.12Now that we finally have our A matrix we can apply it to our V matrix and get the final result.A1.000.000.000.000.000.000.650.350.000.000.000.000.350.250.410.000.000.000.280.160.320.250.000.000.230.130.230.180.230.000.200.100.220.160.200.12×Vthe-0.30.30.30.3-0.3-0.3cat-0.40.0-0.10.3-0.4-0.3sat-0.10.30.30.4-0.30.2on-0.10.10.20.3-0.3-0.0the-0.30.30.30.3-0.3-0.3mat-0.2-0.1-0.20.3-0.5-0.2=Othe-0.30.30.30.3-0.3-0.3cat-0.30.20.20.3-0.4-0.3sat-0.30.20.20.3-0.3-0.1on-0.20.20.20.3-0.3-0.1the-0.20.20.20.3-0.3-0.1mat-0.20.20.20.3-0.3-0.1O now represents each token as a mixture of the ones in V. In one round of attention we found similarities between pairs of token projections. But now that we have these mixed token representations the more we pass them through attention layers the more they mix and the more likely we'll find interesting relationships fall out of our data.By the way we didn't need a nonlinearity like a ReLU because we used a softmax during mixing. We also only needed to learn 3 weight matrices Wk, Wq, and Wv. (Actually you can have a Wo at the very end to project O one more time)In general LLMs are just variations on stacking these modules in a very deep neural network. You can stack them vertically, you can also do some tricks to stack them horizontally. You may wonder how we will predict the next token for our sentence. In practice you take the last row of the O matrix and send it through another weight matrix to predict the next token. But you can do other things too. You can translate, predict sentiment, or anything you think of. Notice how subtle architectural decisions matter for different use cases. For example our masking was useful for chatbot next token prediction, but if you were to translate a whole sentence it would make more sense to not mask at all because it is very helpful to mix past tokens with future ones.We will see later that this architecture creates new challenges. In particular when new tokens are generated the size of the K/V matrices will grow leading to new challenges for training and inference.
An interactive walkthrough by Lisa DlimaAttention Is All You Need ↗