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
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.
Objectives
MultiHeadAttention from scratch as nn.Module with correct shapesWarm-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.
Section
Part 1 of 4
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.
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.
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) \]
| symbol | shape | role |
|---|---|---|
| W^Q_i, W^K_i, W^V_i | (d_model, d_k) | per-head projection matrices |
| d_k = d_model / h | scalar | head dimension — keeps total compute fixed |
| W^O | (d_model, d_model) | output projection merges all heads |
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.
| symbol | shape | role |
|---|---|---|
| W^Q_i, W^K_i, W^V_i | (d_model, d_k) | per-head projection matrices |
| d_k = d_model / h | scalar | head dimension — keeps total compute fixed |
| W^O | (d_model, d_model) | output projection merges all heads |
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.
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.
Section
Part 2 of 4
Concept
Q = X·W_Q, K = X·W_K, V = X·W_V — shapes (B, T, d_model)(B, T, h, d_k) then transpose to (B, h, T, d_k)scores_i = Q_i·K_i^T / sqrt(d_k), softmax, then ·V_i(B, T, h, d_k), reshape to (B, T, d_model)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.
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.
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.
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.
| tensor | shape | note |
|---|---|---|
| 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 |
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.
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.
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.
Section
Part 3 of 4
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.
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.
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.
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.
| h | d_k | params (d=512) |
|---|---|---|
| 1 | 512 | 1,048,576 |
| 2 | 256 | 1,048,576 |
| 4 | 128 | 1,048,576 |
| 8 | 64 | 1,048,576 |
| 16 | 32 | 1,048,576 |
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.
| h | d_k | params (d=512) |
|---|---|---|
| 1 | 512 | 1,048,576 |
| 2 | 256 | 1,048,576 |
| 4 | 128 | 1,048,576 |
| 8 | 64 | 1,048,576 |
| 16 | 32 | 1,048,576 |
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.
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.
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².
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.
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.)
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.
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.
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.
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.
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.
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.
Q=X@W_Q, K=X@W_K, V=X@W_V — shapes (B, T, d_model) eachreshape(B, T, h, d_k).transpose(1,2) → (B, h, T, d_k)(Q_h @ K_h.T) / sqrt(d_k) → (B, h, T, T)softmax(scores, dim=-1) @ V_h → (B, h, T, d_k) contexttranspose(1,2).view(B, T, d_model) then @ W_O4 · d_model² parameters, invariant to hWhy: 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.
Pattern
Q=X@W_Q, K=X@W_K, V=X@W_V — shapes (B, T, d_model) eachreshape(B, T, h, d_k).transpose(1,2) → (B, h, T, d_k)(Q_h @ K_h.T) / sqrt(d_k) → (B, h, T, T)softmax(scores, dim=-1) @ V_h → (B, h, T, d_k) contexttranspose(1,2).view(B, T, d_model) then @ W_O4 · d_model² parameters, invariant to hEdge 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:
Q=X@W_Q, K=X@W_K, V=X@W_V — shapes (B, T, d_model) eachreshape(B, T, h, d_k).transpose(1,2) → (B, h, T, d_k)(Q_h @ K_h.T) / sqrt(d_k) → (B, h, T, T)softmax(scores, dim=-1) @ V_h → (B, h, T, d_k) contexttranspose(1,2).view(B, T, d_model) then @ W_O4 · d_model² parameters, invariant to hElimination
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.
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.
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?
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.
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.
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)?
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.
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.
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.
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?
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.
Section
Project
Concept
Build MultiHeadAttention from scratch — four nn.Linear layers, one forward method. Three milestones: skeleton + param count → forward pass + shapes → h=1 equivalence.
| # | milestone | key check |
|---|---|---|
| 1 | Define module skeleton, verify params = 4·d_model² | sum(p.numel() for p in model.parameters()) |
| 2 | Complete forward: project → split → attend → concat → project | output.shape == (B, T, d_model) |
| 3 | h=1 equivalence: confirm MHA(h=1) == single-head attention | torch.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.
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.
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_model | h | expected params | actual params |
|---|---|---|---|
| 8 | 2 | 256 | 256 |
| 32 | 4 | 4,096 | 4,096 |
| 512 | 8 | 1,048,576 | 1,048,576 |
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_model | h | expected params | actual params |
|---|---|---|---|
| 8 | 2 | 256 | 256 |
| 32 | 4 | 4,096 | 4,096 |
| 512 | 8 | 1,048,576 | 1,048,576 |
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)| tensor | expected shape | actual shape |
|---|---|---|
| output | (2, 5, 8) | (2, 5, 8) |
| attn_w | (2, 2, 5, 5) | (2, 2, 5, 5) |
| params | 4·8²=256 | 256 |
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.
| tensor | expected shape | actual shape |
|---|---|---|
| output | (2, 5, 8) | (2, 5, 8) |
| attn_w | (2, 2, 5, 5) | (2, 2, 5, 5) |
| params | 4·8²=256 | 256 |
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))| check | result |
|---|---|
| out_mha shape | (1, 4, 8) |
| out_single shape | (1, 4, 8) |
| torch.allclose | True |
| conclusion | h=1 MHA ≡ single-head attention |
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.
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_model | h | params | output shape | attn shape |
|---|---|---|---|---|
| 8 | 2 | 256 | (2, 6, 8) | (2, 2, 6, 6) |
| 512 | 8 | 1,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.
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_model | h | params | output shape | attn shape |
|---|---|---|---|---|
| 8 | 2 | 256 | (2, 6, 8) | (2, 2, 6, 6) |
| 512 | 8 | 1,048,576 | (2, 6, 512) | (2, 8, 6, 6) |
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).
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.
Recap
MultiHeadAttention as nn.Module with correct view/transpose ordering| concept | the one thing to remember |
|---|---|
| why multi-head | h parallel subspaces learn diverse relationship types simultaneously |
| head dimension | d_k = d_model / h — smaller heads, same total work |
| parameter count | 4·d_model² — always, regardless of h |
| split order | view(B,T,h,d_k).transpose(1,2) — reshape before transpose |
| h=1 special case | degenerates to single-head attention — Lesson 80 is a special case |
Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.