USAAIO Lesson 111, the Phase 3 capstone review. It derives scaled dot-product attention and analyzes its complexity, covers sinusoidal positional encoding and BPE tokenization step by step, and implements and debugs multi-head attention. It then covers the IoU formula and its computation, the U-Net skip-connection architecture, and a comparison of ViT with CNNs. All the values were computed with torch 2.7.1+cpu and numpy 2.2.6 on synthetic data. The lesson runs to 28 slides.
Subject: Machine Learning · 56 slides · code lesson
Open the interactive version of this deck · Homework for this lesson
Title
USAAIO · Lesson 111 · Week 38
Capstone review before Phase 4 (generative models). Flash-derive attention, trace BPE, compute IoU, debug multi-head attention, compare ViT vs CNN — all exam-level, all numbers real.
Objectives
sqrt(d_k) factor is non-optionalPE[pos, 2i] and PE[pos, 2i+1] for concrete (pos, i) valuesWarm-up
Discussion prompt
Before we open Lesson 111: Transformers, NLP & CV — Phase 3 Review: without looking back, what was the main idea of Semantic Segmentation & U-Net, 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:
pixel-wise classification via semantic segmentation, U-Net encoder-decoder with skip connections (concat not add), upsampling methods (bilinear, ConvTranspose2d, pixel shuffle), Dice loss and Focal loss for class-imbalanced masks, and TinyUNet trained from scratch.
Section
Part 1 of 4
Concept
Every transformer layer (Lessons 88-92) is built on one core operation. Given queries Q, keys K, values V (all (T, d_k)):
\[ \text{Attention}(Q, K, V) = \text{softmax}\!\left(\frac{QK^\top}{\sqrt{d_k}}\right) V \]
QK^T: shape (T, T) — each token's query dot-products with every token's key/ sqrt(d_k): prevents softmax saturation when d_k is large (verified below)softmax(...): rows sum to 1 — weights over the value sequence× V: weighted sum of value vectors, shape (T, d_k) — same as input| quantity | shape | meaning |
|---|---|---|
| Q (queries) | (T, d_k) | what each token is looking for |
| K (keys) | (T, d_k) | what each token advertises |
| V (values) | (T, d_k) | what each token sends if selected |
| QK^T | (T, T) | raw attention logits |
| softmax out | (T, T) | attention weights — rows sum to 1 |
| output | (T, d_k) | context-aware token representations |
Discrimination
Sort into buckets
Sort these by shape, from memory, without looking back at Scaled dot-product attention — the formula. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.
Estimation
Predict first
Three tokens, d_k = 4. Compute the attention matrix and read off the weight table. All values from torch.manual_seed(0) — verified.
Commit before you compute: what does Trace: 3-token attention step by step come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: scores[0] = [0.3803, 0.0517, -0.7319]; after softmax → [0.4881, 0.3514, 0.1605]
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. Token 0's query most strongly matches Key 0 (raw score 0.3803), so it borrows most from Value 0.
Worked example
Three tokens, d_k = 4. Compute the attention matrix and read off the weight table. All values from torch.manual_seed(0) — verified.
import torch
torch.manual_seed(0)
T, dk = 3, 4
Q = torch.randn(T, dk)
K = torch.randn(T, dk)
V = torch.randn(T, dk)
scores = Q @ K.T / (dk ** 0.5) # (3, 3)
weights = torch.softmax(scores, dim=-1)
out = weights @ V
print('scores:\n', scores.numpy().round(4))
print('weights:\n', weights.numpy().round(4))
print('row 0 sum:', weights[0].sum().item()) # 1.0
print('out shape:', out.shape)scores[0] = [0.3803, 0.0517, -0.7319]; after softmax → [0.4881, 0.3514, 0.1605]
Why: Token 0's query most strongly matches Key 0 (raw score 0.3803), so it borrows most from Value 0. Row sum = 1.0000 — softmax guarantee.
| token | raw score to tok 0 | raw score to tok 1 | raw score to tok 2 | weight on tok 0 | weight on tok 1 | weight on tok 2 |
|---|---|---|---|---|---|---|
| 0 | 0.3803 | 0.0517 | -0.7319 | 0.4881 | 0.3514 | 0.1605 |
| 1 | -0.4697 | -0.7661 | -0.7052 | 0.3947 | 0.2935 | 0.3119 |
| 2 | 0.4168 | 0.2572 | -0.6368 | 0.4543 | 0.3873 | 0.1584 |
Pattern
Step through it
Step through Trace: 3-token attention step by step one row at a time. What is driving the change, and what would the row after the last one be?
Concept
The QK^T matrix multiply costs O(T²·d_k) — this is the fundamental bottleneck of the transformer. Doubling sequence length quadruples the attention cost.
| T (seq len) | T² | rel cost vs T=128 |
|---|---|---|
| 128 | 16,384 | 1× |
| 512 | 262,144 | 16× |
| 1024 | 1,048,576 | 64× |
| 2048 | 4,194,304 | 256× |
ViT-B/16 on 224×224 images: T = 196 patches + 1 CLS = 197 tokens — 197² = 38,809 per head per layer. This is why long-context transformers (GPT-4, Gemini) need FlashAttention or linear approximations.
Comparison
Comparison matrix
From Attention complexity: O(T²d): refill the T² column from what you know. The rest of the table is as it appeared.
| T (seq len) | T² | rel cost vs T=128 |
|---|---|---|
| 128 | 16,384 | 1× |
| 512 | 262,144 | 16× |
| 1024 | 1,048,576 | 64× |
| 2048 | 4,194,304 | 256× |
Anomaly
Predict first
A student writes this, and it looks reasonable:
Compute raw attention without scaling: scores = Q @ K.T. Seems fine — softmax still produces a valid distribution.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: With d_k=64 the raw scores have std ≈ sqrt(64)=8, producing extreme logits.
Scale by 1/sqrt(d_k) before softmax: scores = Q @ K.T / (d_k ** 0.5).
Why: With d_k=64 the raw scores have std ≈ sqrt(64)=8, producing extreme logits. Verified: max weight → 1.000000, entropy → 0.36. The distribution collapses to near-one-hot — gradients through softmax vanish almost everywhere.
Trap
Compute raw attention without scaling: scores = Q @ K.T. Seems fine — softmax still produces a valid distribution.
weights = softmax(Q @ K.T) # no scale
Why: With d_k=64 the raw scores have std ≈ sqrt(64)=8, producing extreme logits. Verified: max weight → 1.000000, entropy → 0.36. The distribution collapses to near-one-hot — gradients through softmax vanish almost everywhere.
Scale by 1/sqrt(d_k) before softmax: scores = Q @ K.T / (d_k ** 0.5).
weights = softmax(Q @ K.T / (d_k ** 0.5)) # d_k=64 → scale=8.0
Why: Verified: max weight = 0.8209, entropy = 1.88. The distribution stays spread — all tokens receive meaningful gradients during backprop. The scale factor matches std(Q@K.T) ≈ sqrt(d_k) when Q,K entries are ~N(0,1).
Section
Part 2 of 4
Concept
Transformers are permutation-invariant — without injecting position, [tok_1, tok_2, tok_3] and [tok_3, tok_1, tok_2] produce identical outputs. The original 'Attention is All You Need' solution: add fixed sinusoids to each token embedding.
\[ \text{PE}[\text{pos},\, 2i] = \sin\!\left(\frac{\text{pos}}{10000^{2i/d}}\right), \quad \text{PE}[\text{pos},\, 2i+1] = \cos\!\left(\frac{\text{pos}}{10000^{2i/d}}\right) \]
| pos | dim 0 (sin) | dim 1 (cos) | dim 2 (sin) | dim 3 (cos) |
|---|---|---|---|---|
| 0 | 0.0000 | 1.0000 | 0.0000 | 1.0000 |
| 1 | 0.8415 | 0.5403 | 0.0998 | 0.9950 |
| 2 | 0.9093 | -0.4161 | 0.1987 | 0.9801 |
| 3 | 0.1411 | -0.9900 | 0.2955 | 0.9553 |
Trade off
Comparison matrix
From Sinusoidal positional encoding (Lesson 88 recap): every row here is a choice with a cost. Fill the dim 0 (sin) column, then say which row you would actually pick and what you give up for it.
| pos | dim 0 (sin) | dim 1 (cos) | dim 2 (sin) | dim 3 (cos) |
|---|---|---|---|---|
| 0 | 0.0000 | 1.0000 | 0.0000 | 1.0000 |
| 1 | 0.8415 | 0.5403 | 0.0998 | 0.9950 |
| 2 | 0.9093 | -0.4161 | 0.1987 | 0.9801 |
| 3 | 0.1411 | -0.9900 | 0.2955 | 0.9553 |
Estimation
Predict first
BPE (Byte-Pair Encoding, Sennrich 2016) builds a vocabulary by iteratively merging the most frequent adjacent symbol pair in the corpus. Start: every character is its own token.
Commit before you compute: what does BPE tokenization — 2 merge steps by hand come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: Step 1: merge ('l', 'o') → 'lo' (count=2). Step 2: merge ('lo', 'w') → 'low' (count=2)
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. All tied pairs (lo, ow, we, es, st, t</w>) count=2; Python's max picks the first alphabetically among ties.
Worked example
BPE (Byte-Pair Encoding, Sennrich 2016) builds a vocabulary by iteratively merging the most frequent adjacent symbol pair in the corpus. Start: every character is its own token.
corpus = ['l o w </w>', 'l o w e r </w>',
'n e w e s t </w>', 'w i d e s t </w>']
def get_pairs(vocab):
pairs = {}
for word in vocab:
syms = word.split()
for i in range(len(syms) - 1):
p = (syms[i], syms[i+1])
pairs[p] = pairs.get(p, 0) + 1
return pairs
pairs1 = get_pairs(corpus)
top1 = max(pairs1, key=pairs1.get)
print('Step 1 top pair:', top1, ' count:', pairs1[top1])
# perform merge
corpus2 = [w.replace(top1[0]+' '+top1[1], top1[0]+top1[1])
for w in corpus]
pairs2 = get_pairs(corpus2)
top2 = max(pairs2, key=pairs2.get)
print('Step 2 top pair:', top2, ' count:', pairs2[top2])
print('Corpus after 2 merges:', corpus2)Step 1: merge ('l', 'o') → 'lo' (count=2). Step 2: merge ('lo', 'w') → 'low' (count=2)
Why: All tied pairs (lo, ow, we, es, st, t</w>) count=2; Python's max picks the first alphabetically among ties. After merge 1, 'lo w' becomes the new top pair.
| merge # | pair merged | count | new token | example word after |
|---|---|---|---|---|
| 1 | ('l', 'o') | 2 | 'lo' | 'lo w </w>' |
| 2 | ('lo', 'w') | 2 | 'low' | 'low </w>' |
| ... | continue until vocab_size reached | ... | ... | ... |
Reverse engineer
Discussion prompt
Work backwards. The example finished here:
Step 1: merge ('l', 'o') → 'lo' (count=2). Step 2: merge ('lo', 'w') → 'low' (count=2)
What was it asked to do, and what must it have been given? Reconstruct the problem from its answer.
Hint: Every quantity in the result had to enter somewhere. Account for each one.
Answer:
BPE (Byte-Pair Encoding, Sennrich 2016) builds a vocabulary by iteratively merging the most frequent adjacent symbol pair in the corpus. Start: every character is its own token.
Anomaly
Predict first
A student writes this, and it looks reasonable:
BPE finds the most frequent bigram within each word separately, merges it, then repeats.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: BPE counts pairs across the entire corpus and merges the globally most frequent pair in one step.
BPE counts each adjacent symbol pair across all words in the corpus, finds the globally most frequent pair, and merges it everywhere in one step.
Why: BPE counts pairs across the entire corpus and merges the globally most frequent pair in one step. The merge affects every occurrence of that pair in every word simultaneously.
Trap
BPE finds the most frequent bigram within each word separately, merges it, then repeats.
Count pairs per word independently; merge the top pair in each word
Why: Wrong. BPE counts pairs across the entire corpus and merges the globally most frequent pair in one step. The merge affects every occurrence of that pair in every word simultaneously.
BPE counts each adjacent symbol pair across all words in the corpus, finds the globally most frequent pair, and merges it everywhere in one step.
pairs = get_pairs(full_corpus); top = max(pairs, key=pairs.get); merge top everywhere
Why: In the toy corpus, ('l','o') appears in 'low' and 'lower' — count=2 corpus-wide. It ties with 5 other pairs but is selected as the global merge. One step, one pair, applied to all words.
Break the constraint
Discussion prompt
The rule this trap just fixed:
In the toy corpus, ('l','o') appears in 'low' and 'lower' — count=2 corpus-wide. It ties with 5 other pairs but is selected as the global merge. One step, one pair, applied to all words.
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:
BPE counts pairs across the entire corpus and merges the globally most frequent pair in one step. The merge affects every occurrence of that pair in every word simultaneously.
Section
Part 3 of 4
Concept
IoU measures how well two bounding boxes overlap. It is the backbone of object-detection metrics (mAP) and NMS (Lesson 60).
\[ \text{IoU}(A, B) = \frac{|A \cap B|}{|A \cup B|} = \frac{\text{area of intersection}}{\text{area of A} + \text{area of B} - \text{area of intersection}} \]
| box1 [x1,y1,x2,y2] | box2 | A1 | A2 | inter | union | IoU |
|---|---|---|---|---|---|---|
| [0,0,4,4] | [2,2,6,6] | 16 | 16 | 4 | 28 | 0.1429 |
| [0,0,4,4] | [0,0,4,4] | 16 | 16 | 16 | 16 | 1.0000 |
| [0,0,4,4] | [5,5,9,9] | 16 | 16 | 0 | 32 | 0.0000 |
| [0,0,6,6] | [2,2,4,4] | 36 | 4 | 4 | 36 | 0.1111 |
Pattern
Step through it
Step through Intersection over Union (IoU) one row at a time. What is driving the change, and what would the row after the last one be?
Estimation
Predict first
Implement IoU as a function, then state the NMS decision rule. Predict iou([0,0,4,4], [2,2,6,6]) before running.
Commit before you compute: what does IoU implementation and NMS connection come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: inter area = (min(4,6)-max(0,2)) * (min(4,6)-max(0,2)) = 2*2 = 4; union = 16+16-4 = 28; IoU = 4/28 = 0.1429
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 boxes share a 2×2 corner region.
Worked example
Implement IoU as a function, then state the NMS decision rule. Predict iou([0,0,4,4], [2,2,6,6]) before running.
def iou(b1, b2):
ix1, iy1 = max(b1[0],b2[0]), max(b1[1],b2[1])
ix2, iy2 = min(b1[2],b2[2]), min(b1[3],b2[3])
inter = max(0, ix2-ix1) * max(0, iy2-iy1)
a1 = (b1[2]-b1[0]) * (b1[3]-b1[1])
a2 = (b2[2]-b2[0]) * (b2[3]-b2[1])
union = a1 + a2 - inter
return inter / union if union > 0 else 0.0
print(iou([0,0,4,4], [2,2,6,6])) # 0.14285714...
print(iou([0,0,4,4], [0,0,4,4])) # 1.0
print(iou([0,0,4,4], [5,5,9,9])) # 0.0
# NMS: suppress box B if IoU(best_box, B) > threshold
iou_threshold = 0.5
boxes = [[0,0,4,4],[1,1,5,5],[8,8,12,12]]
scores = [0.9, 0.8, 0.85]
best = max(range(3), key=lambda i: scores[i])
kept = [best]
for i in range(3):
if i != best and iou(boxes[best], boxes[i]) < iou_threshold:
kept.append(i)
print('NMS kept box indices:', sorted(kept)) # [0, 2]inter area = (min(4,6)-max(0,2)) * (min(4,6)-max(0,2)) = 2*2 = 4; union = 16+16-4 = 28; IoU = 4/28 = 0.1429
Why: The boxes share a 2×2 corner region. Since 0.1429 < 0.5 threshold, NMS would NOT suppress the second box — they are different detections.
| box pair | IoU | NMS decision (threshold=0.5) |
|---|---|---|
| [0,0,4,4] vs [2,2,6,6] | 0.1429 | KEEP both — low overlap |
| [0,0,4,4] vs [1,1,5,5] | 0.3600 (est.) | KEEP — just below 0.5 |
| identical boxes | 1.0000 | SUPPRESS duplicate |
| non-overlapping | 0.0000 | KEEP — no overlap |
Reverse engineer
Discussion prompt
Work backwards. The example finished here:
inter area = (min(4,6)-max(0,2)) * (min(4,6)-max(0,2)) = 2*2 = 4; union = 16+16-4 = 28; IoU = 4/28 = 0.1429
What was it asked to do, and what must it have been given? Reconstruct the problem from its answer.
Hint: Every quantity in the result had to enter somewhere. Account for each one.
Answer:
Implement IoU as a function, then state the NMS decision rule. Predict iou([0,0,4,4], [2,2,6,6]) before running.
Concept
U-Net (Ronneberger 2015) is the standard architecture for pixel-level segmentation. Its shape: a contracting encoder path that halves spatial size, a symmetric expanding decoder that restores it, with skip connections that concatenate encoder features directly into decoder layers.
Conv→Conv→MaxPool repeated; captures progressively abstract featuresUpsample→cat(encoder_feat)→Conv→Conv repeated — skip connections reintroduce spatial detail lost by pooling| stage | channels (MiniUNet) | spatial (input 32×32) | params |
|---|---|---|---|
| enc1 (2×Conv) | 1→8 | 32×32 | 664 |
| enc2 (2×Conv) | 8→16 | 16×16 | 3,488 |
| bottleneck | 16→32 | 8×8 | 13,888 |
| dec2 (cat+2×Conv) | 48→16 | 16×16 | 9,248 |
| dec1 (cat+2×Conv) | 24→8 | 32×32 | 2,320 |
| head (1×1) | 8→1 | 32×32 | 9 |
| TOTAL | 29,617 |
Sorting
Sort into buckets
These are the pieces of Lesson 111: Transformers, NLP & CV — Phase 3 Review, out of order. Put each one back under the part of the lesson it belongs to.
Concept
A core exam topic: when do transformers beat CNNs on vision tasks, and why? The answer lies in inductive bias and data scale.
| property | CNN (ResNet, Lesson 55) | ViT (Lesson 92) |
|---|---|---|
| locality bias | YES — conv kernel is local | NO — all patches attend to all |
| translation equiv. | YES — weight sharing | NO — learned pos embeddings |
| O(input) attention | O(T·k²) per conv layer | O(T²·d_k) per encoder layer |
| data regime <10k | competitive | weaker — needs to learn spatial structure |
| data regime >100M | strong but plateaus | matches or beats CNN |
| fine-tuning | ImageNet → domain transfer | same but interpolate pos embeddings at new resolution |
Practical rule: if your dataset has < ~10k labeled examples, prefer a CNN (or a pretrained ViT fine-tuned with very low LR). At 100M+ samples or with massive pretraining (JFT-300M), ViT is superior. U-Net occupies a different niche — segmentation, not classification.
Comparison
Comparison matrix
From ViT vs CNN — the Phase 3 synthesis: refill the CNN (ResNet, Lesson 55) column from what you know. The rest of the table is as it appeared.
| property | CNN (ResNet, Lesson 55) | ViT (Lesson 92) |
|---|---|---|
| locality bias | YES — conv kernel is local | NO — all patches attend to all |
| translation equiv. | YES — weight sharing | NO — learned pos embeddings |
| O(input) attention | O(T·k²) per conv layer | O(T²·d_k) per encoder layer |
| data regime <10k | competitive | weaker — needs to learn spatial structure |
| data regime >100M | strong but plateaus | matches or beats CNN |
| fine-tuning | ImageNet → domain transfer | same but interpolate pos embeddings at new resolution |
Anomaly
Predict first
A student writes this, and it looks reasonable:
The skip connection adds the encoder feature map to the decoder feature map: decoder_feat = upsample(x) + encoder_feat.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: Wrong for U-Net. Adding requires the two tensors to have identical shapes (same channel count).
U-Net skip connections concatenate along the channel dimension: torch.cat([upsample(x), encoder_feat], dim=1). Channel count doubles — handled by the next conv block.
Why: Wrong for U-Net. Adding requires the two tensors to have identical shapes (same channel count). U-Net's encoder and decoder usually have different channel counts at the same spatial level, so element-wise addition would require a projection or fail with a shape mismatch.
Trap
The skip connection adds the encoder feature map to the decoder feature map: decoder_feat = upsample(x) + encoder_feat.
x = self.up(x) + e1 # add residual-style
Why: Wrong for U-Net. Adding requires the two tensors to have identical shapes (same channel count). U-Net's encoder and decoder usually have different channel counts at the same spatial level, so element-wise addition would require a projection or fail with a shape mismatch.
U-Net skip connections concatenate along the channel dimension: torch.cat([upsample(x), encoder_feat], dim=1). Channel count doubles — handled by the next conv block.
x = torch.cat([self.up(x), e1], dim=1) # channels: 32+16=48
Why: In MiniUNet: bottleneck outputs 32 channels; after upsample + cat with enc2 (16 channels), the dec2 input has 48 channels. The subsequent Conv block maps 48→16. Verified: MiniUNet with this structure has 29,617 params and output shape (2,1,32,32).
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.
(T, d_k)):; The QK^T matrix multiply costs O(T²·d_k) — this is the fundamental bottleneck of the transformer. Doubling sequence length quadruples the attention cost.; IoU measures how well two bounding boxes overlap. It is the backbone of object-detection metrics (mAP) and NMS (Lesson 60).scores = Q @ K.T. Seems fine — softmax still produces a valid distribution.; BPE finds the most frequent bigram within each word separately, merges it, then repeats.Section
Summary
Constraint
Discussion prompt
Run Phase 3 implementation and debugging checklist with this step confiscated:
BPE: count all adjacent-symbol pairs across the full corpus; merge the globally most frequent pair; repeat until target vocab size. Merge affects every occurrence everywhere.
Is it still possible? If it is, say what takes its place and what it costs you. If it is not, say exactly what that step was providing that nothing else does.
Hint: A step you can drop for free was never load-bearing. If you cannot drop it, name the thing that goes wrong the moment it is gone.
Answer:
softmax(Q@K.T / sqrt(d_k)) @ V. d_k = d_model / n_heads. Never omit the scale — without it, scores spread by sqrt(d_k) and softmax…PE[p, 2i] = sin(p / 10000^(2i/d)), PE[p, 2i+1] = cos(...). Position 0 → all zeros for sin, all ones for cos. No…inter / (A1 + A2 - inter). Intersection corners: (max(x1s), max(y1s)) to (min(x2s), min(y2s)). NMS suppresses boxes with IoU > threshold vs the…torch.cat([up(x), enc_feat], dim=1) — concatenation, NOT addition. Channel count sums; the following conv block absorbs it.Pattern
softmax(Q@K.T / sqrt(d_k)) @ V. d_k = d_model / n_heads. Never omit the scale — without it, scores spread by sqrt(d_k) and softmax entropy collapses, killing gradients.PE[p, 2i] = sin(p / 10000^(2i/d)), PE[p, 2i+1] = cos(...). Position 0 → all zeros for sin, all ones for cos. No learned params.inter / (A1 + A2 - inter). Intersection corners: (max(x1s), max(y1s)) to (min(x2s), min(y2s)). NMS suppresses boxes with IoU > threshold vs the kept box.torch.cat([up(x), enc_feat], dim=1) — concatenation, NOT addition. Channel count sums; the following conv block absorbs it.Edge cases
Discussion prompt
Phase 3 implementation and debugging checklist 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:
softmax(Q@K.T / sqrt(d_k)) @ V. d_k = d_model / n_heads. Never omit the scale — without it, scores spread by sqrt(d_k) and softmax…PE[p, 2i] = sin(p / 10000^(2i/d)), PE[p, 2i+1] = cos(...). Position 0 → all zeros for sin, all ones for cos. No…inter / (A1 + A2 - inter). Intersection corners: (max(x1s), max(y1s)) to (min(x2s), min(y2s)). NMS suppresses boxes with IoU > threshold vs the…torch.cat([up(x), enc_feat], dim=1) — concatenation, NOT addition. Channel count sums; the following conv block absorbs it.Elimination
Eliminate the wrong options
A transformer encoder with d_model=256, n_heads=8 processes sequences of length T=512. The dominant cost per head per layer from the QK^T matrix multiply is proportional to which expression?
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: QK^T has shape (T, T). Each of the T×T entries requires d_k multiplications and additions, giving O(T²·d_k) per head. With T=512 and d_k=d_model/n_heads=256/8=32: 512²×32 = 8,388,608. The full model cost multiplies by n_heads (and by the number of layers), but per-head per-layer the count is T²·d_k.
Check
Work it out before clicking.
Check your understanding
A transformer encoder with d_model=256, n_heads=8 processes sequences of length T=512. The dominant cost per head per layer from the QK^T matrix multiply is proportional to which expression?
Answer: A
Why: QK^T has shape (T, T). Each of the T×T entries requires d_k multiplications and additions, giving O(T²·d_k) per head. With T=512 and d_k=d_model/n_heads=256/8=32: 512²×32 = 8,388,608. The full model cost multiplies by n_heads (and by the number of layers), but per-head per-layer the count is T²·d_k.
Prediction
Predict first
Using the sinusoidal positional encoding formula with d_model=8, what is PE[pos=0, dim=0]?
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: 0.0
Why: PE[pos, 2i] = sin(pos / 10000^(2i/d)). At pos=0 and i=0 (dim=0): PE[0, 0] = sin(0 / 10000^0) = sin(0) = 0.0. Verified: the PE table shows pos=0 row is all [0.0, 1.0, 0.0, 1.0, ...] — sin=0, cos=1 at every frequency.
Check
Apply the formula directly — no code needed.
Check your understanding
Using the sinusoidal positional encoding formula with d_model=8, what is PE[pos=0, dim=0]?
Answer: A
Why: PE[pos, 2i] = sin(pos / 10000^(2i/d)). At pos=0 and i=0 (dim=0): PE[0, 0] = sin(0 / 10000^0) = sin(0) = 0.0. Verified: the PE table shows pos=0 row is all [0.0, 1.0, 0.0, 1.0, ...] — sin=0, cos=1 at every frequency.
Elimination
Eliminate the wrong options
Box A = [0, 0, 6, 6] (area 36). Box B = [2, 2, 4, 4] (area 4). What is IoU(A, B)?
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: Intersection corners: (max(0,2), max(0,2)) = (2,2) to (min(6,4), min(6,4)) = (4,4). Intersection area = (4-2)×(4-2) = 4. Union = 36 + 4 - 4 = 36. IoU = 4/36 = 1/9 ≈ 0.1111. Box B is fully inside Box A (it is a containment case), so the intersection equals Box B's area (4) and the union equals Box A's area (36). Verified in the trace table.
Check
Compute the intersection corners first, then the areas.
Check your understanding
Box A = [0, 0, 6, 6] (area 36). Box B = [2, 2, 4, 4] (area 4). What is IoU(A, B)?
Answer: A
Why: Intersection corners: (max(0,2), max(0,2)) = (2,2) to (min(6,4), min(6,4)) = (4,4). Intersection area = (4-2)×(4-2) = 4. Union = 36 + 4 - 4 = 36. IoU = 4/36 = 1/9 ≈ 0.1111. Box B is fully inside Box A (it is a containment case), so the intersection equals Box B's area (4) and the union equals Box A's area (36). Verified in the trace table.
Section
Project
Concept
Build multi-head attention from scratch using only nn.Linear and torch.softmax — no nn.MultiheadAttention allowed. Three milestones: single-head → multi-head → debugger.
| # | milestone | key tool / check |
|---|---|---|
| 1 | Single-head attention: implement Q/K/V projections, compute scaled dot-product, return (T, d_k) output | torch.softmax, row sums = 1.0 |
| 2 | Multi-head: split d_model into n_heads, run heads in parallel (or loop), concatenate + project | output shape (B, T, d_model) |
| 3 | Debugger: write a function that checks for the top-3 MHA bugs — shape errors, missing scale, NaN weights | assert statements, entropy check |
Build rules: d_model=16, n_heads=2, d_k=8, T=5, B=2. Print shapes after every step. Verify row sums = 1. Compare output to nn.MultiheadAttention on the same input.
Counterexample
Discussion prompt
Build multi-head attention from scratch using only nn.Linear and torch.softmax — no nn.MultiheadAttention allowed. Three milestones: single-head → multi-head → debugger.
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: d_model=16, n_heads=2, d_k=8, T=5, B=2. Print shapes after every step. Verify row sums = 1. Compare output to nn.MultiheadAttention on the same input.
Worked example
Your turn: implement single_head_attn(Q, K, V, d_k). Predict the shape of scores = Q @ K.transpose(-2,-1) when Q is (B, T, d_k). What must you divide by?
Hint: Q @ K.transpose(-2,-1) gives (B, T, T) — the attention matrix. Divide by d_k**0.5 before softmax.
import torch, torch.nn as nn
def single_head_attn(Q, K, V, d_k):
scores = Q @ K.transpose(-2, -1) / (d_k ** 0.5) # (B,T,T)
weights = torch.softmax(scores, dim=-1) # rows sum to 1
return weights @ V # (B,T,d_k)
torch.manual_seed(42)
B, T, dk = 2, 5, 8
Q = torch.randn(B, T, dk)
K = torch.randn(B, T, dk)
V = torch.randn(B, T, dk)
out = single_head_attn(Q, K, V, dk)
print('output shape:', out.shape) # (2, 5, 8)
weights = torch.softmax(Q@K.transpose(-2,-1)/(dk**0.5), dim=-1)
print('row sums:', weights[0].sum(dim=-1).tolist())| tensor | shape | formula |
|---|---|---|
| Q, K, V | (2, 5, 8) | B=2, T=5, d_k=8 |
| Q @ K.T | (2, 5, 5) | each query × all keys |
| / sqrt(8) | (2, 5, 5) | scale = 2.8284 |
| softmax | (2, 5, 5) | rows sum to 1.0 |
| weights @ V | (2, 5, 8) | weighted value sum |
Discrimination
Sort into buckets
Sort these by shape, from memory, without looking back at Milestone 1 — single-head attention. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.
Worked example
Your turn: extend to n_heads=2, d_model=16. Project x into Q, K, V with shape (B, T, d_model), then split into heads. Predict the shape after reshape(-1, T, d_k) when B=2, n_heads=2.
Hint: (B, T, n_heads, d_k).transpose(1,2) → (B, n_heads, T, d_k). Reshape to (B*n_heads, T, d_k) to reuse single-head code, then reverse.
class MultiHeadAttn(nn.Module):
def __init__(self, d_model, n_heads):
super().__init__()
self.h, self.dk = n_heads, d_model // n_heads
self.Wq = nn.Linear(d_model, d_model, bias=False)
self.Wk = nn.Linear(d_model, d_model, bias=False)
self.Wv = nn.Linear(d_model, d_model, bias=False)
self.Wo = nn.Linear(d_model, d_model, bias=False)
def forward(self, x):
B, T, d = x.shape
def split(W):
return W(x).view(B, T, self.h, self.dk).transpose(1, 2) # (B,h,T,dk)
Q, K, V = split(self.Wq), split(self.Wk), split(self.Wv)
sc = Q @ K.transpose(-2,-1) / (self.dk**0.5) # (B,h,T,T)
w = torch.softmax(sc, dim=-1)
out = (w @ V).transpose(1,2).reshape(B, T, d) # (B,T,d)
return self.Wo(out)
torch.manual_seed(42)
mha = MultiHeadAttn(16, 2)
x = torch.randn(2, 5, 16)
out = mha(x)
print('MHA output shape:', out.shape) # (2, 5, 16)| step | shape | note |
|---|---|---|
| x input | (2, 5, 16) | B=2, T=5, d_model=16 |
| Wq(x).view split | (2, 5, 2, 8) | split d into 2 heads × d_k=8 |
| transpose(1,2) | (2, 2, 5, 8) | (B, n_heads, T, d_k) |
| Q@K.T (attn scores) | (2, 2, 5, 5) | per-head attention matrix |
| w@V output per head | (2, 2, 5, 8) | each head's context |
| transpose+reshape | (2, 5, 16) | concat heads back |
| Wo output | (2, 5, 16) | final projection |
Trade off
Comparison matrix
From Milestone 2 — multi-head attention: every row here is a choice with a cost. Fill the shape column, then say which row you would actually pick and what you give up for it.
| step | shape | note |
|---|---|---|
| x input | (2, 5, 16) | B=2, T=5, d_model=16 |
| Wq(x).view split | (2, 5, 2, 8) | split d into 2 heads × d_k=8 |
| transpose(1,2) | (2, 2, 5, 8) | (B, n_heads, T, d_k) |
| Q@K.T (attn scores) | (2, 2, 5, 5) | per-head attention matrix |
| w@V output per head | (2, 2, 5, 8) | each head's context |
| transpose+reshape | (2, 5, 16) | concat heads back |
| Wo output | (2, 5, 16) | final projection |
Concept
import torch, torch.nn as nn
class MultiHeadAttn(nn.Module):
def __init__(self, d_model, n_heads):
super().__init__()
assert d_model % n_heads == 0
self.h, self.dk = n_heads, d_model // n_heads
self.Wq = nn.Linear(d_model, d_model, bias=False)
self.Wk = nn.Linear(d_model, d_model, bias=False)
self.Wv = nn.Linear(d_model, d_model, bias=False)
self.Wo = nn.Linear(d_model, d_model, bias=False)
def forward(self, x):
B, T, d = x.shape
def split_heads(W):
return W(x).view(B, T, self.h, self.dk).transpose(1,2)
Q,K,V = split_heads(self.Wq), split_heads(self.Wk), split_heads(self.Wv)
sc = Q @ K.transpose(-2,-1) / (self.dk**0.5)
w = torch.softmax(sc, dim=-1)
assert not torch.isnan(w).any(), 'NaN in attention weights — check scale!'
ctx = (w @ V).transpose(1,2).reshape(B, T, d)
return self.Wo(ctx)
def debug_mha(x, d_model, n_heads):
mha = MultiHeadAttn(d_model, n_heads)
assert x.ndim == 3, 'input must be 3D (B, T, d_model)'
out = mha(x)
assert out.shape == x.shape, f'shape mismatch: {out.shape}'
print(f'OK: MHA(d={d_model}, h={n_heads}): {x.shape} -> {out.shape}')
return out
torch.manual_seed(42)
debug_mha(torch.randn(2, 5, 16), 16, 2) # OK
debug_mha(torch.randn(1, 10, 64), 64, 4) # OK| common MHA bug | symptom | fix |
|---|---|---|
| missing sqrt(d_k) scale | NaN loss or near-zero grads after a few steps | divide by d_k**0.5 before softmax |
| wrong transpose axis | shape error in Q@K.T, e.g. (B,h,T,d_k) @ (B,h,T,d_k) | use .transpose(-2,-1) not .T on batched tensors |
| forget output projection Wo | model attends but can't mix head information | add nn.Linear(d_model, d_model) after reshape |
| CLS at wrong index | classification uses x[:,-1] instead of x[:,0] | always prepend CLS; use x[:,0] for head input |
Comparison
Comparison matrix
From The full program + show it off: refill the symptom column from what you know. The rest of the table is as it appeared.
| common MHA bug | symptom | fix |
|---|---|---|
| missing sqrt(d_k) scale | NaN loss or near-zero grads after a few steps | divide by d_k**0.5 before softmax |
| wrong transpose axis | shape error in Q@K.T, e.g. (B,h,T,d_k) @ (B,h,T,d_k) | use .transpose(-2,-1) not .T on batched tensors |
| forget output projection Wo | model attends but can't mix head information | add nn.Linear(d_model, d_model) after reshape |
| CLS at wrong index | classification uses x[:,-1] instead of x[:,0] | always prepend CLS; use x[:,0] for head input |
Connect it up
Draw it
One page, no notation unless you need it: draw how these connect — Attention — derivation and complexity · NLP — positional encoding & BPE · Computer Vision — IoU, U-Net, ViT · Pattern — the Phase 3 master recipe · Your turn: Phase 3 capstone implementation. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.
Recap
softmax(QK^T / sqrt(d_k)) V from scratch and explain why the scale factor is non-negotiable| topic | the one exam fact |
|---|---|
| attention scale | without sqrt(d_k): softmax collapses, entropy → 0, gradients vanish |
| PE position 0 | all sin dims = 0, all cos dims = 1 (independent of d or freq) |
| BPE | global pair count across full corpus; one merge per step, everywhere |
| IoU | inter / (A1 + A2 - inter); containment: IoU = smaller_area / larger_area |
| U-Net skip | torch.cat (concatenate), NOT element-wise add; channels sum |
| ViT complexity | O(T²) in sequence length; ViT-B/16: T=197, 197²=38,809 per head per layer |
Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.