Multi-head attention
Overview
Section titled “Overview”| Module | L5.3 · build · Python · Pass 5 · 2 to 3 h, plus your graded tests (rung R5) |
| You build | python/tinyllm/xfmr/mha.py: MultiHeadAttention with split_heads, merge_heads, head_mask, attend, forward, load_packed_in_proj; and your own oracle tests in python/tests/l5-3-mha/ |
| Contract | course/contracts/py/tinyllm/xfmr/mha.pyi |
| Tests | course/tests/L5.3/test_mha.py (what they check: section 4); the oracle is torch’s nn.MultiheadAttention with copied weights; your tests are graded by mutation, threshold 0.80 with every pitfall fault required |
| Needs | L5.1 scaled dot-product attention · L0.4 Linear, Module · L0.2 reshape, transpose · L0.1 Tensor · M06.3 PCG32 (or --ref-deps) |
| Used by | L5.5 encoder and decoder layers · L6.1 GPT blocks · L6.2 BERT layers |
| Milestone | MS-L5 ({tinyllm} train transformer, then translate --beam 4) |
| Optional depth | Vaswani et al., “Attention Is All You Need” (2017), section 3.2.2; Michel, Levy, and Neubig, “Are Sixteen Heads Really Better than One?” (NeurIPS 2019) |
Key Takeaways
Section titled “Key Takeaways”- A head attends inside its own slice of the features, so one position can read two places at once; one wide head cannot (
test_hand_example_two_heads,test_hand_example_one_head_mixes_everything). - Splitting is a reshape to
[B, T, H, d_head]and a transpose to[B, H, T, d_head]; merging undoes both, in that order (test_split_and_merge_layout). - Each head scales by , the width it compares, not (
test_golden_torch). - Self-attention and cross-attention are the same module: queries from
x_q, keys and values fromx_kv(test_cross_attention_reads_keys_from_x_kv). - A per-sequence mask needs a head axis, or sequence ‘s mask silently lands on head (
test_mask_shapes_agree).
How to work this chapter
Section titled “How to work this chapter”ol start L5.3 # stubs mha.py; prints your test path and rung (R5)ol tests L5.3 # the course tests# write your oracle tests in python/tests/l5-3-mha/ (section 4 says what to cover), then:ol check L5.3 # course tests and the mutation grade of your testsol mutate L5.3 # the full grade, cached by your test files' hashol diff L5.3 # after passing: your code against the reference1. Why now
Section titled “1. Why now”L5.1 gave you one attention: every query compares itself with every key through one dot product over all features, and gets one weight vector. Take the addition task of MS-L5: to write the tens digit of the sum, the decoder has to look at the tens digit of and the tens digit of , four positions apart. A single softmax can split its mass between the two, but then it averages them into one blurred vector. Vaswani et al. split the model width into heads that each run their own attention with their own projections, and concatenate the results: head 0 can find ‘s digit while head 1 finds ‘s. Every model of Parts 5 and 6 (L5.5’s encoder-decoder, L6.1’s GPT, L6.2’s BERT) calls this one module, so its layout and names are fixed here once.
2. Principles
Section titled “2. Principles”| Symbol | Meaning | Type / shape |
|---|---|---|
| batch, query length, key length | int | |
model width d_model | int | |
number of heads n_heads | int, divides | |
width of one head d_head | int | |
| the inputs queries and keys/values come from | [B, Tq, d], [B, Tk, d] | |
the four projections (q_proj, k_proj, v_proj, out_proj) | float32[d, d] each, plus biases [d] | |
| head ‘s slice of the projected queries, keys, values | [B, Tq, dh], [B, Tk, dh] | |
| mask, True = may attend | bool, broadcasts to [B, H, Tq, Tk] |
2.1 One head
Section titled “2.1 One head”One head is L5.1 after learned projections:
With the module is exactly this followed by ; test_one_head_is_projected_attention computes it in numpy from the module’s own weights.
2.2 Heads
Section titled “2.2 Heads”With heads the projections stay , but their output is read in slices of width : head owns features to . Each head attends in its own slice, scaled by because a dot product of terms with unit-variance entries has variance :
In code all heads run in one batched call. A projected [B, T, d] tensor is reshaped to [B, T, H, d_h] (the slices become an axis) and transposed to [B, H, T, d_h], so L5.1 sees H more batch rows. Merging reverses it: transpose back to [B, T, H, d_h], then reshape to [B, T, d]. Reshaping [B, H, T, d_h] straight to [B, T, d] is legal numpy and wrong: it glues together the rows of different positions. The cost is the same as one wide head: four projections and attentions of width .
2.3 Masks for every head
Section titled “2.3 Masks for every head”Masks come from L5.2 and mean “True = may attend”. A causal mask is [Tq, Tk] and applies to every sequence and head; a padding mask is per sequence. The module accepts [Tq, Tk], [B, Tq, Tk] (it inserts the head axis at 1), and any shape that already broadcasts to [B, H, Tq, Tk], such as [B, 1, 1, Tk] key padding. Numpy aligns shapes from the right, so a [B, Tq, Tk] mask without the inserted axis broadcasts against [B, H, Tq, Tk] with its batch axis on the head axis: when nothing fails and the masks are simply wrong.
2.4 Self- and cross-attention
Section titled “2.4 Self- and cross-attention”attend(x_q, x_kv, mask) takes queries from x_q and both keys and values from x_kv. Self-attention passes the same tensor twice; the decoder’s cross-attention (L5.5) passes its own states as x_q and the encoder’s output as x_kv. The output has x_q’s length. Keys are a set: permuting the positions of x_kv (with its mask) leaves the output unchanged, and permuting x_q permutes the output rows. Position information therefore has to come from outside, which is L5.4’s job.
2.5 Dropout and torch’s packed layout
Section titled “2.5 Dropout and torch’s packed layout”Attention dropout zeroes attention weights at random during training (rate dropout, drawn from PCG32(0).substream("dropout") by L5.1) and is off in eval mode. torch’s nn.MultiheadAttention keeps stacked in one in_proj_weight of shape [3d, d], rows in the order q, k, v; load_packed_in_proj copies them into the three Linears.
3. Worked example by hand
Section titled “3. Worked example by hand”, , , and every projection the identity with zero bias, so , . One query and two keys:
| features 0, 1 (head 0) | features 2, 3 (head 1) | |
|---|---|---|
Head 0. Scores ; weights and ; context .
Head 1. Scores ; weights and ; context .
Merge. Head 0’s context fills features 0 and 1, head 1’s fills 2 and 3: , and leaves it. This is test_hand_example_two_heads.
With one head of width 4 the same numbers give scores , weights for every feature, and output : head 0’s preference for is outvoted (test_hand_example_one_head_mixes_everything).
4. The interface
Section titled “4. The interface”class MultiHeadAttention(Module): def __init__(self, d_model: int, n_heads: int, dropout: float = 0.0, bias: bool = True, rng=None) def split_heads(self, x: Tensor) -> Tensor # [B, T, d] -> [B, H, T, dh] def merge_heads(self, x: Tensor) -> Tensor # [B, H, T, dh] -> [B, T, d] def head_mask(self, mask, B, Tq, Tk) -> Optional[NDArray] def attend(self, x_q: Tensor, x_kv: Tensor, mask=None) -> tuple[Tensor, Tensor] # out, weights [B, H, Tq, Tk] def forward(self, x_q: Tensor, x_kv: Tensor, mask=None) -> Tensor def load_packed_in_proj(self, weight, bias=None) -> None # torch's [3d, d], rows q | k | vParameters, in state_dict order: q_proj.weight, q_proj.bias, k_proj.*, v_proj.*, out_proj.* (L0.4 Linears, drawn from one rng).
What the tests check
Section titled “What the tests check”| Test | KIND | Checks | Why it matters downstream |
|---|---|---|---|
test_hand_example_two_heads | unit | section 3: weights per head and the merged output | you and the test agree on the layout and the scale |
test_hand_example_one_head_mixes_everything | unit | the same numbers with one head | why heads exist |
test_golden_torch | golden | torch nn.MultiheadAttention: output, per-head weights, every gradient, self (causal + padding) and cross (padding) | L5.5, L6.1, L6.2 load torch-trained weights |
test_gradcheck_inputs_and_parameters | gradcheck | both inputs and all 8 parameter tensors against frozen central differences | every model trains through this backward |
test_split_and_merge_layout | property | head owns features ; merge inverts split | positions never mix |
test_heads_are_independent | property | changing head 1’s query rows leaves head 0’s weights bitwise equal | heads are separate subspaces |
test_one_head_is_projected_attention | differential | equals numpy softmax then | the definition, independently |
test_mask_shapes_agree | property | [Tq, Tk], [B, Tq, Tk], [B, H, Tq, Tk] masks agree, with | the decoder’s causal-and-padding masks |
test_masked_keys_are_never_read | property | masked keys weigh exactly 0; their values cannot change the output | padding in every batch |
test_cross_attention_reads_keys_from_x_kv | property | output length from x_q; x_kv positions are a set | L5.5’s cross-attention |
test_dropout_only_in_training | unit | eval mode equals no dropout | evaluation and decoding are deterministic |
test_parameter_names_shapes_and_init | unit | the eight names in order, seeded init, bias=False | checkpoint keys of Parts 5 to 7 |
test_load_packed_in_proj | unit | torch’s packed rows are q, k, v | loading torch weights (L5.5’s golden test) |
test_validation | boundary | heads must tile ; shapes; bool masks | wiring bugs fail loudly |
Your graded tests (rung R5)
Section titled “Your graded tests (rung R5)”Rung R5 asks for oracles. Write multi-head attention again in float64 numpy from the module’s own state_dict (project, reshape, transpose, softmax with the mask, merge, project) and compare attend’s output and weights with it, for self- and cross-attention, with a per-sequence [B, Tq, Tk] mask where . Add the section 3 numbers, a split/merge check, eval-mode dropout, the packed load order, and a gradient check with your own tinyllm.num.gradcheck (M04.1) in float64. Import only names in contracts/py. ol check L5.3 requires a mutation score of at least 0.80 with every pitfall fault (s01 to s07) killed.
5. Pitfalls
Section titled “5. Pitfalls”| Pitfall | Symptom | Caught by |
|---|---|---|
| 1. merging (or splitting) heads with a plain reshape | positions mix; torch disagrees; nothing crashes | test_split_and_merge_layout, test_golden_torch (mutants s01, s02) |
| 2. scaling by | softmax too flat by a factor | test_hand_example_two_heads, test_golden_torch (mutant s03) |
3. projecting keys from x_q | cross-attention reads the decoder, or crashes when lengths differ | test_cross_attention_reads_keys_from_x_kv (mutant s04) |
4. a [B, Tq, Tk] mask without a head axis | with head gets sequence ‘s mask | test_mask_shapes_agree (mutant s05) |
| 5. dropout in eval mode | decoding is random, evaluation noisy | test_dropout_only_in_training (mutant s06) |
6. reading torch’s packed in_proj as k, q, v | a loaded checkpoint attends with swapped roles | test_load_packed_in_proj (mutant s07) |
| the output projection skipped | heads are never mixed; shapes still fit | test_one_head_is_projected_attention (mutant s08) |
6. Where it’s used next
Section titled “6. Where it’s used next”| Direction | Module | How it uses this |
|---|---|---|
| Back | L5.1 | scaled_dot_product_attention runs all heads in one batched call |
| Back | L0.4 | Linear for the four projections, Module for registration |
| Back | L0.2 | reshape and transpose for split and merge |
| Back | L0.1 | the Tensor every input and parameter is |
| Back | M06.3 | PCG32 for the default init and dropout streams |
| Forward | L5.5 | encoder self-attention, decoder masked self-attention and cross-attention |
| Forward | L6.1 | GPT-2’s causal self-attention (c_attn split into q, k, v) |
| Forward | L6.2 | BERT’s bidirectional self-attention |
| Forward | L7.5 | grouped-query attention shares heads between query heads |
If you skip this module, ol check L5.5 stops with needs L5.3: build it, or pass --ref-deps.
Going further
Section titled “Going further”| Your piece | Production equivalent | What it adds | Where to look |
|---|---|---|---|
split_heads / merge_heads | PyTorch F.scaled_dot_product_attention on [B, H, T, d] | fused kernels (FlashAttention, memory-efficient) behind one call | torch/nn/functional.py, aten/src/ATen/native/transformers/ |
load_packed_in_proj | HF GPT2Attention c_attn, Llama’s separate q/k/v | fused QKV matmul for speed, split views | transformers/models/gpt2/modeling_gpt2.py |
| one K, V per head | multi-query and grouped-query attention | fewer K, V heads: a smaller KV cache | Shazeer (2019); Ainslie et al. (2023); L7.5 |
| attention dropout | most modern LLMs train with 0 | dropout matters for small data, not trillions of tokens | Llama and Mistral configs |