Level 14 · Theory · runs in your browser

Self-attention

How does one word “look at” the other words?

After level 13, every word is a list of numbers. But the list for “bank” is the same in “river bank” and “bank account”. A word needs to use information from the words around it. Attention is how it does that: every word builds a new vector by mixing all the words in the sentence, in proportions it computes itself.

This is the most important part of the Transformer. This level builds it in small steps, always with three words and two numbers per word.

1. Why not just read left to right?

Before attention, the usual way to read a sentence was a recurrent network (RNN). It reads one word at a time and keeps a single summary of everything so far, the hidden state h. Each new word xtx_t updates it: ht=tanh⁡(w ht−1+xt)h_t = \tanh(w\,h_{t-1} + x_t) (a one-number version of the RNN in side trip N3).

That design has two problems:

  1. Fixed size. By the 50th word, everything about the 49 words before it has to fit into the same few numbers.
  2. Early words get weaker and weaker. The effect of the first word reaches the end through every later step, and each step multiplies it by w (and by the slope of tanh, which is never above 1). So when the model gets the end of a sentence wrong, almost none of the correction reaches the first word, and training can’t learn to use it.
Number In an RNN with w = 0.5, ignoring tanh, each later word multiplies word 1’s effect on the hidden state by 0.5. After 3 more words, what is word 1’s effect multiplied by?
🔒 Answer the question above to unlock

Here is how fast it shrinks with w = 0.5:

words laterthe first word’s effect is multiplied by
10.5
20.25
50.0312
100.000977
508.88 × 10⁻¹⁶

After a few dozen words, the effect of the start of the sentence is almost zero. Attention does the opposite: every word looks at every other word directly, in one step, whatever the distance. For the 50th word, reading the first word is exactly as easy as reading the 49th. The rest of this level builds that.

The Classic networks branch builds an RNN and an LSTM (levels N3 and N4) and shows this shrinking in full. Side trip N5 measures the problem on a real task: a summary of 32 numbers reverses 0% of 20-digit strings. N5 also computes a first attention with a single question vector. This level adds the matrices W_Q, W_K, W_V and lets every word ask.

2. A weighted average

Three words: cat = [2, 0], dog = [1, 1], car = [0, 2]. Attention does three things:

  1. Scores. Dot every word with every word: scores = X @ X.T. A big dot product means “these two point in the same direction”. The lab then divides every score by √2. Take that as a fixed scaling for now; section 5 explains why.
  2. Weights. Turn each row of scores into proportions that add up to 1, with softmax.
  3. Output. Each word’s new vector is the weighted average of all the words: output = weights @ X.

Drag the dots. Every table below changes with them.

Each word’s output is a weighted average of all the words

Drag a word’s dot (or focus it and use the arrow keys). Point at or tap a weight to see who looks at whom.

12xycatdogcar
how much cat looks at each word cat′ = cat’s output, the weighted mix (dashes: each word’s weight)
Draw the output of
X (3 words × 2 numbers)
xy
cat
dog
car
↓ dot every word with every word, divide by √2
scores = X Xᵀ / √2
catdogcar
cat2.83?0
dog?1.411.41
car01.412.83
↓ softmax each row, so it adds up to 1
weights
catdogcar
cat0.770.190.05
dog0.330.330.33
car0.050.190.77
↓ mix the word vectors with those weights
output = weights @ X
xy
cat?0.28
dog11
car0.281.72
cat’s output = 0.77×cat + 0.19×dog + 0.05×car = [?, 0.28] Cat’s output appears in the picture once you compute its x below.
Number Before dividing by √2, what is the score of cat looking at dog? (cat = [2, 0], dog = [1, 1])
🔒 Answer the question above to unlock

Softmax turns a row of scores into weights. It takes e to the power of each score and divides by the total, so big scores get most of the weight, and every weight is positive. Cat’s row of scores is [2.83, 1.41, 0] after scaling, which becomes the weights [0.768, 0.187, 0.045].

Number cat = [2, 0], dog = [1, 1], car = [0, 2]. Cat’s attention weights over (cat, dog, car) are [0.768, 0.187, 0.045]. What is the first number (x) of cat’s output? (2 decimals)
ChooseIn an attention weight table for the words cat, dog, car (in that order), rows are words and columns are words. weights[0][2] = 0.045. What does that number say?
ChooseThree words attend to each other with scores = X @ X.T / √2, then softmax: cat = [2, 0], dog = [1, 1], car = [0, 2]. If car becomes [5, 5] (cat and dog stay the same), what happens to car’s row of weights?
I got stuck here Why is score(cat, dog) always the same as score(dog, cat)?

Because both come from the same dot product: cat · dog = dog · cat. When every word uses its own vector to ask and to answer, the score table is symmetric: the same on both sides of the diagonal.

That is a real limitation. “The cat chased the dog” needs cat → dog to be different from dog → cat. Section 3 fixes it.

Now write it. The browser has np and softmax ready.

CodeWrite the scores line of attention.

Enter keeps the indent · Tab indents · Esc then Tab leaves the editor · ⌘/Ctrl + Enter runs

Try it

Press “Car becomes [5, 5]”. Car’s row of weights jumps to almost [0, 0, 1]. Now press “All three the same”: every row becomes [0.33, 0.33, 0.33]. When no score is much bigger than the others, attention is just a plain average.

3. Q, K, V: three versions of each word

So far each word used the same vector for three jobs. In a real Transformer, each word makes three versions of itself, by multiplying with three learned matrices:

The score of “a looks at b” is now a’s query · b’s key.

🔒 Answer the question above to unlock

Then everything is as before: attention(Q, K, V) = softmax(Q @ K.T / √d_k) @ V.

Three roles for every word: a question (Q), a label (K) and content (V)

Edit any number, or try a preset. Point at or tap a score to read how it was made.

1 the words and the three projection matrices

X
cat
dog
car
W_Q
W_K
W_V

2 every word gets a query, a key and a value

Q = X W_Q
cat20
dog11
car02
K = X W_K
cat00
dog10
car20
V = X W_V
cat20
dog11
car02

3 score each query against each key, softmax, mix the values

scores = Q Kᵀ
catdogcar
cat024
dog012
car?00
not symmetric
weights = softmax(scores / √2)
catdogcar
cat0.050.190.77
dog0.140.280.58
car0.330.330.33
output = weights V
cat0.281.72
dog0.561.44
car11

One-way keys: each key is [y, 0], so cat can look at car, but car does not look at cat.

Pick a score. The question comes from the row word’s Q, the label from the column word’s K.

Press “Swap the keys”: cat now looks mostly at car. Then press “One-way keys”: W_K = [[0, 0], [1, 0]], so each word’s key is its y number moved to the front, [y, 0]. Look at the badge under the scores.

Number Take the “one-way keys” example: W_Q is the identity and W_K = [[0, 0], [1, 0]], so K = X @ W_K turns each word [x, y] into the key [y, 0]. cat = [2, 0] and car = [0, 2]. What is the score of car looking at cat, car’s query · cat’s key?
Shape A sentence has 5 words. Q and K are both (5, 4). What shape is Q @ K.T?
I got stuck here If W_V is not in the scores, does it matter at all?

The scores (and the weights) only use Q and K. W_V does not change who gets looked at. It changes what is copied. Press “V keeps only x”: the weights stay exactly the same, but every output loses its y number, because V only carries x. Q and K choose which words to read; V decides what each of those words gives.

Go deeper Why not just use one matrix?

The score q_a · k_b equals x_a W_Q W_Kᵀ x_bᵀ. So the scores only depend on the product W_Q W_Kᵀ, a 2×2 matrix here. If that product is symmetric, the scores are symmetric too. That is why “Swap the keys” still gives a symmetric table: the swap matrix is symmetric. “One-way keys” is not, so cat can look at car without car looking back.

Splitting the product into W_Q and W_Kᵀ also lets the model compare words in a smaller space (d_k numbers) than the word vectors themselves (d_model numbers). Level 15 uses that to run several attentions side by side.

4. Nobody designs these matrices

The matrices in section 3 were written by hand. In a real model nobody writes them: they are learned, the same way level 2 learned a line. Here is the smallest version. Two words, apple = [1, 0] and sweet = [0, 1]. The task: sweet’s output should be apple’s content, [1, 0]. W_K and W_V are the identity, and W_Q = [[0, 0], [c, 0]]: only the bottom-left number, c, can change. Sweet’s query is then [0, 1] @ W_Q = [c, 0], so c sets how strongly sweet looks at apple. The lab measures the difference from the target with level 2’s squared error, added over both numbers: loss = (out_x − 1)² + (out_y − 0)². Gradient descent changes c to make the loss smaller.

🔒 Answer the question above to unlock

Attention is learned: gradient descent changes one number, c, in W_Q

Press a step button and let training change c, or set c yourself with the slider.

W_Q
00
c = 0.000
weights (who looks at whom)
applesweet
apple0.500.50
sweet??
steps
0
sweet’s query
[0.00, 0]
sweet’s output
[?, ?]
target
[1, 0]
loss
0.50000
d loss / d c
-0.3536
loss for every value of the number c
00.511.5-202468c →loss

The curve appears as far as the ball has gone.

c = 0.00 · loss 0.50000 · slope -0.3536, so the next step moves c by +0.3536
the current c (a step moves it) its slope, d loss / d c (the gradient) loss 0: the target, never quite reached
Number At c = 0, sweet’s query is [0, 0]. Its keys are apple = [1, 0] and sweet = [0, 1]. How much weight does sweet put on apple?
Predict firstGuess before you try it: if you keep pressing “200 steps”, what happens to c?

Are Q, K and V trained in advance? No. Two different kinds of numbers live in a model, and only one kind is trained.

kindexampleswhere it comes fromsaved with the model?
parametersW_Q, W_K, W_V, the embedding tablestart random; training changes themyes, fixed after training
activationsQ, K, V, scores, weights, outputscomputed from the input X every time the model runsno: deleted after each run (training keeps them until the backward pass)

So W_Q is stored inside the model; Q is not. Give the model a new sentence and the same fixed W_Q produces a new Q. (Level 18 shows one exception: while writing a long text, K and V of earlier words are kept for a while, the KV cache.)

Running a trained model on new input is called inference: the parameters stay fixed, the input passes through every layer once, and the model gives its output. Attention is one step inside every layer of that run. Training runs the same steps, then compares the output with the target, then changes the parameters.

ChooseThe model is trained and saved. A user types a new sentence. Which numbers are different from the last time the model ran?
I got stuck here Is the target an attention pattern? Do we tell the model where to look?

No. The loss compares sweet’s output with [1, 0]. The weights never appear in the loss. Training changes c until the output is right. The weights change too, but only because c changed.

The target comes from the task. In this lab it was given. In a language model it is already in the text: the next word (level 11). Nobody ever writes down “this word should look at that word”.

Go deeper Train W_V too: same loss, different attention

The lab trains only c and keeps W_V fixed. What if W_Q, W_K and W_V are all trained? demo.py does it four times:

what is trainedsweet’s weight on appleloss
W_Q and W_K (W_V fixed to the identity)1.000.0000
all three, starting from the identity0.520.0000
all three, random start (seed 0)0.160.0000
all three, random start (seed 1)0.730.0000

Every run produces the right output, but the attention weights are completely different. When W_V can change, V can make “look at sweet” give [1, 0] too. So an attention map is one of many ways to get the same output. A heatmap shows what one trained model does, not what it must do. Read heatmaps with care.

5. Why divide by √d_k

So far every score was divided by √d_k = √2 before softmax. (d_k is the number of numbers in each query and key, here 2.) What happens without it?

🔒 Answer the question above to unlock
Number Without dividing by √2, cat’s scores are cat·cat, cat·dog, cat·car = [4, 2, 0]. Use e⁴ ≈ 54.6, e² ≈ 7.4, e⁰ = 1. What weight does cat put on itself after softmax? (2 decimals)

Check it: uncheck “divide by √d_k” in the first lab and watch cat’s row.

Without the division the weights get sharper. With two numbers per word that is a small effect. With 64 or 512 numbers it is a big one: a dot product adds up d_k terms, so it grows with d_k. Softmax of big numbers puts nearly all the weight on one word, and then the gradients that train W_Q and W_K almost vanish.

Go deeper How big do the scores get?

Take random query and key vectors whose numbers have a standard deviation (std) of 1. Their dot product has std √d_k:

d_kstd of q · kafter dividing by √d_k
21.421.01
648.081.01
51222.551.00

Dividing by √d_k puts the scores back to std 1, whatever d_k is. demo.py prints this table.

6. Write one attention head

You now know every piece of one attention head: three projections, scores, scaling, softmax, a weighted average. Put them together in one function. The browser has np and softmax ready.

🔒 Answer the question above to unlock
CodeWrite one attention head from start to end: the three projections, the scaled scores, softmax, the weighted average. Several lines.

Enter keeps the indent · Tab indents · Esc then Tab leaves the editor · ⌘/Ctrl + Enter runs

Every word here may look at every word, including the words after it. Level 15 adds masks, so a word can be stopped from looking ahead or at padding, and splits attention into several heads.

Recap

a summary for when you finish the level

The key formulas and common mistakes appear here once you clear the level.

You can now

  • Compute one word’s scores, softmax weights and output by hand for 2–3 words.
  • Make Q, K and V from X with three matrices and say the shape of every step.
  • Write one attention head in NumPy from X to the output.

Keep in mind

  • Attention: softmax(Q @ K.T / √d_k) @ V
  • Q = X @ W_Q, K = X @ W_K, V = X @ W_V
  • (L, d_k) @ (d_k, L) → scores (L, L): row i = the word that looks, column j = the word it looks at
  • Each row of weights adds up to 1; the output is a weighted average of the rows of V
  • Divide by so the scores keep std about 1 for any

Common mistakes

  • Writing Q @ K or Q * K.T for the scores: every query dotted with every key is Q @ K.T.
  • Reading rows and columns of the weight table swapped: weights[0][2] is how much word 0 takes from word 2.

Press ? for keyboard shortcuts

Reading mode · every part open, no stars