Breaking down "Attention Is All You Need"
and the Transformer
The transformer is a model architecture that essentially created the entire AI boom we are seeing now. Despite countless innovations in the field and model, all of it dates back to August 2nd 2023 in the heart of Google's DeepMind Lab.
The way a transformer works is by predicting the next token. So for a simple example, imagine you have the sentence "I just dropped my mechanical pencil I can't believe my lead ___"
The fill in the blank here is obviously "broke." and we can think of that as a token that the model is going to predict.
With this slight introduction I want to spend the rest of this blog breaking down the self-attention mechanism and the encoder and decoder design. I assume basic familiarity with ChatGPT and the general idea of what a transformer is, but not the math and exact implementation we will go into. For more preliminary information there are resources at the bottom.
So before the transformer and self-attention we had a very sequential based approach (RNN, LSTM, etc). In this approach we would take each token (think of a token as a word in a sentence), feed it through a model, store some data in memory, and then keep doing that with every next token. Now this works and was a solid approach to this problem, however it wasn't scalable and was unable to handle large context windows. Imagine a 10 page essay. Training would all be sequential so it would take forever to train the model and also with this kind of an approach by the time you go to the end of the essay the model would have no clue what happened in the very start. As you might have seen the scale of LLM's has gone up an absurd amount in the past few years. What we are finding is that scale eventually trumps everything, even in adjacent fields such as robotics. This means that we needed a parallelizable approach. Luckily the attention mechanism was not only parallelizable but also tacked on the benefit of handling large context windows by relating every token to each other (I will explain this part more). Now with all of this said, this blog will hone in on that attention mechanism and how it works under the hood.
Take the input example we had from before: "I just dropped my mechanical pencil I can't believe my lead ___". The first step is to tokenize this (try an actual tokenizer).
The model can now take each of these tokens and look it up in a massive dictionary to get its vector representation. As a mental model, imagine each token starting as a giant sparse vector with around 50,000 possible slots. Almost everything is 0, and one position lights up to say which token it is. Let's get some example vectors for our tokens.
You can think of maybe the 43,403 slot to be pencil and if we have the word pencil in our sentence that slot would be 1 and everything else 0. In practice this is not how we represent tokens and it is more than just 1's and 0's but it is a good start. Now another aspect I want to break down is the intuition for how these vectors operate. Think of a toy 2D space where the x axis is gender-ish meaning and the y axis is occupation. Then similar words land near each other, and meaningful differences become directions you can move in.
Now these are again just very basic representations of a much more complex system underneath but this should get you thinking about how these 50,000-dimensional vectors can start to hold some value.
Once every token has a vector, attention compares every token to every other token. The table below is a toy version of that idea: the same tokens go across the top and down the side, and each blob is the attention score between that pair.
Let me break this down a bit, since the idea can be easy to miss. We take all of these vectors and multiply them against one another using a dot product. Tokens that relate to one another will tend to have a higher dot product, which means they can attend to one another.
We cannot just multiply two raw vectors together, because that would not tell us enough about what relationship we want to measure. We need learned transformations in between. This brings us to the query and key matrices. We take each token vector in a sentence and multiply it by a query matrix. We then take those same token vectors and multiply them by a key matrix. You can think of the query matrix as asking a question, or a collection of questions embedded in a vector, while the key matrix provides the corresponding answers. After extensive training on a huge amount of text, the model learns patterns about language and embeds them in high-dimensional vectors. For example, a query might be trained to ask, "Are there any adjectives that describe me?" when it is multiplied by a noun, while a matching key might answer, "I am an adjective, and I describe you!" When a key matches a query in that high-dimensional space, their dot product becomes larger.
rows: keys K โข columns: queries Q
green pairs share a feature, so their Q ยท K score is higher
mirrored query-key pairs have the same score in this toy example
each complete query column is normalized to add up to 1.00
As you can see in this table, the words pencil and lead attend to one another, as do I and dropped, and my and pencil. This lets us look at the table and see which words relate to one another and which words are important to each other. This is only possible after we learn the query and key matrices. Another implementation detail as you can see in the italizes above, the values you get after the dot product of the vectors is a large number and we apply a softmax to bring all of these values below 1 and all the values for a certain word will add up to one.
Now just to be very very clear. Go back to the table above (the closest one). Each word is a token vector. So "dropped" is a vector that embeds its meaning (this is important because soon the vectors that represent each token will change!). We take "dropped" and "I" and "my" and all the other tokens and multiple them by a Q matrix which is our query matrix. After each token embedding (e1, e2, e3, etc) is multiplied by the matrix (also this just occured to me that you must have a good understanding of lin alg for this!!) we get our Q_i vectors. Extend this process out to our key matrix and boom we now have two vectors that represent the Query and the Key for each token embedding. Again I think the reasoning behind how this works and why this works is BIG DATA, I'm not too deep into LLMs so take that aspect with a grain of salt but if you need some reasoning for why we are doing. Anyways moving on we now take the dot product of these two vectors and that tells us how much they relate to one another. Backing up as to why, two vectors that are same will have large dot products and the further away they are the smaller or even more negative their dot products will become. We take the whole grid with all of its dot products and then normalize it to a scale between 0-1 such that each column adds up to 1. Now why would we do that? Think about it this way. We now have for any single token a list going down of how much each other token matters to it (or attends to it in LLM world). So for the token "pencil" we can see that "I" attends 0.09, "dropped" is 0.12, "my" is 0.29, "pencil" is 0.10, and "lead" is 0.40. Ok so what does that tell us. That tells us that "lead" is the most important in this context. Cool now we can move onto the next step which is creating the Value vector for each token!