Explainer · interactive · transformers

How self-attention works

Each token builds its new vector as a weighted average of the tokens' values. The weights come from how well its query matches each token's key.

The problem it solves: the vector for "it" is the same in every sentence. In "The animal didn't cross the street because it was too tired", "it" must mix in information from "animal" to know what it means. Self-attention does this with queries, keys, and values that all come from the same sentence.

Level 1 · what it is

One token, attending

Below, four tokens of that sentence have small vectors made by hand, so you can see every number. A key k says what a token offers and a query q what a token looks for, each with two entries: (animate, place). The key of "animal" is (2, 0); the query of "it" is (2, 0), so it looks for something animate.

The score of a token is q · k divided by √d, where d is the number of entries in q and k (here 2). Softmax turns the scores into weights. A value v says what a token passes on; here each value marks one token, so the output of "it" (Σ weight × value) reads as its share of each token.

The query decides what a token finds. The keys did not change; only the question "it" asks did. Try it with the slider.

The animal didn't cross the street because it was too tired.

Remove a part, see what breaks

tokenkey k (animate, place)q · kscoreweightvalue v
Output of "it": Σ weight × value, as shares of each token's value
Go deeper

How, why, and in practice

Level 2 · how

Query · key → softmax → weighted sum

Each token vector x makes three vectors with three learned matrices: a query q = x WQ (what it looks for), a key k = x WK (what it offers), and a value v = x WV (what it passes on). For the token "it":

  1. Take the dot product of its query with every key. The query (2, 0) gives 4, 0, 1, 0.
  2. Divide by √d, where d is the number of entries in q and k. Here d = 2, so the scores are 2.83, 0, 0.71, 0.
  3. Apply softmax to the scores. The weights are 0.808, 0.048, 0.097, 0.048: positive, and they sum to 1.
  4. Add up the values, each times its weight. The output is (0.808, 0.048, 0.048): "it" now carries 81% of the animal's value. The 0.097 that "it" gives itself adds nothing, because its value here is (0, 0, 0).

Every token does the same at the same time, with its own query. The output goes on to the next layer.

Level 3 · why

Why each part is there

derivedWeights from the content

The simplest mix gives every token the same weight. Then "it" gets 0.25 of the animal and 0.25 of the street, and it cannot tell which noun it refers to. Scores from query · key let each token choose.

derivedA query separate from the key

What a token looks for is not what it offers. Without its own WQ, the query of "it" would equal its key, (0.5, 0.5) here, and it matches animal and street equally: q · k is 1 for both, weights 0.313 each, a tie. A separate WQ lets "it" ask for something animate.

factSoftmax, not a plain division

Dividing by the row sum fails once a score is negative: one and minus one sum to zero. Softmax always gives positive weights that sum to 1.

derivedDivide by √d

If the d numbers in a query and a key are independent, with mean 0 and variance 1, their dot product has variance d. Its standard deviation is √d: 8 at d = 64. Softmax then puts almost all the weight on one token. The demo below simulates that growth: it takes four scores, 1.2, −0.3, 0.5, −0.9, as they would be at d = 1, and spreads them by √d, as at dimension d. Not divided, the top weight reaches 0.996 at d = 64; dividing by √d undoes the spread. Then training gets almost no gradient for the other tokens. Dividing by √d keeps the spread near 1: the top weight stays 0.543 at every d.

Weights of four tokens, scores not divided by √d

Level 4 · in practice

Matrices, heads, and the mask

Attention(Q, K, V) = softmax(Q Kᵀ / √d) V

Stack the token vectors as the rows of a matrix X. Then Q = X WQ, K = X WK, V = X WV, and Q Kᵀ holds every score at once: row i is token i's query against every key. Softmax runs along each row.

Heads. A layer runs several heads in parallel, each with its own WQ, WK, WV, so each head can look for something different. Their outputs are joined end to end and projected back. GPT-2 small has 12 heads of d = 64 per layer; 12 × 64 = 768, the number of entries in its token vectors.

Causal mask. A decoder (GPT-style) predicts the next token, so a token may attend only to itself and earlier tokens: when it generates, the later tokens do not exist yet, and in training, seeing them would give away the answer. The mask sets the scores for later tokens to −∞ before softmax, so their weights are 0. "It" comes before "tired": with the mask its weights are 0.848, 0.050, 0.102, and 0 for "tired". Without the mask (an encoder, such as BERT), "it" also gives 0.048 to "tired".

What to trust

What is fact, and what is this page's simplification

factThe formulas.

Query, key, value, softmax(Q Kᵀ / √d) V, heads, and the mask are the transformer's attention (Vaswani et al., 2017).

constructedEvery vector on this page.

Made by hand: keys and queries with d = 2 named entries (animate, place), values that mark one token each. A trained model learns its vectors, and its dimensions have no names.

interpretation"A query is what a token looks for."

It describes what the math allows. A trained head need not use it that way.

observationSome heads track a clear relation.

In trained models, some heads attend to the previous token, or from a pronoun to its noun. Many heads have no simple description.

In short

Summary

  1. Each token's output is a weighted average of the values of the tokens it can see.
  2. The weights are softmax(query · key / √d): the query decides what a token finds.
  3. Separate queries and keys let a token look for something other than itself; √d keeps softmax from putting all the weight on one token.
  4. So "it" ends up carrying 81% of the animal's value: its query asks for something animate, and the animal's key offers it.

Not covered: how WQ, WK, WV are learned; how the model knows token order (positional encodings); the MLP, residual connections, and layer norm.