Lesson 81: Multi-Head Attention

USAAIO Lesson 81, from Week 28 of Phase 3, building multi-head attention from scratch. It covers projecting Q, K, and V into h parallel subspaces, running scaled dot-product attention in each head, and concatenating before the W_O projection. It proves that the parameter count is 4·d_model² regardless of h, and shows that h=1 reduces to single-head attention. You implement MultiHeadAttention as an nn.Module and verify every shape. The lesson runs to 28 slides.

Subject: Machine Learning · 55 slides · code lesson

Open the interactive version of this deck · Homework for this lesson

What this lesson covers

The lesson, slide by slide

1. Multi-Head Attention

Title

USAAIO · Lesson 81 · Week 28 (Phase 3: Transformers)

Split Q/K/V into h parallel subspaces, run scaled dot-product attention in each, then concat + project. Same parameter budget as single-head — vastly richer representations.

2. By the end of this lesson you can

Objectives

  1. Explain why splitting into h heads captures diverse relationship types simultaneously
  2. Trace the project → split → attend → concat → project dataflow for any (d_model, h)
  3. Prove that MHA parameter count = 4·d_model² regardless of h
  4. Implement MultiHeadAttention from scratch as nn.Module with correct shapes
  5. Verify that h=1 MHA with d_k=d_model is identical to single-head attention

3. What survived from Scaled Dot-Product Attention from Scratch?

Warm-up

Discussion prompt

Before we open Lesson 81: Multi-Head Attention: without looking back, what was the main idea of Scaled Dot-Product Attention from Scratch, and what could you do by the end of it that you could not do before?

Hint: One sentence for the idea, one for the skill. If the second one is blank, that is the part to revisit.

Answer:

implement scaled_dot_product_attention(Q, K, V, mask) from scratch in PyTorch — QKV shape analysis, the 1/sqrt(d_k) scaling rationale, causal (lower-triangular) masking, padding masks, numerical stability (FP16 overflow and the max-subtraction fix), O(T^2 * d_k) complexity, and verification against F.scaled_dot_product_attention.

4. Why multiple heads?

Section

Part 1 of 4

5. Single-head attention's blind spot

Concept

Recall (Lesson 80): scaled dot-product attention computes softmax(QK^T / sqrt(d_k)) V. With one head, every query position attends to all key positions — but in a single representation subspace.

One subspace can only optimize for one type of relationship at a time. In language: the word "bank" may need one head to attend to "river" (semantic) and another to the verb "sat" (syntactic) simultaneously.

Multi-head attention runs h independent attention functions in parallel, each in a lower-dimensional subspace, then merges the results.

6. Break it if you can: Single-head attention's blind spot

Counterexample

Discussion prompt

Recall (Lesson 80): scaled dot-product attention computes softmax(QK^T / sqrt(d_k)) V. With one head, every query position attends to all key positions — but in a single representation subspace.

That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.

Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.

Answer:

Multi-head attention runs h independent attention functions in parallel, each in a lower-dimensional subspace, then merges the results.

7. The multi-head idea

Concept

\[ \text{MultiHead}(Q,K,V) = \text{Concat}(\text{head}_1, \ldots, \text{head}_h)\,W^O \]

\[ \text{head}_i = \text{Attention}\!\left(QW^Q_i,\; KW^K_i,\; VW^V_i\right) \]

symbolshaperole
W^Q_i, W^K_i, W^V_i(d_model, d_k)per-head projection matrices
d_k = d_model / hscalarhead dimension — keeps total compute fixed
W^O(d_model, d_model)output projection merges all heads

8. Fill in: shape for The multi-head idea

Comparison

Comparison matrix

From The multi-head idea: refill the shape column from what you know. The rest of the table is as it appeared.

symbolshaperole
W^Q_i, W^K_i, W^V_i(d_model, d_k)per-head projection matrices
d_k = d_model / hscalarhead dimension — keeps total compute fixed
W^O(d_model, d_model)output projection merges all heads

9. What each head learns

Intuition

Think of each head as a specialist analyst reviewing the same document but with a different lens. Head 1 may focus on syntactic dependencies, Head 2 on co-reference, Head 3 on proximity relationships — all simultaneously, using the same input X.

The W^O projection at the end is the integrator: it takes all the specialist reports (each of dimension d_k) and combines them into a single coherent representation of dimension d_model.

In practice, what each head specializes on is learned from data — you don't assign roles manually. But visualizing attention heatmaps (your homework) reveals interpretable patterns.

10. By analogy: What each head learns

Analogy

Discussion prompt

Explain What each head learns by analogy to something with no Machine Learning in it at all — a queue, a recipe, a map, a bank balance, whatever fits. Then say where your analogy breaks.

Hint: An analogy that never breaks is not an analogy, it is the same idea wearing a hat. Find the seam — that is the part that is actually new.

Answer:

The W^O projection at the end is the integrator: it takes all the specialist reports (each of dimension d_k) and combines them into a single coherent representation of dimension d_model.

11. Dataflow — project, split, attend, concat

Section

Part 2 of 4

12. Step-by-step dataflow

Concept

  1. Project: Q = X·W_Q, K = X·W_K, V = X·W_V — shapes (B, T, d_model)
  2. Split: reshape to (B, T, h, d_k) then transpose to (B, h, T, d_k)
  3. Attend: for each head i: scores_i = Q_i·K_i^T / sqrt(d_k), softmax, then ·V_i
  4. Concat: transpose back to (B, T, h, d_k), reshape to (B, T, d_model)
  5. Project: output = concat · W_O — back to (B, T, d_model)

Steps 2–4 are implemented as a single batched matrix multiply — PyTorch handles all h heads simultaneously. No Python loop over heads needed.

13. Teach it back: Step-by-step dataflow

Explain it

Discussion prompt

Explain Step-by-step dataflow to a student a year behind you. No notation, no jargon they have not met — and it still has to be true.

Hint: If your explanation needs a symbol they have never seen, you are describing the notation rather than the idea.

Answer:

Steps 2–4 are implemented as a single batched matrix multiply — PyTorch handles all h heads simultaneously. No Python loop over heads needed.

14. Guess the shape of the answer: Trace: d_model=8, h=2, d_k=4, seq=4, batch=1

Estimation

Predict first

Use nn.MultiheadAttention(embed_dim=8, num_heads=2, bias=False) to print the exact shape at every step.

Commit before you compute: what does Trace: d_model=8, h=2, d_k=4, seq=4, batch=1 come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.

Correct: output (1,4,8) — same shape as input; attention weights (1,2,4,4) — one 4×4 matrix per head

Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. MHA preserves the (B,T,d_model) shape through the full projection cycle.

15. Trace: d_model=8, h=2, d_k=4, seq=4, batch=1

Worked example

Use nn.MultiheadAttention(embed_dim=8, num_heads=2, bias=False) to print the exact shape at every step.

import torch, torch.nn as nn, math
torch.manual_seed(42)
D, H, T = 8, 2, 4
DK = D // H   # d_k = 4
mha = nn.MultiheadAttention(D, H, batch_first=True, bias=False)
X = torch.randn(1, T, D)          # (B=1, T=4, D=8)
out, aw = mha(X, X, X, need_weights=True, average_attn_weights=False)
print('input :', X.shape)           # (1,4,8)
print('output:', out.shape)         # (1,4,8)
print('attn  :', aw.shape)          # (1,2,4,4) = (B,h,T,T)
print('params:', sum(p.numel() for p in mha.parameters()),
      '== 4*8^2 =', 4*D**2)

output (1,4,8) — same shape as input; attention weights (1,2,4,4) — one 4×4 matrix per head

Why: MHA preserves the (B,T,d_model) shape through the full projection cycle. The per-head attention maps are (T×T) softmax distributions.

tensorshapenote
X (input)(1, 4, 8)B=1, T=4, d_model=8
Q, K, V(1, 4, 8)after W_Q/W_K/W_V
Q_heads, K_heads(1, 2, 4, 4)B, h, T, d_k after split
attn_weights(1, 2, 4, 4)one (T×T) per head
context(1, 2, 4, 4)attn @ V_heads
concat(1, 4, 8)merge h=2 heads
output(1, 4, 8)after W_O — matches input shape

16. Which is which, by shape

Discrimination

Sort into buckets

Sort these by shape, from memory, without looking back at Trace: d_model=8, h=2, d_k=4, seq=4, batch=1. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.

(1, 4, 8)
X (input); Q, K, V; concat; output
(1, 2, 4, 4)
Q_heads, K_heads; attn_weights; context
g1
shape is "(1, 4, 8)" for X (input), Q, K, V, concat, output — that is what the table on "Trace: d_model=8, h=2, d_k=4, seq=4…" records, and it is the single property separating this group from the rest.
g2
shape is "(1, 2, 4, 4)" for Q_heads, K_heads, attn_weights, context — that is what the table on "Trace: d_model=8, h=2, d_k=4, seq=4…" records, and it is the single property separating this group from the rest.

17. Scaling keeps computation constant

Concept

Each head works in a d_k = d_model / h dimensional space. With h=8 and d_model=512, each head uses d_k=64 instead of 512. The total work across all h heads equals one full-size single-head attention.

\[ h \times \underbrace{d_k^2}_{\text{per head QK}^\top} = h \times \left(\frac{d_{\text{model}}}{h}\right)^2 = \frac{d_{\text{model}}^2}{h} \]

The h heads run in parallel on modern hardware, so the wall-clock time is the same as single-head. You get h perspectives for free — the key architectural insight of the Transformer paper.

18. Teach it back: Scaling keeps computation constant

Explain it

Discussion prompt

Explain Scaling keeps computation constant to a student a year behind you. No notation, no jargon they have not met — and it still has to be true.

Hint: If your explanation needs a symbol they have never seen, you are describing the notation rather than the idea.

Answer:

The h heads run in parallel on modern hardware, so the wall-clock time is the same as single-head. You get h perspectives for free — the key architectural insight of the Transformer paper.

19. Parameter count proof

Section

Part 3 of 4

20. Parameter count is always 4·d_model²

Concept

The h per-head matrices W^Q_i, W^K_i, W^V_i each have shape (d_model, d_k). Stacked, each set forms a single (d_model, d_model) matrix. Plus W^O.

\[ \underbrace{h \cdot d_{\text{model}} \cdot d_k}_{W_Q} + \underbrace{h \cdot d_{\text{model}} \cdot d_k}_{W_K} + \underbrace{h \cdot d_{\text{model}} \cdot d_k}_{W_V} + \underbrace{d_{\text{model}}^2}_{W_O} = 4 \, d_{\text{model}}^2 \]

Since d_k = d_model / h, each W term is h · d_model · (d_model/h) = d_model². So 3·d_model² + d_model² = 4·d_model² — independent of h.

21. By analogy: Parameter count is always 4·d_model²

Analogy

Discussion prompt

Explain Parameter count is always 4·d_model² by analogy to something with no Machine Learning in it at all — a queue, a recipe, a map, a bank balance, whatever fits. Then say where your analogy breaks.

Hint: An analogy that never breaks is not an analogy, it is the same idea wearing a hat. Find the seam — that is the part that is actually new.

Answer:

The h per-head matrices W^Q_i, W^K_i, W^V_i each have shape (d_model, d_k). Stacked, each set forms a single (d_model, d_model) matrix. Plus W^O.

22. Guess the shape of the answer: Verify param count for any h

Estimation

Predict first

Run the scratch MultiHeadAttention module (code you'll write in the project) for h=1, 2, 4, 8, 16 and confirm the count never changes.

Commit before you compute: what does Verify param count for any h come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.

Correct: Every row prints 1,048,576 — the count never changes with h

Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. The split into per-head matrices is equivalent to a single (d_model, d_model) projection.

23. Verify param count for any h

Worked example

Run the scratch MultiHeadAttention module (code you'll write in the project) for h=1, 2, 4, 8, 16 and confirm the count never changes.

import torch.nn as nn, math
class MHA(nn.Module):
    def __init__(self, d, h):
        super().__init__()
        self.W_Q = nn.Linear(d, d, bias=False)
        self.W_K = nn.Linear(d, d, bias=False)
        self.W_V = nn.Linear(d, d, bias=False)
        self.W_O = nn.Linear(d, d, bias=False)
        self.h = h; self.d_k = d // h
    def forward(self, x):
        pass  # see project section

d = 512
for h in [1, 2, 4, 8, 16]:
    m = MHA(d, h)
    p = sum(x.numel() for x in m.parameters())
    print(f'h={h:2d}  d_k={d//h:3d}  params={p:,}')

Every row prints 1,048,576 — the count never changes with h

Why: The split into per-head matrices is equivalent to a single (d_model, d_model) projection. h is a structural choice, not a parameter-budget choice.

hd_kparams (d=512)
15121,048,576
22561,048,576
41281,048,576
8641,048,576
16321,048,576

24. What each one costs: Verify param count for any h

Trade off

Comparison matrix

From Verify param count for any h: every row here is a choice with a cost. Fill the params (d=512) column, then say which row you would actually pick and what you give up for it.

hd_kparams (d=512)
15121,048,576
22561,048,576
41281,048,576
8641,048,576
16321,048,576

25. Something is wrong here: thinking more heads = more parameters

Anomaly

Predict first

A student writes this, and it looks reasonable:

With h=8 heads, each head has its own W^Q_i, W^K_i, W^V_i — that's 3×8=24 weight matrices, so MHA with h=8 has 8× the parameters of h=1.

It is wrong. Say what breaks — and say it before you turn the page.

Correct: Each per-head matrix is (d_model, d_k) = (512, 64), not (512, 512).

More heads = smaller per-head dimension, not more total parameters. d_k = d_model / h.

Why: Each per-head matrix is (d_model, d_k) = (512, 64), not (512, 512). Eight of them stack to the same total as one (512, 512) matrix. The parameter count is 4·512² regardless of h.

26. Trap: thinking more heads = more parameters

Trap

The trap

With h=8 heads, each head has its own W^Q_i, W^K_i, W^V_i — that's 3×8=24 weight matrices, so MHA with h=8 has 8× the parameters of h=1.

Claim MHA(d=512, h=8) has 8× the params of MHA(d=512, h=1)

Why: Each per-head matrix is (d_model, d_k) = (512, 64), not (512, 512). Eight of them stack to the same total as one (512, 512) matrix. The parameter count is 4·512² regardless of h.

The fix

More heads = smaller per-head dimension, not more total parameters. d_k = d_model / h.

h=1: one (512, 512) W_Q; h=8: eight (512, 64) W^Q_i — same product, same params

Why: h·(d_model·d_k) = h·d_model·(d_model/h) = d_model². Doubling h halves d_k — the total is always 4·d_model².

27. Break it on purpose: thinking more heads = more parameters

Break the constraint

Discussion prompt

The rule this trap just fixed:

More heads = smaller per-head dimension, not more total parameters. d_k = d_model / h.

Now break it on purpose. Build a case that violates it and follow the consequences until something visibly fails. Where does the failure first show up — and would you have noticed it if you had not been looking?

Hint: The dangerous rules are the ones whose violation still produces an answer. If yours fails loudly, try to find one that fails quietly.

Answer:

Each per-head matrix is (d_model, d_k) = (512, 64), not (512, 512). Eight of them stack to the same total as one (512, 512) matrix. The parameter count is 4·512² regardless of h.

28. h=1 MHA is exactly single-head attention

Concept

When h=1, d_k=d_model. No splitting or concatenation occurs — the reshape and transpose are identity operations. The result is ordinary scaled dot-product attention with W_Q, W_K, W_V, W_O projections.

\[ \text{head}_1 = \text{Attention}(X W^Q_1,\; X W^K_1,\; X W^V_1) \quad \text{with } d_k = d_{\text{model}} \]

Verifiable in code: run torch.allclose(mha_h1(X), single_head(X)) after giving them identical weights — it returns True. (Lesson 80 special case.)

29. Where does each piece belong: Lesson 81: Multi-Head Attention

Sorting

Sort into buckets

These are the pieces of Lesson 81: Multi-Head Attention, out of order. Put each one back under the part of the lesson it belongs to.

Why multiple heads?
Single-head attention's blind spot; The multi-head idea; What each head learns
Dataflow — project, split, attend, concat
Step-by-step dataflow; Trace: d_model=8, h=2, d_k=4, seq=4, batch=1; Scaling keeps computation constant
Parameter count proof
Parameter count is always 4·d_model²; Verify param count for any h; h=1 MHA is exactly single-head attention
s1
Why multiple heads? is where Lesson 81: Multi-Head Attention puts Single-head attention's blind spot, The multi-head idea, What each head learns. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.
s2
Dataflow — project, split, attend, concat is where Lesson 81: Multi-Head Attention puts Step-by-step dataflow, Trace: d_model=8, h=2, d_k=4, seq=4, batch=1, Scaling keeps computation constant. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.
s3
Parameter count proof is where Lesson 81: Multi-Head Attention puts Parameter count is always 4·d_model², Verify param count for any h, h=1 MHA is exactly single-head attention. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.

30. Something is wrong here: wrong reshape order for head splitting

Anomaly

Predict first

A student writes this, and it looks reasonable:

To split Q of shape (B, T, D) into h heads, reshape to (B, T·h, d_k) — grouping consecutive time steps together per head.

It is wrong. Say what breaks — and say it before you turn the page.

Correct: This scrambles positions: head 0 gets tokens [0, h, 2h, ...] instead of all tokens.

Reshape to (B, T, h, d_k) then .transpose(1, 2) to (B, h, T, d_k).

Why: This scrambles positions: head 0 gets tokens [0, h, 2h, ...] instead of all tokens. You must reshape to (B, T, h, d_k) first, THEN transpose to (B, h, T, d_k) so each head sees the full sequence.

31. Trap: wrong reshape order for head splitting

Trap

The trap

To split Q of shape (B, T, D) into h heads, reshape to (B, T·h, d_k) — grouping consecutive time steps together per head.

Q.view(B, T*h, d_k) — interleaves positions across heads

Why: This scrambles positions: head 0 gets tokens [0, h, 2h, ...] instead of all tokens. You must reshape to (B, T, h, d_k) first, THEN transpose to (B, h, T, d_k) so each head sees the full sequence.

The fix

Reshape to (B, T, h, d_k) then .transpose(1, 2) to (B, h, T, d_k).

Q.view(B, T, h, d_k).transpose(1, 2) — each head gets all T positions

Why: The reshape packs d_k features per head per position. The transpose brings the head axis before the sequence axis so batched matmul treats each head independently over the full T-length sequence.

32. Which of these survive contact with Lesson 81: Multi-Head Attention?

Two truths and a lie

Sort into buckets

Some of these hold up and some are the exact mistakes this lesson is built to prevent. Sort them.

Holds up
Multi-head attention runs h independent attention functions in parallel, each in a lower-dimensional subspace, then merges the results.; Steps 2–4 are implemented as a single batched matrix multiply — PyTorch handles all h heads simultaneously. No Python loop over heads needed.; The h per-head matrices W^Q_i, W^K_i, W^V_i each have shape (d_model, d_k). Stacked, each set forms a single (d_model, d_model) matrix. Plus W^O.
Breaks
With h=8 heads, each head has its own W^Q_i, W^K_i, W^V_i — that's 3×8=24 weight matrices, so MHA with h=8 has 8× the parameters of h=1.; To split Q of shape (B, T, D) into h heads, reshape to (B, T·h, d_k) — grouping consecutive time steps together per head.
sound
These are stated as this lesson states them — each one survives the edge cases Lesson 81: Multi-Head Attention puts it through.
flawed
Each of these is lifted from a trap in this deck: reasonable-sounding, and wrong in a way that only shows up once you rely on it.

33. Rebuild the recipe: The multi-head attention recipe

Ranking

Put in order

These are the steps of The multi-head attention recipe, scrambled. Put them back in order before the next slide shows you.

  1. Project: Q=X@W_Q, K=X@W_K, V=X@W_V — shapes (B, T, d_model) each
  2. Split: reshape(B, T, h, d_k).transpose(1,2) → (B, h, T, d_k)
  3. Score: (Q_h @ K_h.T) / sqrt(d_k) → (B, h, T, T)
  4. Attend: softmax(scores, dim=-1) @ V_h → (B, h, T, d_k) context
  5. Merge: transpose(1,2).view(B, T, d_model) then @ W_O
  6. Count: 4 · d_model² parameters, invariant to h

Why: This is the order the recipe itself gives. Recalling the sequence without the slide in front of you is the difference between recognising the method and being able to run it — most of what goes wrong in practice is a step done out of turn.

34. The multi-head attention recipe

Pattern

  1. Project: Q=X@W_Q, K=X@W_K, V=X@W_V — shapes (B, T, d_model) each
  2. Split: reshape(B, T, h, d_k).transpose(1,2) → (B, h, T, d_k)
  3. Score: (Q_h @ K_h.T) / sqrt(d_k) → (B, h, T, T)
  4. Attend: softmax(scores, dim=-1) @ V_h → (B, h, T, d_k) context
  5. Merge: transpose(1,2).view(B, T, d_model) then @ W_O
  6. Count: 4 · d_model² parameters, invariant to h

35. Where does it stop working: The multi-head attention recipe

Edge cases

Discussion prompt

The multi-head attention recipe works on the cases you have just seen. Push it to the edge: what is the most degenerate input it still handles — empty, zero, one item, everything equal — and what is the first case where it stops being true? Name the case, not just "it breaks".

Hint: Try the smallest legal input, then the largest, then the one where two things collide. Methods are specified at their edges; the middle takes care of itself.

Answer:

  1. Project: Q=X@W_Q, K=X@W_K, V=X@W_V — shapes (B, T, d_model) each
  2. Split: reshape(B, T, h, d_k).transpose(1,2) → (B, h, T, d_k)
  3. Score: (Q_h @ K_h.T) / sqrt(d_k) → (B, h, T, T)
  4. Attend: softmax(scores, dim=-1) @ V_h → (B, h, T, d_k) context
  5. Merge: transpose(1,2).view(B, T, d_model) then @ W_O
  6. Count: 4 · d_model² parameters, invariant to h

36. Rule out three: Check yourself — parameter count

Elimination

Eliminate the wrong options

A MultiHeadAttention layer has d_model=256 and h=8 (bias=False). How many learnable parameters does it have?

3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.

  • A. 262,144 (4 · 256²)
  • B. 32,768 (4 · 256² / 8)
  • C. 2,097,152 (4 · 256² · 8)
  • D. 65,536 (256² for W_O only)

Survives elimination: A

Why: Total = 4·d_model² = 4·256² = 262,144. Each of W_Q, W_K, W_V, W_O is (d_model, d_model); splitting into h heads does not change the total because d_k=d_model/h and h·d_k=d_model.

37. Check yourself — parameter count

Check

Work it out before clicking.

Check your understanding

A MultiHeadAttention layer has d_model=256 and h=8 (bias=False). How many learnable parameters does it have?

  • A. 262,144 (4 · 256²) (correct)
  • B. 32,768 (4 · 256² / 8)
  • C. 2,097,152 (4 · 256² · 8)
  • D. 65,536 (256² for W_O only)

Answer: A

Why: Total = 4·d_model² = 4·256² = 262,144. Each of W_Q, W_K, W_V, W_O is (d_model, d_model); splitting into h heads does not change the total because d_k=d_model/h and h·d_k=d_model.

Why B tempts people
Divides by h=8 as if smaller heads mean proportionally fewer total params — but h·(d_model·d_k)=d_model² regardless of h.
Why C tempts people
Multiplies by h=8, treating each head as a full (d_model, d_model) matrix rather than a (d_model, d_k) slice.
Why D tempts people
Counts only W_O — forgets W_Q, W_K, W_V, each of which contributes another d_model².

38. Answer it before you see the options: Check yourself — attention weight shape

Prediction

Predict first

In a MultiHeadAttention layer with B=4, T=12, d_model=64, h=8, what is the shape of the per-head attention weight tensor returned by the layer (before averaging)?

Answer it in your own words, now, with nothing to choose from. The options are on the next slide — and picking the right one off a list is an easier skill than producing it.

Correct: (4, 8, 12, 12)

Why: Shape is (B, h, T, T) = (4, 8, 12, 12). Each of the h=8 heads produces a (T×T)=(12×12) softmax attention matrix per sample in the batch.

39. Check yourself — attention weight shape

Check

Trace the shape through the attention computation.

Check your understanding

In a MultiHeadAttention layer with B=4, T=12, d_model=64, h=8, what is the shape of the per-head attention weight tensor returned by the layer (before averaging)?

  • A. (4, 8, 12, 12) (correct)
  • B. (4, 12, 12)
  • C. (4, 8, 12, 8)
  • D. (4, 12, 64)

Answer: A

Why: Shape is (B, h, T, T) = (4, 8, 12, 12). Each of the h=8 heads produces a (T×T)=(12×12) softmax attention matrix per sample in the batch.

Why B tempts people
(4,12,12) is the average-over-heads shape that nn.MultiheadAttention returns when average_attn_weights=True (the default) — it collapses the h dimension.
Why C tempts people
(4,8,12,8) confuses T=12 with d_k=d_model/h=8 in the last dimension — the attention matrix rows index over query positions (T), columns over key positions (T), not d_k.
Why D tempts people
(4,12,64) has the shape of the output tensor, not the attention weights.

40. Rule out three: Check yourself — reshape order

Elimination

Eliminate the wrong options

Q has shape (B, T, d_model). After projecting with W_Q, which code correctly splits it into h heads of dimension d_k=d_model/h?

3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.

  • A. Q.view(B, T, h, d_k).transpose(1, 2)
  • B. Q.view(B, T*h, d_k)
  • C. Q.view(B, h, T, d_k)
  • D. Q.transpose(1, 2).view(B, h, T, d_k)

Survives elimination: A

Why: view(B, T, h, d_k) packs d_k features per head per position, then transpose(1,2) moves the head axis before the sequence axis to give (B, h, T, d_k). Each head sees all T positions with d_k features.

41. Check yourself — reshape order

Check

One line of code — one correct answer.

Check your understanding

Q has shape (B, T, d_model). After projecting with W_Q, which code correctly splits it into h heads of dimension d_k=d_model/h?

  • A. Q.view(B, T, h, d_k).transpose(1, 2) (correct)
  • B. Q.view(B, T*h, d_k)
  • C. Q.view(B, h, T, d_k)
  • D. Q.transpose(1, 2).view(B, h, T, d_k)

Answer: A

Why: view(B, T, h, d_k) packs d_k features per head per position, then transpose(1,2) moves the head axis before the sequence axis to give (B, h, T, d_k). Each head sees all T positions with d_k features.

Why B tempts people
view(B, T*h, d_k) merges T and h into one axis — subsequent matmuls would treat the merged axis as the sequence length, mixing head and position indices.
Why C tempts people
view(B, h, T, d_k) directly is wrong: view reinterprets contiguous memory in row-major order, placing head-0's first position's d_k values correctly, but the remaining values are drawn from the wrong memory locations — use view+transpose, not direct view to (B,h,T,d_k).
Why D tempts people
Transposing first gives (B, d_model, T), then view(B, h, T, d_k) would split the d_model axis over the T dimension — the shape may look right but the data is scrambled.

42. Your turn: implement MHA

Section

Project

43. Project: MultiHeadAttention as nn.Module

Concept

Build MultiHeadAttention from scratch — four nn.Linear layers, one forward method. Three milestones: skeleton + param count → forward pass + shapes → h=1 equivalence.

#milestonekey check
1Define module skeleton, verify params = 4·d_model²sum(p.numel() for p in model.parameters())
2Complete forward: project → split → attend → concat → projectoutput.shape == (B, T, d_model)
3h=1 equivalence: confirm MHA(h=1) == single-head attentiontorch.allclose(out_mha, out_single)

Build rules: one nn.Linear per weight matrix (no bias); use .view() + .transpose() for splits; no Python loop over heads — let batched matmul handle h heads simultaneously.

44. Break it if you can: Project: MultiHeadAttention as nn.Module

Counterexample

Discussion prompt

Build MultiHeadAttention from scratch — four nn.Linear layers, one forward method. Three milestones: skeleton + param count → forward pass + shapes → h=1 equivalence.

That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.

Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.

Answer:

Build rules: one nn.Linear per weight matrix (no bias); use .view() + .transpose() for splits; no Python loop over heads — let batched matmul handle h heads simultaneously.

45. Milestone 1 — module skeleton + param count

Worked example

Your turn: define the __init__ with four nn.Linear layers. Predict the param count for d_model=32, h=4.

Hint: nn.Linear(d_model, d_model, bias=False) — each matrix is (d_model, d_model); total = 4·d_model².

import torch.nn as nn
class MultiHeadAttention(nn.Module):
    def __init__(self, d_model, h):
        super().__init__()
        assert d_model % h == 0
        self.h = h
        self.d_k = d_model // h
        self.W_Q = nn.Linear(d_model, d_model, bias=False)
        self.W_K = nn.Linear(d_model, d_model, bias=False)
        self.W_V = nn.Linear(d_model, d_model, bias=False)
        self.W_O = nn.Linear(d_model, d_model, bias=False)
    def forward(self, Q, K, V): pass  # next milestone

m = MultiHeadAttention(32, 4)
print(sum(p.numel() for p in m.parameters()),
      '== 4*32^2 =', 4*32**2)
d_modelhexpected paramsactual params
82256256
3244,0964,096
51281,048,5761,048,576

46. Fill in: expected params for Milestone 1 — module skeleton + param count

Comparison

Comparison matrix

From Milestone 1 — module skeleton + param count: refill the expected params column from what you know. The rest of the table is as it appeared.

d_modelhexpected paramsactual params
82256256
3244,0964,096
51281,048,5761,048,576

47. Milestone 2 — complete forward pass

Worked example

Your turn: implement forward. Predict the shape of attn_w for B=2, T=5, d_model=8, h=2.

Hint: project → view(B, T, h, d_k).transpose(1,2) → score /sqrt(d_k) → softmax dim=-1 → @ V_h → transpose(1,2).contiguous().view(B, T, d_model) → W_O.

import torch, torch.nn.functional as F, math
class MultiHeadAttention(nn.Module):
    def __init__(self, d_model, h):
        super().__init__()
        assert d_model % h == 0
        self.h = h; self.d_k = d_model // h
        self.W_Q = nn.Linear(d_model, d_model, bias=False)
        self.W_K = nn.Linear(d_model, d_model, bias=False)
        self.W_V = nn.Linear(d_model, d_model, bias=False)
        self.W_O = nn.Linear(d_model, d_model, bias=False)
    def forward(self, Q, K, V):
        B, T, D = Q.shape
        def split(x): return x.view(B,T,self.h,self.d_k).transpose(1,2)
        q,k,v = split(self.W_Q(Q)), split(self.W_K(K)), split(self.W_V(V))
        scores = q @ k.transpose(-2,-1) / math.sqrt(self.d_k)
        attn_w = F.softmax(scores, dim=-1)
        ctx = (attn_w @ v).transpose(1,2).contiguous().view(B,T,D)
        return self.W_O(ctx), attn_w

torch.manual_seed(99)
model = MultiHeadAttention(8, 2)
out, aw = model(torch.randn(2,5,8), torch.randn(2,5,8), torch.randn(2,5,8))
print('output:', out.shape, '  attn_w:', aw.shape)
tensorexpected shapeactual shape
output(2, 5, 8)(2, 5, 8)
attn_w(2, 2, 5, 5)(2, 2, 5, 5)
params4·8²=256256

48. What each one costs: Milestone 2 — complete forward pass

Trade off

Comparison matrix

From Milestone 2 — complete forward pass: every row here is a choice with a cost. Fill the expected shape column, then say which row you would actually pick and what you give up for it.

tensorexpected shapeactual shape
output(2, 5, 8)(2, 5, 8)
attn_w(2, 2, 5, 5)(2, 2, 5, 5)
params4·8²=256256

49. Milestone 3 — h=1 equivalence

Worked example

Your turn: set h=1 in your module and confirm its output equals single-head attention with identical weights. Predict: does torch.allclose return True?

Hint: manually copy weights from your h=1 MHA into standalone W_Q/W_K/W_V/W_O matrices and run single-head attention. Both paths share the same computation.

import torch, torch.nn.functional as F, math
torch.manual_seed(5)
D = 8
mha_h1 = MultiHeadAttention(D, h=1)
X = torch.randn(1, 4, D)
out_mha, _ = mha_h1(X, X, X)
# Single-head attention using the same weights
WQ = mha_h1.W_Q.weight.T    # (d_model, d_model) in .T form
WK = mha_h1.W_K.weight.T
WV = mha_h1.W_V.weight.T
WO = mha_h1.W_O.weight.T
Q = X @ WQ; K = X @ WK; V = X @ WV
scores = Q @ K.transpose(-2,-1) / math.sqrt(D)
out_single = F.softmax(scores, dim=-1) @ V @ WO
print('allclose:', torch.allclose(out_mha, out_single, atol=1e-5))
checkresult
out_mha shape(1, 4, 8)
out_single shape(1, 4, 8)
torch.allcloseTrue
conclusionh=1 MHA ≡ single-head attention

50. Which is which, by result

Discrimination

Sort into buckets

Sort these by result, from memory, without looking back at Milestone 3 — h=1 equivalence. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.

(1, 4, 8)
out_mha shape; out_single shape
True
torch.allclose
h=1 MHA ≡ single-head attention
conclusion
g1
result is "(1, 4, 8)" for out_mha shape, out_single shape — that is what the table on "Milestone 3 — h=1 equivalence" records, and it is the single property separating this group from the rest.
g2
result is "True" for torch.allclose — that is what the table on "Milestone 3 — h=1 equivalence" records, and it is the single property separating this group from the rest.
g3
result is "h=1 MHA ≡ single-head attention" for conclusion — that is what the table on "Milestone 3 — h=1 equivalence" records, and it is the single property separating this group from the rest.

51. The full program

Concept

import torch, torch.nn as nn, torch.nn.functional as F, math

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model, h):
        super().__init__()
        assert d_model % h == 0
        self.h = h; self.d_k = d_model // h
        self.W_Q = nn.Linear(d_model, d_model, bias=False)
        self.W_K = nn.Linear(d_model, d_model, bias=False)
        self.W_V = nn.Linear(d_model, d_model, bias=False)
        self.W_O = nn.Linear(d_model, d_model, bias=False)
    def forward(self, Q, K, V):
        B, T, D = Q.shape
        def split(x): return x.view(B,T,self.h,self.d_k).transpose(1,2)
        q,k,v = split(self.W_Q(Q)), split(self.W_K(K)), split(self.W_V(V))
        aw = F.softmax(q@k.transpose(-2,-1)/math.sqrt(self.d_k), dim=-1)
        ctx = (aw@v).transpose(1,2).contiguous().view(B,T,D)
        return self.W_O(ctx), aw

torch.manual_seed(0)
for d, h in [(8,2), (512,8)]:
    m = MultiHeadAttention(d, h)
    p = sum(x.numel() for x in m.parameters())
    x = torch.randn(2, 6, d)
    out, aw = m(x, x, x)
    print(f'd={d:3d} h={h} params={p:,} out={tuple(out.shape)} aw={tuple(aw.shape)}')
d_modelhparamsoutput shapeattn shape
82256(2, 6, 8)(2, 2, 6, 6)
51281,048,576(2, 6, 512)(2, 8, 6, 6)

This is the exact module used inside every Transformer encoder and decoder layer — from BERT to GPT to T5. The Transformer paper (Vaswani et al. 2017) uses d_model=512, h=8, d_k=64 as the base configuration.

52. Fill in: attn shape for The full program

Comparison

Comparison matrix

From The full program: refill the attn shape column from what you know. The rest of the table is as it appeared.

d_modelhparamsoutput shapeattn shape
82256(2, 6, 8)(2, 2, 6, 6)
51281,048,576(2, 6, 512)(2, 8, 6, 6)

53. Show it off

Concept

Out loud, slides closed: (1) explain the five dataflow steps of MHA — project, split, attend, concat, project; (2) prove that the parameter count is 4·d_model² for any h; (3) state why h=1 MHA is the same as single-head attention.

Stretch (homework): visualize attention heatmaps from different heads on a real sentence (torchtext or manual tokenization); prove the parameter count algebraically on paper; add masking to make your module work as a decoder self-attention layer. Next: Lesson 82 — the full Transformer encoder block (MHA + LayerNorm + FFN residual connections).

54. Connect it up: Lesson 81: Multi-Head Attention

Connect it up

Draw it

One page, no notation unless you need it: draw how these connect — Why multiple heads? · Dataflow — project, split, attend, concat · Parameter count proof · Your turn: implement MHA. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.

55. What you can do now

Recap

conceptthe one thing to remember
why multi-headh parallel subspaces learn diverse relationship types simultaneously
head dimensiond_k = d_model / h — smaller heads, same total work
parameter count4·d_model² — always, regardless of h
split orderview(B,T,h,d_k).transpose(1,2) — reshape before transpose
h=1 special casedegenerates to single-head attention — Lesson 80 is a special case

Sources

  1. USAAIO Year-Long Master Lesson Plan, Lesson 81 (Week 28 — Transformer Architecture: Multi-Head Attention) — Barron · USAAIO Round 2 Preparation, 2026
  2. MultiHeadAttention parameter count, attention scores, softmax, concat/project shapes verified with torch 2.7.1+cpu and numpy 2.2.6, June 2026 — Real execution, verified

Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.

Book on Wyzant · Text (657) 465-8108