Multi-Head Attention
Goal
Split the embedding channels into heads, run causal attention per head, recombine them, and prove that the residual-stream shape is preserved.
One attention head gives one learned routing pattern. Multi-head attention partitions the feature channels so several attention lookups can happen in parallel.
For n_embd = 12 and n_head = 3:
head_dim = 12 / 3 = 4
(B,T,12)
→ split/rearrange → (B,3,T,4)
→ attention per head
→ concatenate → (B,T,12)
The sequence is not split into different token sets. Every head processes the same positions under the causal rule; the heads use different feature subspaces/projections.
Same token positions, three parallel routing views
All three heads below receive the same four token positions. The synthetic strong-link summaries differ because each head has a different learned projection; the head outputs are then concatenated and projected back to residual width.
(B,T,12)(B,3,T,4)t12t23t34t4t4→t4· 50%t4→t2· 30%
t12t23t34t4t4→t3· 55%t4→t1· 25%
t12t23t34t4t4→t2· 45%t4→t3· 35%
(B,T,12)The percentages above are only the strongest synthetic links shown for one query position; they are not a complete attention distribution. The important relationship is structural: every row contains the same T positions, while each head can route information differently.
Why several heads can be useful
Imagine one query position needs two kinds of evidence: a nearby syntax cue and a more distant content cue. A single head produces one normalized routing pattern for its feature subspace. Multiple heads give the model several separately learned routing patterns before their results are recombined.
Do not turn that into the stronger claim that “head 1 always learns syntax and head 2 always learns meaning.” Heads are learned components, and their roles are not guaranteed to stay human-readable.
The shape trace matters because it prevents a common misconception:
(B,T,C) = (2,5,12)
→ (B,H,T,D) = (2,3,5,4)
The five token positions are still present in every head. What changed is how the 12 feature channels are represented across three head dimensions. After attention, those head feature outputs are rearranged back to 12 channels so the residual stream can continue as (2,5,12).
When a multi-head implementation fails, label axes B,H,T,D explicitly. Many silent bugs come from a tensor having the right numbers but the wrong semantic axis order.
Heads split feature capacity, then recombine it
Suppose C=12 and H=3. A common design uses D=4 features per head.
Each head computes its own attention output of shape (B,T,4). Concatenating three heads restores (B,T,12), and a learned output projection mixes information across the concatenated head features.
Heads are therefore not three independent mini-models whose predictions are averaged. They are parallel feature-routing components inside one layer.
Different heads can learn different useful patterns, but forcing a fixed human story onto each head is risky. Reliable checks are more concrete: correct reshape order, learned projections, legal masking in every head, and restoration of the residual-stream width before addition.
Predict
Trace every reshape boundary
Run the notebook lab and record shapes after projection, head split, attention, transpose/concatenation, and output projection. Change n_head to another valid divisor and verify input/output remain (B,T,C).
Loading lab…
Then try an invalid combination such as n_embd=10, n_head=3. A clear early divisibility failure is better than allowing a later mysterious reshape error.
Quick Check
Explain it back
Explain why “more heads” is not the same as “more sequence length,” and name the divisibility invariant you would test before running attention.
Key Takeaways
- Heads partition feature capacity, not token positions.
- Each head keeps the causal rule.
- Head dimensions must divide the embedding width in the simple implementation.
- Recombination returns to the residual-stream shape
(B,T,C).
Next Lesson
Next, preserve that residual stream explicitly by adding each sublayer's update back to its input.
References
- Vaswani et al., Attention Is All You Need.
- PyTorch, MultiheadAttention.
Completion is stored locally on this device.