Self-Attention Primer
Self-attention drawn small enough to see: one six-token sentence, two dimensions per vector, and one plane that every query, key and value lives on. The score matrix, the √d divisor, softmax, the causal mask and the weighted sum — every number on the page is computed by the figure beside it, including the one that shows the mask being applied wrongly.
Three vectors out of one
A token arrives as a single vector. Attention's first move is to split it into a query, a key and a value.
The sentence is the cat ate and it ran, and the page follows one token through it: the pronoun it. Each token is drawn as a two-dimensional vector rather than the 64 a real head uses, so that every quantity here is something you can point at on a plane. Nothing in the argument depends on the width.
Three learned matrices turn that one vector into three: W_Q gives the token its query, W_K its key, and W_V its value. Install them one at a time with the slider:
Notice that the three arrows leave in three different directions from one input. The query is what this token is looking for, the key is what it advertises to everyone else, and the value is what it hands over if anyone takes it up. The names are borrowed from retrieval, and the analogy survives the whole page.
The three matrices are fixed for the layer, so the three arrows are a rigid function of the input: the query, the key and the value all move when the token does. Drag the grey input vector anywhere on the plane:
Watch the query swing hard as you push the input upward while the value barely turns. That is what three matrices buy: what a word looks for, what it advertises and what it contributes are three different functions of it, and one vector cannot carry all three without making them the same job.
Every other token has a key of its own in the same plane. Step through the six and watch which key each query throws its longest shadow onto:
At the winner is the key of cat, at 1.72 against 1.14 for it's own key. Nobody told the model that it refers to cat: gradient descent shaped W_Q and W_K until pronoun queries lean towards noun keys, because that is what lowered the loss.
If the two maps were one map, a query would be its own key and asking would be the same as being asked. Blend W_Q into W_K and watch the two directions fall together:
Untouched, cat scores ate at 0.51 while ate scores cat at −0.32: the verb wants its subject far more than the subject wants its verb. At both read 1.16 and the relation has become a similarity. Language is not symmetric, and keeping the two maps apart is what buys the direction.
A score is an alignment
Two vectors in, one number out. That number is how far the query leans along the key.
The arithmetic is a dot product: multiply matching slots, add them up. The geometry is a shadow. Drop a perpendicular from the query's tip onto the key's line, and the score is the length of that shadow times the length of the key. Drag the query and watch the number:
Notice where the sign turns over. The score is positive while the query leans the same way as the key, exactly zero at a right angle, and negative once it leans away — and a negative score is not an error, it is a token actively voting the other way.
One query is asked against all six keys at once, which turns the plane into a row of six numbers. Step the query through the sentence and read the row off the columns:
Because the six scores are raw dot products they are unbounded in both directions: peaks at 3.02 on its own key, while runs to −0.45 on cat. Nothing yet says which of them is large.
And that is the trap in a raw dot product: it is a product of two lengths, so a key can win by being long rather than by pointing anywhere useful. Stretch the key of and and watch it climb:
At the function word and outscores cat without having turned one degree. Real Transformers keep this in check with LayerNorm before the projections, and some recent models normalise q and k themselves — QK-norm — precisely because a single overlong key can capture every row of the matrix.
Every query meets every key
One row per query, one column per key. That square is the whole of Q · Kᵀ.
There is no search and no shortlist. Every token's query is dotted against every token's key, including its own, and the results are laid out as a square whose side is the length of the sentence. Fill it in one row at a time:
Each square's side is the size of the score, so the picture is read without reading a digit. Six tokens cost ; a thousand-token context costs a million of them, per head, per layer. That is the quadratic everyone complains about — though in GPT-2 small the four projections still cost more than it until n passes 2·d_model, which is 1,536 tokens.
A row is one query against the whole sentence — the six numbers §02 read off the columns, now stacked with everyone else's. Walk down the rows:
has its largest square under cat at 1.22, and is nearly flat: a determiner has nothing in particular to look for, and an almost-flat row is how attention says I have no preference.
It is tempting to read the square as a similarity table, and it is not one. Flip it about the diagonal and watch the picture change:
Because W_Qᵀ W_K is not symmetric, cat → ate at 0.51 and ate → cat at −0.32 are different numbers about the same pair. The verb reaches for its subject; the subject votes against the verb. A similarity matrix could not tell those apart, and that asymmetry is the whole reason there are two maps and not one.
From scores to a mixture
Six unbounded numbers go in. Six positive weights that add up to one come out.
Softmax exponentiates each score and divides by the total. Exponentiating makes every entry positive whatever the sign of the score; dividing by the total makes the row sum to one. The teal staircase adds the weights up left to right and has to land on the top edge:
Notice how soft the mixture is. In the largest weight is 0.30 on cat and the smallest is 0.11 on the — a factor of three, not a lookup. Attention almost never picks one token; it pours a little of everything and more of something.
The six are not six independent decisions. They share one denominator, so raising any one score must lower every other weight. Drag cat's score along the row and watch the other five give way:
Watch the readout on the right: whatever you do to one score, the sum stays 1.000. That is the invariant of the whole mechanism, and it is worth writing as an assertion — all(w > 0) and sum(w) == 1 at every row, which is what makes the output a weighted average rather than an arbitrary linear combination.
The floor of that assertion is strict. A score can be pushed as far down as you like and the weight will shrink towards zero without ever arriving. Push it and read the exponent:
At the weight is 1.2e-14 — invisible in the bar, still positive in the arithmetic. That is why a mask has to use −∞ rather than a very negative number, and it is also why softmax is safe to implement as exp(s − max(s)): the shift cancels in the ratio, so the stable version is the same function, not an approximation of it.
Why there is a √d in the formula
The divisor is not a fudge factor. It is the spread of the thing being divided.
A dot product of two d-dimensional vectors is a sum of d products. If the components are independent with mean 0 and variance 1, each product has variance 1 and the sum has variance d — so the spread of a score grows as √d. Slide the width and read it off the curve:
The width axis is logarithmic — each tick is a four-fold widening — and the spread still doubles from one tick to the next, so the curve bends upward even here. The amber rule at 1 is where the spread would have to sit for softmax to behave. At the spread is 8.00 — eight times too wide.
Eight times too wide does not mean eight times noisier. Softmax is an exponential, so multiplying every score by eight raises the biggest one to the eighth power against its rivals. Push the width up and watch the mixture collapse:
By one weight already holds 0.991 of the mixture, and by 64 it holds all of it. That is worse than an arbitrary choice: the gradient of softmax through a weight is proportional to w(1 − w), so a row pinned at 1.000 passes almost nothing back and stops learning which token it should have chosen.
Dividing every score by √d puts the spread back at exactly 1 whatever the width. Run the same slider with the divisor in place:
The mixture does not move: 0.572 at every width from 1 to 256, because the scale that was growing has been divided out. Dividing by d instead would over-correct and squeeze the spread to 1/√d, flattening every row towards a uniform average. Our toy head has d = 2, so its divisor is 1.414; GPT-2 small runs d = 64 per head and divides by exactly 8.
The mask, and the way it is usually broken
A language model is trained to predict the next token. Nothing in attention stops it reading that token first.
The square from §03 lets position 0 attend to position 5, which during training is a token the model is being asked to guess. A decoder therefore deletes everything above the diagonal before the softmax. Take the rows out one at a time and watch the future go:
The dashed squares are the scores that were computed and then thrown away — , which is why an unfused implementation does 36 dot products to keep 21. The mask is fixed, has no parameters, and is the only part of the whole mechanism that knows anything about position.
What survives is renormalised: the row still has to sum to one, so the tokens that remain simply divide the whole mixture between them. Step the asking position down the sentence:
At the mixture is 1.00 on the — the first token has exactly one candidate, so its output is its own value vector whatever its query says. At the weights are spread over five, with 0.34 on cat.
Now the mistake. A 0/1 mask is a matrix, and multiplying by it looks like the obvious way to apply it. Switch the control from + (−∞) to × 0:
Because a masked score becomes 0 rather than −∞, and exp(0) = 1, ran keeps 0.092 of a mixture it is supposed to be absent from — the model is reading the answer. Nothing throws, no shape is wrong, and the training loss falls faster than usual, which is the tell. The correct line is s = s.masked_fill(m, float("-inf")), and it is worth an assertion that the masked entries of the weight matrix are exactly zero.
The sum, and what it costs
The weights were never the answer. They are the proportions in which the value vectors get poured together.
The last step is one line: multiply each value vector by its weight and add the six together. Because each weight is a fraction of one, each term is a short step in that value's direction. Add the terms one at a time:
Notice that the chain never doubles back, because there are no negative weights to reverse it. Six steps later is the new representation of it: [0.82, 0.81], a token that has been told something about the cat.
That is also the geometric form of §04's assertion. Positive weights summing to one means the output is a convex combination, so it lands inside the polygon the values span. Drag cat's score as hard as you like and try to get out:
You cannot. and the output stops at [1.21, 0.93], a hundredth short of cat's own value — §04's strictly positive floor is exactly why it can approach a corner and never arrive at one. One attention head can only ever return an average of what the sentence already contains — which is why the layer that follows it is a non-linear MLP, and why a stack of attention layers with no MLPs collapses to something close to a single linear map.
And now the thing the mechanism cannot do. The sum runs over a set, so reordering the tokens reorders the weights and nothing else. Shuffle the reading order and keep an eye on the output:
The bars rearrange and [0.82, 0.81] does not move. Self-attention is permutation-equivariant: with no positional information it cannot tell the cat ate and it ran from any shuffle of it, and it fails at this silently — the model trains, the loss falls, and word order simply never enters the representation. Positional encodings exist to break this symmetry, and the causal mask above is the only other thing on the page that knows about order.
The bill is one dot product per ordered pair, twice: once for Q · Kᵀ and once for the mix. That is 2n²d multiply-adds and, if the matrix is materialised, n² scores of memory against the n · d the tokens themselves occupy. Slide the context length:
At one head's score matrix is 2.0 MiB in fp16 — 288 MiB across 12 layers and 12 heads. At it is 32.0 GiB while the tokens it came from are 192.0 MiB: 171 times smaller. FlashAttention exists to never write those bytes down.
Four lines, and what they draw
Every idea on this page is one of these four lines.
s = q @ k.transpose(-2, -1) / d_k**0.5
s = s.masked_fill(mask, float("-inf"))
w = s.softmax(dim=-1)
out = w @ vRun them over the whole sentence and the result is one picture: six rows of weights, each row a query's mixture over the keys, with the future cut away. Switch the mask off and watch the upper triangle come back:
This is the attention map every paper prints, and it is now readable cell by cell: it takes 0.34 of its mixture from cat under the causal mask and 0.30 without it. Four ways to break it, none of which raises: a mask multiplied instead of added, no positional information, the √d divisor dropped, and keys left unnormalised.