Attention and the transformer, explained from scratch
How each token looks at the others to update its meaning, and how stacking that idea builds a transformer.
Por ahora, este artículo solo está disponible en inglés. El resto del sitio está traducido.
By the time a language model starts thinking about your text, each token has been swapped for a list of numbers, its embedding(see What are embeddings?). But that first vector comes from a lookup table. It is the same every time the token appears, whatever the sentence around it. This article explains the idea that fixes that, attention, and how it is packaged into the transformer, the design behind almost every modern language model.
Why context matters
Look at the word “bank” in these two sentences.
A lookup table gives “bank” one vector in both. A person reads “river” and knows at once which bank is meant. Pronouns are even more dependent on context: in “The animal didn’t cross the street because it was too tired”, the word “it” means nothing until you work out that it points back to “the animal”.
So a model needs a way for each token to gather information from other tokens and update its own vector. That is exactly what attention does.
Attention: every token asks a question
Here is the intuition. Each token asks, “which of the other tokens matter to me?” It gives every token it can see a score, turns the scores into percentages, and then takes a blend of those tokens’ information, mixing in more from the ones with higher percentages. These percentages are called attention weights.
- The6%
- animal62%
- didn't0%
- cross0%
- the0%
- street14%
- because6%
- it12%
- washidden
- toohidden
- tiredhidden
“it” puts 62% of its attention on “animal”. Later tokens are hidden from it.
With “it” selected, most of the weight goes to “animal”. After this step, the vector for “it” carries some of the meaning of “animal”, which is what later parts of the model need. Nobody writes these weights by hand in a real model. They are computed from the vectors themselves, using numbers the model learned during training.
Queries, keys and values
How does a token decide who matters? Think of a library search. You type a query (what you are looking for). Each book has a key, like the label on its spine (what it is about). You compare your query with every label, and the better the match, the more you read from that book’s contents, its value.
In a transformer, every token plays all three roles. Its vector is multiplied by three different learned matrices to produce three new vectors:
- query: what this token is looking for;
- key: what this token offers, so others can find it;
- value: the information it hands over if someone attends to it.
Then, for the token doing the looking, the recipe has four steps. Compare its query with every key using a dot product: multiply the two vectors number by number and add up the results. Similar directions give a large number. Scale the scores down by the square root of the vector length. Softmax them: a function that turns any list of numbers into positive numbers adding up to 1. Finally, blend the values, each multiplied by its weight.
| token | key / value | q · k | weight |
|---|---|---|---|
| the | k[0.10, -0.30]v[0.10, 0.10] | · | · |
| animal | k[1.20, 0.90]v[0.90, 0.20] | · | · |
| it | k[0.60, 0.40]v[0.30, 0.60] | · | · |
| output | · | ||
The token “it” has a query [1.00, 1.50]: what it is looking for. Every token it may look at (including itself) has a key, what it offers, and a value, what it will hand over.
Written compactly, for all tokens at once, this is the formula from the 2017 paper that introduced the transformer: softmax(QKᵀ / √d) · V. Q, K and V are the queries, keys and values of every token stacked into tables, and d is the length of a key vector.
Causal masking: no peeking ahead
Text generators like GPT write one token at a time, left to right (From the last vector to the next word covers that loop). During training, the model learns to predict each next token of a text. If a token could see the words after it, it could simply copy the answer.
So these models use a causal mask: before the softmax, the scores for every later token are set to minus infinity, which softmax turns into a weight of exactly zero. Each token can look at itself and everything before it, never after. Drawn as a grid, the weights form a triangle.
| looking token | The | animal | didn't | cross | the | street | because | it | was | too | tired |
|---|---|---|---|---|---|---|---|---|---|---|---|
| The | |||||||||||
| animal | |||||||||||
| didn't | |||||||||||
| cross | |||||||||||
| the | |||||||||||
| street | |||||||||||
| because | |||||||||||
| it | |||||||||||
| was | |||||||||||
| too | |||||||||||
| tired |
Not every transformer is causal. Models built to understand text rather than write it, such as BERT, let every token see the whole sentence in both directions.
Many heads at once
One set of weights can only express one kind of relationship at a time. But a token might care about several things: which noun a pronoun refers to, which verb goes with a subject, what the previous word was.
So a transformer runs several attention computations side by side, each with its own query, key and value matrices. Each one is called a head. Their outputs are joined together and mixed back into a single vector. The original transformer used 8 heads per layer; GPT-2’s smallest version uses 12.
The transformer block
Attention lets tokens share information. A transformer wraps it in a block with three more ingredients:
- A feed-forward network. After attention has mixed information between tokens, each token’s vector goes through a small neural network on its own: the vector is expanded to a longer one (four times longer in the original design), passed through a simple non-linear function, and shrunk back. This is where much of the model’s stored knowledge is thought to live.
- Residual connections. Instead of replacing a token’s vector, each sub-layer’s output is added to it. The vector becomes a running record that every layer edits a little. This also makes very deep stacks much easier to train.
- Normalisation. Before each sub-layer, the vector is rescaled to a standard size so numbers neither explode nor fade as they pass through many layers.
Stacking blocks
A model is many of these blocks stacked on top of each other, each with its own learned numbers. The original transformer had 6 in each half; GPT-2’s smallest version has 12, and large models have dozens more. Every block reads the vectors the previous one produced and refines them. Early layers tend to handle local, surface patterns; later ones build more abstract features.
The shape never changes along the way: one vector per token goes in, one vector per token comes out. At the very top, the vector of the last token is turned into a guess about the next token.
Where is word order?
There is a catch. Attention on its own ignores order: it compares every token with every other, so “dog bites man” and “man bites dog” would look the same. Transformers fix this by giving each token positional information. The original paper added a fixed pattern of sine waves to each embedding, a different pattern for each position. GPT-2 learns a position vector instead. Many newer models rotate the query and key vectors by an angle that depends on position (a method called RoPE), so attention scores reflect how far apart two tokens are.
Further reading
- Attention in transformers, step-by-step3Blue1Brown · video, the best visual walk-through
- Transformers, the tech behind LLMs3Blue1Brown · video, the big picture
- The Illustrated TransformerJay Alammar
- Attention Is All You NeedVaswani et al., 2017 · the original paper
