Level 15 · Theory · runs in your browser

Masks and multi-head attention

How do you stop a word from looking at some words, and look in several ways at once?

In level 14 every word looked at every word, with one set of W_Q, W_K and W_V. A GPT needs two more things:

  1. Masks, so a word can be stopped from looking at some words: the words after it, or padding.
  2. Several heads, so the model can look in several ways at once.

This level adds both and ends with one attention head of a GPT, written by you.

1. Masks: who is not allowed to look

Two situations need some cells of the weight table forced to zero.

A mask is a table of True/False with True = blocked. A blocked score becomes −∞ before softmax, and e^−∞ = 0. The two masks are combined with OR: a cell is blocked if either mask blocks it.

ChooseCat’s scores are [0.34, 0.05, blocked]; the third word is a PAD token. What if a blocked score were set to 0 instead of −∞? Use e^0.34 ≈ 1.40, e^0.05 ≈ 1.05, e^0 = 1. The PAD token would get…
🔒 Answer the question above to unlock

Masks decide which words each word may attend to

Tap a word to turn it into padding. Point at or tap a weight to see why it is blocked.

padding mask (1, 5): one row
00011
↓ copied down every row (broadcast)
00011
00011
00011
00011
00011
OR
causal mask (5, 5)
01111
00111
00011
00001
00000
=
combined (5, 5)
thecatsatPADPAD
the01111
cat00111
sat00011
PAD00011
PAD00011
weights (equal scores, blocked cells set to −∞ before softmax)
thecatsatPADPAD
the1.000.000.000.000.00
cat0.500.500.000.000.00
sat0.330.330.330.000.00
PAD0.330.330.330.000.00
PAD0.330.330.330.000.00

In a batch the padding mask has shape (B, 1, 5); the 1 is the query axis it gets copied along.

12 of 25 cells are open.
1 = blocked: this score becomes −∞, its weight 0 0 = allowed
Try it

In the mask lab, uncheck “causal” and tap the last two words to make them padding. Which columns are now empty? Turn “causal” back on. Which row always has exactly one open cell, whatever you pad?

Shape A sentence has 6 tokens. What shape is its causal mask?
I got stuck here Why is the causal mask a 6×6 table and not just 6 numbers?

Because it is a mask on the weights, and the weight table has one row per word that is looking and one column per word being looked at. For a causal mask, “may word i look at word j” depends on both i and j, so it needs a full 6×6 table.

The padding mask really is just 6 numbers, one per column: a PAD token is blocked for everyone. That is why it is stored as one row, with shape (1, 6), and copied down every row when it is used.

The size comes from the sequence that looks at itself: 6 tokens looking at the same 6 tokens. Side trip N7 has models with two sequences (a source and a target of different lengths). Its section 5 shows which mask each attention uses, and which length sets its size.

Number Causal mask only, 4 tokens. How many of the 16 weights can be non-zero?
🔒 Answer the question above to unlock

One mask for a whole batch

A batch holds several sentences, and each sentence has its own padding. Its scores have shape (B, L, L).

I got stuck here What is the B in (B, 1, L)?

B is the batch size: how many sentences are processed together. Each sentence has its own padding pattern, so the padding mask is (B, L), one row per sentence. The scores are (B, L, L). To make the shapes fit, the mask gets a 1 in the middle, (B, 1, L), and NumPy copies that single row across all L query rows. This copying is called broadcasting.

Shape A batch has 2 sentences of 5 tokens. The padding mask is (2, 1, 5) and the scores are (2, 5, 5). What shape is the mask after broadcasting?

In the lab, the PAD rows still spread their own weights over the real words. That is fine: a PAD token’s output is never used, and the loss skips it. The mask only has to stop real words from looking at PAD.

Masks use two NumPy functions. np.where(mask, a, b) takes a where mask is True and b everywhere else:

blocked = np.array([True, False, False])
np.where(blocked, -np.inf, np.array([0.34, 0.05, 0.9]))   # [-inf, 0.05, 0.9]

To build a causal mask, one more function does the work. np.triu(M, k) (“triangle, upper”) keeps the entries of M on and above diagonal k and sets the rest to 0 (or False). k = 0 is the main diagonal, k = 1 starts one step above it:

np.triu(np.ones((3, 3), dtype=int), k=0)    # [[1, 1, 1],
                                            #  [0, 1, 1],
                                            #  [0, 0, 1]]
CodeBuild the causal mask: True (or 1) wherever a word would look ahead (column > row), False (or 0) elsewhere.

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

2. Many heads

One set of W_Q and W_K gives one way of looking: one score table. But words relate in several ways at once. Multi-head attention runs a few attentions side by side, each on a separate slice of the numbers, and then mixes them.

Here d_model = 4 and there are 2 heads, so each head gets d_k = 2 columns. Heads are numbered from 0, like everything else. Each head outputs 2 numbers per word, and the heads’ outputs are placed next to each other in one table called concat, head 0 first.

Splitting does not add parameters. W_Q, W_K and W_V stay (d_model, d_model): the heads share their columns instead of each getting a full copy, so 8 heads cost the same as 1 big head.

🔒 Answer the question above to unlock
Shape 3 words, 2 heads, each head outputs 2 numbers per word. The heads’ outputs are joined into one table called concat, one row per word. What shape is concat?
Number 2 heads each output 2 numbers per word. Their outputs are placed side by side in one table, concat, head 0 first. Heads and columns both count from 0. At which column of concat does head 1’s output start?
🔒 Answer the question above to unlock

Now pick a head in the lab and follow its columns.

Two heads, each on separate columns, mixed back together by W_O

Pick a head to see which numbers it owns. X and W_O can be edited.

1 the words; W_Q and W_V are the identity, so Q = V = X

X (3 words × d_model 4) = Q = V
0123
cat
dog
car
K = X W_K
0123
cat2001
dog1110
car0211

2 each head attends using only its own two columns

head 0 weights (columns 0, 1)
catdogcar
cat0.770.190.05
dog0.330.330.33
car0.050.190.77
head 1 weights (columns 2, 3)
catdogcar
cat0.200.400.40
dog0.400.200.40
car0.250.250.50

3 the head outputs sit side by side; W_O mixes them into one output

concat(head 0, head 1) (3 × 4)
h0h0h1h1
cat1.720.280.600.80
dog110.800.60
car0.281.720.750.75
W_O (4 × 4)
h0
h0
h1
h1
output = concat W_O
cat1.720.280.600.80
dog110.800.60
car0.281.720.750.75
Head 0 reads columns 0–1 of Q, K and V, writes columns 0–1 of concat, and only rows 0–1 of W_O read it.
the numbers head 0 owns W_O is where the heads get mixed: 2 heads in, 1 output out.

Head 0 uses columns 0 and 1 of Q, K and V, which here are exactly the cat, dog, car of level 14 (section 2), so it gives the same weights. Head 1 uses columns 2 and 3 and, with its swapped keys, a different pattern. In the end there is one output: 2 heads in, 1 output out, after W_O.

I got stuck here Does head 0 only see half of each word?

No. Each head uses all d_model numbers of every word: Q = X @ W_Q mixes every input number into every column of Q. The splitting happens to Q, K and V, not to X. In this lab W_Q is the identity, so the columns of Q happen to be the columns of X. In a trained model, column 0 of Q is a mix of all four numbers of the word.

Go deeper What W_O does

After the heads run, their outputs sit next to each other: concat is (3, 4). Without W_O, head 0’s result would stay in columns 0–1 and head 1’s in columns 2–3, always separate. W_O is a (4, 4) matrix. Output column j is a mix of all four concat columns, so every output number can use both heads. Rows 0–1 of W_O are the weights for head 0’s numbers, rows 2–3 for head 1’s.

Number A model has d_model = 512 and 8 heads. How many numbers does each head get (dₖ)?

3. All of it at once

You now know every piece: three projections, scores, scaling, a mask, softmax, a weighted average. Here they are as one function, with a causal mask. This is one attention head of a GPT.

🔒 Answer the question above to unlock
CodeWrite one attention head with a causal mask, from X and the three weight tables to the output out. Several lines.

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

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

  • Build a causal mask and a padding mask, and say which cells of the weight table they block.
  • Split d_model into heads and say where each head’s output sits in concat.
  • Write one masked attention head in NumPy, from X to the output.

Keep in mind

  • True = blocked; a blocked score becomes −∞ before softmax, and e^−∞ = 0
  • Causal mask, shape (L, L): np.triu(…, k=1) on an all-True table
  • Apply it with np.where(mask, -np.inf, scores)
  • : the heads split the width evenly
  • concat is (L, d_model), head 0 first; W_O mixes the heads

Common mistakes

  • Setting a blocked score to 0 instead of −∞: e⁰ = 1, so the blocked word still gets weight.
  • Using k=0 in np.triu: that blocks the diagonal too, and a word may look at itself.

Press ? for keyboard shortcuts

Reading mode · every part open, no stars