USAAIO Lesson 115, from Phase 3, a complete transformer review. It runs from scaled dot-product attention through multi-head attention, the encoder block - multi-head attention, feed-forward network, layer norm, and residual - and the decoder block, with its masked self-attention, cross-attention, and feed-forward network. It then covers all the major variants, BERT, GPT, T5, ViT, and GNNs, comparing their architectures, and analyzes the O(n²d) complexity, causal masking, sinusoidal positional encoding, and the role of the CLS token. It ends with a PixelBERT implemented from scratch and verified with torch 2.7.1+cpu on the sklearn digits dataset. The lesson runs to 27 slides.
Subject: Machine Learning · 49 slides · code lesson
Open the interactive version of this deck · Homework for this lesson
Title
USAAIO · Lesson 115 · Phase 3
Attention → encoder block → decoder block → BERT / GPT / T5 / ViT / GNN. Every shape, every number, exam-speed drill. Speed target: MHA in 15 min, full encoder in 25 min.
Objectives
1/√d_k is necessarynn.Linear projections and explain why splitting into h heads helpsWarm-up
Discussion prompt
Before we open Lesson 115: Transformer End-to-End Review: without looking back, what was the main idea of Phase 3 Mock Exam — Transformers, NLP, and Computer Vision, 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:
timed mock-exam simulation covering all Phase 3 topics — scaled dot-product attention derivation, multi-head attention parameter counting, causal masking, sinusoidal positional encoding, BERT architecture and fine-tuning, LoRA adapters, greedy vs beam-search decoding, BLEU-4 and ROUGE-N computation, NMS with IoU, and full transformer training CE/perplexity. Worked answer-key walkthrough for each section.
Section
Part 1 of 5
Concept
Every transformer sublayer computes the same operation: given Q (queries), K (keys), V (values) — each (seq, d_k) — produce a weighted sum of values where weights are query–key similarities.
\[ \text{Attn}(Q,K,V) = \text{softmax}\!\left(\frac{QK^\top}{\sqrt{d_k}}\right)V \]
The √d_k scale prevents dot products from growing large in high dimension, which would push softmax into flat or near-one-hot regimes and kill gradients. With d_k=8 the scale factor is 1/√8 = 0.3536 (verified).
| Tensor | Shape | Meaning |
|---|---|---|
| Q | (seq, d_k) | What position i is looking for |
| K | (seq, d_k) | What position j offers |
| V | (seq, d_k) | What position j actually contributes |
| QKᵀ/√d_k | (seq, seq) | Similarity score matrix |
| softmax(·) | (seq, seq) | Row-stochastic attention weights |
| Output | (seq, d_k) | Weighted value sum per position |
Comparison
Comparison matrix
From Scaled dot-product attention: refill the Shape column from what you know. The rest of the table is as it appeared.
| Tensor | Shape | Meaning |
|---|---|---|
| Q | (seq, d_k) | What position i is looking for |
| K | (seq, d_k) | What position j offers |
| V | (seq, d_k) | What position j actually contributes |
| QKᵀ/√d_k | (seq, seq) | Similarity score matrix |
| softmax(·) | (seq, seq) | Row-stochastic attention weights |
| Output | (seq, d_k) | Weighted value sum per position |
Pattern
Predict first
The table runs: 1 | QKᵀ (pre-scale) | (4,4) | raw dot products · 2 | ÷ √8 = 0.3536 | (4,4) | scores dampened · 3 | softmax(row) | (4,4) | [0.1942, 0.4102, 0.3242, 0.0714]
In Trace: attention weights with seq=4, d_k=8, given the rows so far: what is the next one — the row where Step is 4?
Correct: 4 | × V | (4,8) | weighted value sum
| Step | Expression | Shape | Values (row 0) |
|---|---|---|---|
| 1 | QKᵀ (pre-scale) | (4,4) | raw dot products |
| 2 | ÷ √8 = 0.3536 | (4,4) | scores dampened |
| 3 | softmax(row) | (4,4) | [0.1942, 0.4102, 0.3242, 0.0714] |
| 4 | × V | (4,8) | weighted value sum |
Why: The relationship between the columns, not the individual numbers, is what generates the next row. Small concrete tensors let us print every number; seed 42 is reproducible.
Worked example
Fix inputs: Q, K, V each shape (4, 8), seed 42
Why: Small concrete tensors let us print every number; seed 42 is reproducible.
import torch, torch.nn.functional as F
torch.manual_seed(42)
seq, d_k = 4, 8
Q = torch.randn(seq, d_k)
K = torch.randn(seq, d_k)
V = torch.randn(seq, d_k)
scores = Q @ K.T / d_k**0.5 # (4,4)
attn_w = F.softmax(scores, dim=-1)
out = attn_w @ V # (4,8)
print(attn_w[0].tolist())Output: attn_w[0] = [0.1942, 0.4102, 0.3242, 0.0714]
Why: Row 0 sums to 1.0 (softmax row-stochastic). Position 0 attends most to position 1 (weight 0.4102).
| Step | Expression | Shape | Values (row 0) |
|---|---|---|---|
| 1 | QKᵀ (pre-scale) | (4,4) | raw dot products |
| 2 | ÷ √8 = 0.3536 | (4,4) | scores dampened |
| 3 | softmax(row) | (4,4) | [0.1942, 0.4102, 0.3242, 0.0714] |
| 4 | × V | (4,8) | weighted value sum |
Trade off
Comparison matrix
From Trace: attention weights with seq=4, d_k=8: every row here is a choice with a cost. Fill the Values (row 0) column, then say which row you would actually pick and what you give up for it.
| Step | Expression | Shape | Values (row 0) |
|---|---|---|---|
| 1 | QKᵀ (pre-scale) | (4,4) | raw dot products |
| 2 | ÷ √8 = 0.3536 | (4,4) | scores dampened |
| 3 | softmax(row) | (4,4) | [0.1942, 0.4102, 0.3242, 0.0714] |
| 4 | × V | (4,8) | weighted value sum |
Section
Part 2 of 5
Concept
A single attention head computes one similarity function. Multi-head attention projects into h independent subspaces (each d_k = d_model/h), runs h attention heads in parallel, then concatenates and projects back — capturing h different relationship types simultaneously.
\[ \text{MHA}(X) = \text{Concat}(\text{head}_1,\ldots,\text{head}_h)W^O \quad \text{head}_i = \text{Attn}(XW^Q_i, XW^K_i, XW^V_i) \]
| Hyperparameter | Value | Shape consequence |
|---|---|---|
| d_model | 32 | input/output dim |
| h (heads) | 4 | 4 parallel attention functions |
| d_k = d_model/h | 8 | per-head key/query/value dim |
| W_Q, W_K, W_V | Linear(32,32) | 4096 params each |
| W_O | Linear(32,32) | 1024 params |
| Total MHA params | 4 × 1024 = 4096 | verified: 4096 |
Estimation
Predict first
Lines 9-11: the reshape+transpose trick — 32 dims become 4×8, head axis inserted at dim 1.
Commit before you compute: what does MHA implementation: shapes at every 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: Output shape: (1, 6, 32) — same as input. Attention weight tensor is (1, 4, 6, 6).
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 is shape-preserving: (B, S, D) in, (B, S, D) out.
Worked example
Project X → Q, K, V via three Linear(32,32), then reshape to (B, h, S, d_k)
Why: The reshape splits the 32-dim into 4 heads × 8 dims each. .transpose(1,2) moves heads before the sequence dim so batched matmul works per head.
import torch, torch.nn as nn, torch.nn.functional as F
torch.manual_seed(42)
B, S, D, h = 1, 6, 32, 4
d_k = D // h # 8
W_q = nn.Linear(D, D, bias=False)
W_k = nn.Linear(D, D, bias=False)
W_v = nn.Linear(D, D, bias=False)
W_o = nn.Linear(D, D, bias=False)
x = torch.randn(B, S, D)
Q = W_q(x).view(B, S, h, d_k).transpose(1,2) # (1,4,6,8)
K = W_k(x).view(B, S, h, d_k).transpose(1,2)
V = W_v(x).view(B, S, h, d_k).transpose(1,2)
scores = Q @ K.transpose(-2,-1) / d_k**0.5 # (1,4,6,6)
attn = F.softmax(scores, dim=-1)
ctx = (attn @ V).transpose(1,2).contiguous().view(B, S, D) # (1,6,32)
out = W_o(ctx)
print(out.shape) # torch.Size([1, 6, 32])Lines 9-11: the reshape+transpose trick — 32 dims become 4×8, head axis inserted at dim 1.
Output shape: (1, 6, 32) — same as input. Attention weight tensor is (1, 4, 6, 6).
Why: MHA is shape-preserving: (B, S, D) in, (B, S, D) out. The h×S×S attention maps can be inspected per-head to see what each head specialized in.
| Tensor | Shape after op | Note |
|---|---|---|
| x | (1, 6, 32) | input |
| W_q(x).view(…) | (1, 6, 4, 8) | split heads |
| Q after .transpose(1,2) | (1, 4, 6, 8) | head dim first |
| scores Q@Kᵀ/√d_k | (1, 4, 6, 6) | per-head similarity |
| attn (softmax) | (1, 4, 6, 6) | per-head weights |
| ctx after merge | (1, 6, 32) | concatenated heads |
| out = W_o(ctx) | (1, 6, 32) | final projection |
Error analysis
Annotate
Walk the callouts on MHA implementation: shapes at every step. Each one is a place this is easy to get subtly wrong.
.transpose(1,2) moves heads before the sequence dim so batched matmul works per head.Section
Part 3 of 5
Concept
One encoder block applies four operations in order: LayerNorm → MHA → residual, then LayerNorm → FFN → residual. This is the Pre-LN variant (used in GPT-2, modern practice); original 'Attention is All You Need' used Post-LN.
\[ x \leftarrow x + \text{MHA}(\text{LN}(x)) \quad\text{then}\quad x \leftarrow x + \text{FFN}(\text{LN}(x)) \]
| Sublayer | Operation | Shape |
|---|---|---|
| LayerNorm 1 | normalize across d_model | (B, S, 32) |
| MHA | multi-head self-attention | (B, S, 32) |
| Residual add | x = x + MHA(LN(x)) | (B, S, 32) |
| LayerNorm 2 | normalize across d_model | (B, S, 32) |
| FFN | Linear(32,64) → ReLU → Linear(64,32) | (B, S, 32) |
| Residual add | x = x + FFN(LN(x)) | (B, S, 32) |
With d_model=32, h=4, d_ff=64: encoder block has 8,416 parameters (verified). FFN typically uses d_ff = 4×d_model, expanding then contracting.
Counterexample
Discussion prompt
With d_model=32, h=4, d_ff=64: encoder block has 8,416 parameters (verified). FFN typically uses d_ff = 4×d_model, expanding then contracting.
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.
Concept
A decoder block has three sublayers: (1) masked self-attention on the target, (2) cross-attention where Q comes from the decoder but K,V come from the encoder output, (3) FFN. The cross-attention is what lets the decoder 'read' the encoded source.
| Sublayer | Q source | K,V source | Mask |
|---|---|---|---|
| Masked self-attn | decoder tgt | decoder tgt | causal lower-triangular |
| Cross-attention | decoder tgt | encoder output | none (src padding mask only) |
| FFN | — | — | none |
With tgt_seq=4, src_seq=6, d_model=32: decoder maps (B,4,32)×(B,6,32) → (B,4,32). Decoder params with d=32, h=4, d_ff=64: 12,576 (verified — extra cross-attn projections).
Analogy
Discussion prompt
Explain Decoder block: two attention sublayers 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:
With tgt_seq=4, src_seq=6, d_model=32: decoder maps (B,4,32)×(B,6,32) → (B,4,32). Decoder params with d=32, h=4, d_ff=64: 12,576 (verified — extra cross-attn projections).
Anomaly
Predict first
A student writes this, and it looks reasonable:
Wrong: Apply a causal (lower-triangular) mask to the encoder self-attention so early tokens cannot see future ones.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: Feels like 'the encoder processes left-to-right so it should not peek ahead.'
Right: The encoder uses bidirectional (no causal) attention — every position sees every other. Only the decoder's self-attention sublayer uses the causal mask.
Why: Feels like 'the encoder processes left-to-right so it should not peek ahead.'
Trap
Wrong: Apply a causal (lower-triangular) mask to the encoder self-attention so early tokens cannot see future ones.
mask = torch.tril(torch.ones(S, S)) # used in encoder
Why: Feels like 'the encoder processes left-to-right so it should not peek ahead.'
Right: The encoder uses bidirectional (no causal) attention — every position sees every other. Only the decoder's self-attention sublayer uses the causal mask.
Encoder: mask=None. Decoder self-attn: mask = torch.tril(torch.ones(T, T))
Why: The encoder's job is to build a full contextual representation of the source; causal masking would cripple it. The decoder must not peek at future target tokens because at inference time they don't exist yet.
Concept
Attention is permutation-equivariant — it does not know token order. Positional encoding adds a position-dependent signal to each token embedding before the first layer.
\[ PE_{(pos,2i)} = \sin\!\left(\frac{pos}{10000^{2i/d}}\right) \quad PE_{(pos,2i+1)} = \cos\!\left(\frac{pos}{10000^{2i/d}}\right) \]
| pos | d0 (sin) | d1 (cos) | d2 (sin) | d3 (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 |
Low-index dims (d0,d1) oscillate fast (period 2π); high-index dims oscillate slowly (period 2π×10000). This gives every position a unique fingerprint. BERT and ViT use learned positional embeddings instead — both approaches work.
Pattern
Step through it
Step through Sinusoidal positional encoding one row at a time. What is driving the change, and what would the row after the last one be?
Section
Part 4 of 5
Concept
| Model | Architecture | Masking | Pre-train task | Primary use |
|---|---|---|---|---|
| BERT | Encoder-only | Bidirectional | MLM + NSP | Classification, NER, QA |
| GPT | Decoder-only | Causal (LT) | Next-token LM | Text generation |
| T5 | Enc-Dec | Enc=bi, Dec=causal | Span denoising | Seq2seq (translation, summarization) |
| ViT | Encoder-only | Bidirectional | Supervised / MAE | Image classification |
| GNN | Graph message-passing | Neighbor aggregation | Supervised / contrastive | Graph-structured data |
Lesson 88 built the full encoder (Vaswani). Lesson 89 covered BERT's masked LM. Lesson 90 covered GPT's causal LM. Lesson 92 built ViT. This lesson synthesizes all four under one roof.
Comparison
Comparison matrix
From Architecture variants: one comparison table: refill the Primary use column from what you know. The rest of the table is as it appeared.
| Model | Architecture | Masking | Pre-train task | Primary use |
|---|---|---|---|---|
| BERT | Encoder-only | Bidirectional | MLM + NSP | Classification, NER, QA |
| GPT | Decoder-only | Causal (LT) | Next-token LM | Text generation |
| T5 | Enc-Dec | Enc=bi, Dec=causal | Span denoising | Seq2seq (translation, summarization) |
| ViT | Encoder-only | Bidirectional | Supervised / MAE | Image classification |
| GNN | Graph message-passing | Neighbor aggregation | Supervised / contrastive | Graph-structured data |
Concept
BERT and ViT prepend a [CLS] token to the input sequence before any attention layers. By the final encoder layer, the CLS position has attended to every other token — it becomes a pooled sentence/image representation used for classification.
| Position | Token | After final encoder layer |
|---|---|---|
| 0 | [CLS] (learned) | aggregate of full sequence → fed to classifier head |
| 1 | token_1 | contextual embedding of token_1 |
| 2 | token_2 | contextual embedding of token_2 |
| … | … | … |
The CLS embedding is not a mean pool — it is a dedicated slot trained end-to-end to aggregate. Alternative: global average pooling of all positions (works comparably, used in many ViT variants). Both avoid autoregressive decoding for classification.
Concept
ViT (Lesson 92) applies a standard encoder to images by treating P×P pixel patches as tokens. With H=8, P=2: N = (8/2)² = 16 patches. A nn.Conv2d(C, d, kernel=P, stride=P) implements patch embedding in one op.
\[ N = \left(\frac{H}{P}\right)^2 = \left(\frac{8}{2}\right)^2 = 16 \quad\text{sequence length} = N+1 = 17\text{ (CLS)} \]
| Step | Tensor shape | Operation |
|---|---|---|
| Input image | (1, 1, 8, 8) | C=1 grayscale |
| patch_embed (Conv2d) | (1, 32, 4, 4) | 16 patch vectors, d=32 |
| .flatten(2).T(1,2) | (1, 16, 32) | sequence of N=16 tokens |
| cat([CLS, patches]) | (1, 17, 32) | prepend CLS |
| + pos_embed | (1, 17, 32) | learned positional bias |
Concept
A GNN updates each node by aggregating neighbor features, then combining with its own state. Mean aggregation + self-loop on a 3-node graph (edges 0↔1, 1↔2) with 4-dim features:
| Node | Features before | Mean of neighbors | After (self + mean-nbr) |
|---|---|---|---|
| 0 | [1.0, 0.0, 2.0, -1.0] | [0.0, 1.0, 0.0, 1.0] | [1.0, 1.0, 2.0, 0.0] |
| 1 | [0.0, 1.0, 0.0, 1.0] | [1.5, 1.0, 1.5, -0.5] → [1.5, 1.0, 1.5, -0.5] | [1.5, 2.0, 1.5, 0.5] |
| 2 | [2.0, 2.0, 1.0, 0.0] | [0.0, 1.0, 0.0, 1.0] | [2.0, 3.0, 1.0, 1.0] |
Unlike transformers, GNNs aggregate over sparse adjacency (not all pairs) — complexity O(E·d) per layer, where E = number of edges. Transformer's O(n²d) is infeasible on large graphs.
Counterexample
Discussion prompt
A GNN updates each node by aggregating neighbor features, then combining with its own state. Mean aggregation + self-loop on a 3-node graph (edges 0↔1, 1↔2) with 4-dim features:
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:
Unlike transformers, GNNs aggregate over sparse adjacency (not all pairs) — complexity O(E·d) per layer, where E = number of edges. Transformer's O(n²d) is infeasible on large graphs.
Section
Part 5 of 5
Concept
The QKᵀ matmul is O(n²d) — quadratic in sequence length n, linear in head dimension d. This is the transformer's primary scaling bottleneck for long sequences.
| n (seq len) | O(n²) | ×d=64 total ops |
|---|---|---|
| 64 | 4,096 | 262,144 |
| 128 | 16,384 | 1,048,576 |
| 256 | 65,536 | 4,194,304 |
| 512 | 262,144 | 16,777,216 |
At n=512 the attention matrix alone costs 16M multiply-adds per head per layer. This motivates FlashAttention (fused kernel, avoids materializing the full S×S matrix), sparse attention, and linear attention approximations — all USAAIO-relevant topics.
Pattern
Step through it
Step through Attention complexity: O(n²d) one row at a time. What is driving the change, and what would the row after the last one be?
Constraint
Discussion prompt
Run The transformer authoring recipe (exam speed) with this step confiscated:
Causal mask: torch.tril(torch.ones(T,T)) — decoder self-attn only; encoder uses no mask
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:
scores = Q@K.T / d_k**0.5; attn = softmax(scores); out = attn@VLinear(D,D), .view(B,S,h,d_k).transpose(1,2), run attention, .transpose(1,2).view(B,S,D), project outx = x + MHA(LN(x)) then x = x + FFN(LN(x)) — two residuals, two LayerNormstorch.tril(torch.ones(T,T)) — decoder self-attn only; encoder uses no maskx[:,0,:] for classification headnn.Conv2d(C, d, kernel=P, stride=P) → .flatten(2).transpose(1,2)Pattern
scores = Q@K.T / d_k**0.5; attn = softmax(scores); out = attn@VLinear(D,D), .view(B,S,h,d_k).transpose(1,2), run attention, .transpose(1,2).view(B,S,D), project outx = x + MHA(LN(x)) then x = x + FFN(LN(x)) — two residuals, two LayerNormstorch.tril(torch.ones(T,T)) — decoder self-attn only; encoder uses no maskx[:,0,:] for classification headnn.Conv2d(C, d, kernel=P, stride=P) → .flatten(2).transpose(1,2)Edge cases
Discussion prompt
The transformer authoring recipe (exam speed) 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:
scores = Q@K.T / d_k**0.5; attn = softmax(scores); out = attn@VLinear(D,D), .view(B,S,h,d_k).transpose(1,2), run attention, .transpose(1,2).view(B,S,D), project outx = x + MHA(LN(x)) then x = x + FFN(LN(x)) — two residuals, two LayerNormstorch.tril(torch.ones(T,T)) — decoder self-attn only; encoder uses no maskx[:,0,:] for classification headnn.Conv2d(C, d, kernel=P, stride=P) → .flatten(2).transpose(1,2)Elimination
Eliminate the wrong options
With d_k = 64, what is the scale factor applied to QKᵀ before the softmax, and why is it necessary?
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: Dot product of two d_k-dim unit random vectors has variance d_k, so std dev is √d_k. Dividing by √d_k restores unit variance, keeping softmax inputs in the responsive region where gradients are non-vanishing. With d_k=64, scale = 1/8 = 0.125.
Check
Work through this before clicking. Pen and paper — do not compute softmax in your head, just identify the correct formula step.
Check your understanding
With d_k = 64, what is the scale factor applied to QKᵀ before the softmax, and why is it necessary?
Answer: A
Why: Dot product of two d_k-dim unit random vectors has variance d_k, so std dev is √d_k. Dividing by √d_k restores unit variance, keeping softmax inputs in the responsive region where gradients are non-vanishing. With d_k=64, scale = 1/8 = 0.125.
Check
Identify the model given the description. Think about masking and blocks before clicking.
Check your understanding
A model processes a source sequence through N encoder blocks (bidirectional attention), then generates a target token-by-token through N decoder blocks using masked self-attention on the target and cross-attention to the encoder output. Which model does this describe?
Answer: A
Why: T5 is an encoder-decoder (seq2seq) model: the encoder builds a bidirectional representation of the source; the decoder autoregressively generates the target with masked self-attention (causal) and cross-attention to the encoder output. This is the classic Vaswani 2017 architecture adapted for text-to-text pre-training.
Elimination
Eliminate the wrong options
A 12-layer transformer encoder processes sequences of length n=1024 with d_model=768 and h=12 heads (d_k=64 each). If you double the sequence length to n=2048, by what factor does the cost of the attention computation increase?
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: Attention cost is O(n²d). Doubling n multiplies n² by 4, so cost increases 4×. Concretely: at n=1024, QKᵀ is 1024² = 1,048,576 entries per head per layer; at n=2048 it is 2048² = 4,194,304 — exactly 4×. The d factor and h and L factors remain constant.
Check
Combine complexity and masking knowledge. No calculator — reason from the formulas.
Check your understanding
A 12-layer transformer encoder processes sequences of length n=1024 with d_model=768 and h=12 heads (d_k=64 each). If you double the sequence length to n=2048, by what factor does the cost of the attention computation increase?
Answer: A
Why: Attention cost is O(n²d). Doubling n multiplies n² by 4, so cost increases 4×. Concretely: at n=1024, QKᵀ is 1024² = 1,048,576 entries per head per layer; at n=2048 it is 2048² = 4,194,304 — exactly 4×. The d factor and h and L factors remain constant.
Pattern
Predict first
The table runs: 0 | 2.3599 | 0.1125 | random init, ~10-class chance · 1 | 2.3490 | 0.1125 | loss falling, weights adjusting · 2 | 2.3391 | 0.1125 | stable decay pattern · 3 | 2.3304 | 0.0625 | acc fluctuates — tiny dataset · 4 | 2.3226 | 0.0500 | loss strictly decreasing
In PixelBERT: CLS classification on digits (full…, given the rows so far: what is the next one — the row where Step is 5?
Correct: 5 | 2.3156 | 0.0875 | longer training needed for high acc
| Step | Loss | Acc | Note |
|---|---|---|---|
| 0 | 2.3599 | 0.1125 | random init, ~10-class chance |
| 1 | 2.3490 | 0.1125 | loss falling, weights adjusting |
| 2 | 2.3391 | 0.1125 | stable decay pattern |
| 3 | 2.3304 | 0.0625 | acc fluctuates — tiny dataset |
| 4 | 2.3226 | 0.0500 | loss strictly decreasing |
| 5 | 2.3156 | 0.0875 | longer training needed for high acc |
Why: The relationship between the columns, not the individual numbers, is what generates the next row. This exercises the full BERT-style pipeline in < 30 lines: patch projection, CLS prepend, learned positional embedding, encoder block, linear head.
Worked example
Treat each 8×8 digit image as 8 tokens of 8 pixels each, project to d=16, prepend CLS, run one encoder block, classify on CLS
Why: This exercises the full BERT-style pipeline in < 30 lines: patch projection, CLS prepend, learned positional embedding, encoder block, linear head.
import torch, torch.nn as nn, torch.nn.functional as F
from sklearn.datasets import load_digits
torch.manual_seed(0)
class EncoderBlock(nn.Module):
def __init__(self, d, h, d_ff):
super().__init__()
self.mha = nn.MultiheadAttention(d, h, batch_first=True)
self.ff = nn.Sequential(nn.Linear(d,d_ff), nn.ReLU(), nn.Linear(d_ff,d))
self.ln1 = nn.LayerNorm(d); self.ln2 = nn.LayerNorm(d)
def forward(self, x):
a,_ = self.mha(self.ln1(x), self.ln1(x), self.ln1(x))
x = x + a
return x + self.ff(self.ln2(x))
class PixelBERT(nn.Module):
def __init__(self):
super().__init__()
self.proj = nn.Linear(8, 16)
self.cls = nn.Parameter(torch.randn(1,1,16))
self.pos = nn.Parameter(torch.randn(1,9,16))
self.enc = EncoderBlock(16, 2, 32)
self.head = nn.Linear(16, 10)
def forward(self, x):
B = x.size(0)
t = self.proj(x.view(B,8,8))
t = torch.cat([self.cls.expand(B,-1,-1), t], 1) + self.pos
return self.head(self.enc(t)[:,0,:])
digits = load_digits()
X = torch.tensor(digits.data[:80]/16., dtype=torch.float32)
y = torch.tensor(digits.target[:80])
opt = torch.optim.Adam(PixelBERT().parameters(), lr=1e-3)
model = PixelBERT()
opt = torch.optim.Adam(model.parameters(), lr=1e-3)
for step in range(6):
loss = F.cross_entropy(model(X), y)
loss.backward(); opt.step(); opt.zero_grad()
acc = (model(X).argmax(1)==y).float().mean().item()
print(f"step {step}: loss={loss.item():.4f} acc={acc:.4f}")Training output (6 steps, seed 0, 80 samples from load_digits)
Why: Loss monotonically decreases, confirming the forward pass is correct and gradients flow through CLS → encoder → projection.
| Step | Loss | Acc | Note |
|---|---|---|---|
| 0 | 2.3599 | 0.1125 | random init, ~10-class chance |
| 1 | 2.3490 | 0.1125 | loss falling, weights adjusting |
| 2 | 2.3391 | 0.1125 | stable decay pattern |
| 3 | 2.3304 | 0.0625 | acc fluctuates — tiny dataset |
| 4 | 2.3226 | 0.0500 | loss strictly decreasing |
| 5 | 2.3156 | 0.0875 | longer training needed for high acc |
Trade off
Comparison matrix
From PixelBERT: CLS classification on digits (full…: every row here is a choice with a cost. Fill the Acc column, then say which row you would actually pick and what you give up for it.
| Step | Loss | Acc | Note |
|---|---|---|---|
| 0 | 2.3599 | 0.1125 | random init, ~10-class chance |
| 1 | 2.3490 | 0.1125 | loss falling, weights adjusting |
| 2 | 2.3391 | 0.1125 | stable decay pattern |
| 3 | 2.3304 | 0.0625 | acc fluctuates — tiny dataset |
| 4 | 2.3226 | 0.0500 | loss strictly decreasing |
| 5 | 2.3156 | 0.0875 | longer training needed for high acc |
Step zero
Discussion prompt
Your turn — multi-head attention in 15 minutes — before any calculation: what is the plan? Name the moves in order, in plain English, without doing the arithmetic.
Hint: It starts with: Milestone 1 (3 min): Define __init__ — four `nn.Linear(d_model…
Answer:
__init__ — four nn.Linear(d_model, d_model, bias=False) for W_q, W_k, W_v, W_o. Store h and d_k = d_model//h.forward(x), project and reshape — W_q(x).view(B,S,h,d_k).transpose(1,2) for Q,K,V.scores = Q@K.T(-2,-1) / d_k**0.5, softmax on dim=-1, ctx = attn@V, then .transpose(1,2).contiguous().view(B,S,D), return…mha(torch.randn(2,6,32)) should return shape (2,6,32). Count params: 4×32×32 = 4096.Worked example
Speed drill task: implement MultiHeadAttention from scratch using only nn.Linear, F.softmax, and tensor ops. No nn.MultiheadAttention. Target: working code in 15 minutes, matching the shapes in the trace table below.
Milestone 1 (3 min): Define __init__ — four nn.Linear(d_model, d_model, bias=False) for W_q, W_k, W_v, W_o. Store h and d_k = d_model//h.
Why: All four projections are d_model×d_model — no per-head parameters, the split happens in the reshape.
Milestone 2 (5 min): In forward(x), project and reshape — W_q(x).view(B,S,h,d_k).transpose(1,2) for Q,K,V.
Why: .transpose(1,2) moves h before S so batched matmul Q@K.transpose(-2,-1) operates per-head (dim -2 and -1 are S and d_k).
Milestone 3 (5 min): Compute scores = Q@K.T(-2,-1) / d_k**0.5, softmax on dim=-1, ctx = attn@V, then .transpose(1,2).contiguous().view(B,S,D), return W_o(ctx).
Why: .contiguous() is required before .view() because .transpose() makes the tensor non-contiguous in memory.
Milestone 4 (2 min): Verify — mha(torch.randn(2,6,32)) should return shape (2,6,32). Count params: 4×32×32 = 4096.
Why: Shape-preserving and param count match both confirm correctness before wiring into an encoder block.
| Milestone | Target | Check |
|---|---|---|
| 1 — init | 4 Linear(32,32,bias=False) | param count = 4096 |
| 2 — reshape Q | shape (B,h,S,d_k) = (2,4,6,8) | transpose(1,2) after view |
| 3 — scores | shape (B,h,S,S) = (2,4,6,6) | dim=-1 softmax |
| 4 — output | shape (B,S,D) = (2,6,32) | .contiguous().view needed |
Comparison
Comparison matrix
From Your turn — multi-head attention in 15 minutes: refill the Check column from what you know. The rest of the table is as it appeared.
| Milestone | Target | Check |
|---|---|---|
| 1 — init | 4 Linear(32,32,bias=False) | param count = 4096 |
| 2 — reshape Q | shape (B,h,S,d_k) = (2,4,6,8) | transpose(1,2) after view |
| 3 — scores | shape (B,h,S,S) = (2,4,6,6) | dim=-1 softmax |
| 4 — output | shape (B,S,D) = (2,6,32) | .contiguous().view needed |
Connect it up
Draw it
One page, no notation unless you need it: draw how these connect — Scaled Dot-Product Attention — the atomic unit · Multi-Head Attention — h projections in parallel · Encoder & Decoder Blocks — the full architecture · BERT / GPT / T5 / ViT / GNN — the variant landscape · Complexity & Implementation Drill. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.
Recap
Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.