# 15. Masks and multi-head attention

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

LLM by Hand · Theory · runs in your browser · interactive page: https://llm.liko.page/learn/masks-and-heads/

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.

- **No looking ahead.** When a model writes text one word at a time, word 3 must not use word 4: word 4 is what it is
  trying to predict. The **causal mask** blocks every cell above the diagonal.
- **No looking at padding.** Sentences in a batch have different lengths, so short ones are filled with PAD tokens
  until they all have the same length. Nobody should look at those. The **padding mask** blocks their columns.

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.

**Predict.** Cat’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…

A. 0% of the weight
B. about 29% of the weight
C. exactly 50% of the weight

*Answer it on the page to check your work.*

*[Interactive lab: Mask — open the page to use it]*

**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?

**Question.** A sentence has 6 tokens. What shape is its causal mask?

*Answer it on the page to check your work.*

**If you are stuck: 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](/learn/encoder-decoder/) 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.

**Question.** Causal mask only, 4 tokens. How many of the 16 weights can be non-zero?

*Answer it on the page to check your work.*

### 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).

**If you are stuck: 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.

**Question.** 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?

*Answer it on the page to check your work.*

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:

```python
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:

```python
np.triu(np.ones((3, 3), dtype=int), k=0)    # [[1, 1, 1],
                                            #  [0, 1, 1],
                                            #  [0, 0, 1]]
```

**Code question.** Build the causal mask: True (or 1) wherever a word would look ahead (column > row), False (or 0) elsewhere.

Fill in the blank (`____`):

```python
def causal_mask(L):
    return ____

print(causal_mask(4).astype(int))
```

*Answer it on the page to check your work.*

## 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.

**Question.** 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?

*Answer it on the page to check your work.*

**Question.** 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 it on the page to check your work.*

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

*[Interactive lab: Multi head — open the page to use it]*

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`.

**If you are stuck: 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.

**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.

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

*Answer it on the page to check your work.*

## 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.

**Code question.** Write one attention head with a causal mask, from X and the three weight tables to the output out. Several lines.

Fill in the blank (`____`):

```python
def masked_attention(X, W_Q, W_K, W_V):
    L, d_k = X.shape[0], W_Q.shape[1]
    # 1. Q, K, V   2. scaled scores
    # 3. causal mask: True = blocked, blocked scores become -inf
    # 4. softmax, then @ V
    ____
    return out

X = np.array([[2.0, 0.0], [1.0, 1.0], [0.0, 2.0]])
W_K = np.array([[0.0, 1.0], [1.0, 0.0]])
print(masked_attention(X, np.eye(2), W_K, np.eye(2)))
```

*Answer it on the page to check your work.*

## 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.
