06-01 Multi-Head Attention
Why this matters
Section 3.2.2 projects Q, K, and V into multiple subspaces so different heads can attend to different patterns.
Intuition first (no jargon)
Multiple smaller attentions run in parallel, then concatenate and project back to model width.
Paper grounding
- The paper defines
head_i = Attention(QW_i^Q, KW_i^K, VW_i^V). - Multi-head output is
Concat(head_1, ..., head_h)W^O.
In the paper, queries, keys, and values are projected h times with learned linear projections before per-head attention.
Code walkthrough
jsexport function splitHeads(x, numHeads) {} export function combineHeads(heads) {} export function multiHeadAttention(x, params, mask) {}
Your task
Implement head split/combine and full multi-head attention.
- Reshape
[T][dModel]into[H][T][dHead]. - Run scaled masked attention per head.
- Concatenate heads and project back to
dModel.
Hints
- Validate
dModel % numHeads === 0. - Keep head dimension naming consistent.
- Test split then combine round-trip.
Check your thinking
- Why can heads specialize?
- What does output projection do after concatenation?
- Why is shape bookkeeping critical here?
Stretch (optional)
Return per-head attention maps for visualization.
Likely test focus
- Shape correctness at each stage.
- Split/combine round-trip integrity.
- Deterministic output with fixed params.
What should improve
Your implementation now supports split-attend-concatenate multi-head computation.
Bridge to next lesson
Next lesson: add token-wise feed-forward transformation.