Lesson 2 of 7 · 48 min
Attention from scratch
Scaled dot-product attention, derived. Why Q/K/V are three projections of the same input, why we divide by √d_k (the exact variance argument), what multi-head actually is (a shape split, not extra compute), the exact tensor shapes through one head, and the quadratic cost that defines everything downstream.
The one operation the whole field is built on
Attention(Q, K, V) = softmax(Q·Kᵀ / √d_k)·V. Read it as a soft, content-based dictionary lookup. For each token we form three vectors by learned linear projections of the same input: a query (“what am I looking for?”), a key (“what do I contain that others might want?”), and a value (“what information do I deliver if attended to?”). The dot product q·k scores how well a query matches each key; softmax over the keys turns scores into weights summing to 1; we then take that weighted sum of the value vectors. Self-attention means Q, K, V all come from the same sequence; cross-attention means queries come from one sequence, keys/values from another.
Let’s build GPT: from scratch, in code, spelled outAndrej KarpathyWhy divide by √d_k — the exact variance argument
q and k is independent with mean 0 and variance σ². Then the dot product q·k = Σ qᵢkᵢ (a sum of d_k independent zero-mean terms) has mean 0 and variance d_k·σ² — it grows linearly with the head dimension. Feed large-variance logits into softmax and it saturates: one logit dominates, the distribution becomes nearly one-hot, and the gradient through softmax becomes vanishingly small. Dividing by √d_k rescales the variance back to σ², keeping softmax in its responsive regime where gradients flow.1WHY 1/sqrt(d_k), NOT 1/d_k, NOT 123 q, k components: i.i.d., mean 0, variance s^24 q . k = sum of d_k terms -> mean 0, VARIANCE = d_k * s^2 (grows with d_k)56 divide by sqrt(d_k): Var(q . k / sqrt(d_k)) = s^2 (back to unit scale)78 too big a logit -> softmax saturates -> ~one-hot -> gradient ~ 0 -> no learning910 Vaswani: large d_k "pushes softmax into regions of extremely small gradients."11 So the divisor is exact variance correction, not a vibe.Multi-head: a shape split, not more compute
MultiHead(Q,K,V) = Concat(head₁,…,head_h)·Wᴼ, where headᵢ = Attention(Q·Wᵢ_Q, K·Wᵢ_K, V·Wᵢ_V). Each head projects into a smaller d_k = d_model/h subspace (Vaswani base: d_model=512, h=8, so d_k=d_v=64), runs scaled dot-product attention in parallel, then all heads are concatenated back to d_model and mixed by Wᴼ. The total parameter count and per-token FLOPs match a single d_model-wide attention exactly.h smaller heads lets each specialize on a different relation (one head tracks subject→verb, another verb→object, another local n-grams) and the averaging happens within a head, not across all of them. Training drives heads to differentiate because identical heads waste capacity. Interview angle. “Why not one big head with the same FLOPs?” → the resolution-collapse argument (Shazeer): one head loses effective resolution by averaging soft pointers; multiple heads preserve distinct relational subspaces — and this is exactly the property MQA/GQA later trade away for serving speed.Key idea
The exact tensor shapes through one block
.contiguous() after a transpose, or a double-softmax. Take GPT-2 124M: d_model=768, n_head=12, d_head=64, input x of shape (B, T, 768).1TENSOR SHAPES, ONE BLOCK (GPT-2 124M: d_model=768, n_head=12, d_head=64)23 x (B, T, 768)4 -> Q,K,V projections each (B, T, 768)5 -> reshape per head (B, 12, T, 64) split 768 = 12 * 646 -> Q @ K^T (B, 12, T, T) <-- the T x T attention matrix7 -> / sqrt(64), causal mask, softmax over LAST dim (keys)8 -> @ V (B, 12, T, 64)9 -> merge heads (B, T, 768)10 -> output proj W^O (B, T, 768)11 -> residual add (B, T, 768)1213 The (B, n_head, T, T) matrix is what dominates memory and what FlashAttention14 refuses to materialize in HBM (Lesson 6).1import torch, torch.nn.functional as F23def attention(q, k, v, causal=True):4 # q,k,v: (B, n_head, T, d_head)5 d_head = q.size(-1)6 att = (q @ k.transpose(-2, -1)) / d_head**0.5 # (B, n_head, T, T)7 if causal:8 T = q.size(-2)9 mask = torch.tril(torch.ones(T, T, device=q.device)).bool()10 att = att.masked_fill(~mask, float("-inf")) # block attending to the future11 att = att.softmax(dim=-1) # over KEYS (last dim)12 return att @ v # (B, n_head, T, d_head)1314# nanoGPT does exactly this: reshape (B,T,C)->(B,n_head,T,d_head), score, mask, softmax, @v, merge.mask[i,j] = -inf for j > i before softmax means position i can only attend to positions ≤ i — that’s what makes generation autoregressive and lets the model be trained on all positions in parallel without “seeing the future.” BERT instead uses bidirectional attention (no causal mask) and a different training objective (masked-language-modeling). Interview angle. “What single property distinguishes encoder-only, decoder-only, and encoder-decoder?” → whether decoding is autoregressive; the causal mask is the mechanical expression of it.The quadratic cost — two distinct phenomena
O(T²·d) per layer: forming Q·Kᵀ materializes a T×T matrix and (att·V) multiplies it back. Memory I/O is O(T²·d + T·d²): a naive kernel writes the T×T attention matrix to GPU HBM after softmax and reads it back for the next matmul, so even with free FLOPs the memory traffic is quadratic. As sequences grew, the binding constraint migrated from “too much math” to “too much memory movement” — which is exactly the door FlashAttention walks through (Lesson 6).1THE QUADRATIC, BY THE NUMBERS (per layer, T = sequence length)23 T x T attention matrix grows with T^2:4 T = 1k -> 1,000,000 entries / head5 T = 8k -> 64,000,000 entries / head (64x the 1k case)6 T = 32k -> 1,073,741,824 entries / head (~1B, per head, per layer)78 compute: O(T^2 * d) the matmuls9 HBM I/O: O(T^2 * d + T*d^2) writing/reading the T x T matrix (the real wall)1011 vs RNN: O(T * d^2) total, but SEQUENTIAL (no parallelism) -> the trade we made.Common mistake
“Attention is O(n²) and RNNs are O(n), so transformers are just slower.”
What attention heads actually learn
H query heads (so query-side resolution survives) but share K/V across groups, which is why GQA holds quality near MHA while MQA degrades more. Interview angle. “What does GQA give up relative to MHA?” → some key/value resolution — the distinct K/V subspaces per head — in exchange for an H/G× smaller KV cache; query-head diversity is preserved, which is why the quality hit is small at G=8.Key idea
Interview prep
softmax(Q·Kᵀ/√d_k)·V and the variance argument on a whiteboard.- 01“Explain self-attention.” → content-based mixing → query·key scores → softmax over keys → weighted sum of values; every token reaches every other in one hop; mask the future in decoders.
- 02“Why divide by √d_k?” → q·k has variance d_k·σ²; large logits saturate softmax and kill gradients; √d_k restores unit scale. (Not “normalization.”)
- 03“What is Q/K/V intuitively?” → Q = what I’m looking for, K = what I contain, V = what I deliver; three learned projections of the same input.
- 04“Why multi-head instead of one big head?” → resolution: one softmax averages all positions into mush; heads specialize on different relations at the same FLOPs.
- 05“What are the exact shapes?” → (B,T,d_model) → per-head (B,h,T,d_head) → scores (B,h,T,T) → softmax over keys → @V → merge to (B,T,d_model).
- 06“What’s the cost of attention?” → compute O(T²·d); HBM I/O O(T²·d + T·d²) — name both; the I/O term is the long-context wall.
- 07“Self vs cross attention?” → self: Q,K,V from one sequence; cross: queries from the decoder, keys/values from the encoder.
- 08“Dot-product vs additive attention?” → dot-product is a single matmul (GPU-friendly, needs the √d_k scale); additive (Bahdanau) uses a small MLP — slower, less common today.
Common mistake
The red-flag answer: calling √d_k “just a normalization trick.”
Checkpoint
A teammate removes the 1/√d_k scaling from a custom attention to “simplify,” and now training loss plateaus immediately with tiny gradients. Best diagnosis?
Checkpoint
An interviewer asks why you’d use 12 heads of d_head=64 rather than one head of d=768, given identical FLOPs. Strongest answer?
Checkpoint
Your from-scratch transformer compiles and runs but never learns; the attention weights look nearly uniform and loss is flat. Which bug is most consistent with this?
Checkpoint
Going from a 2k to a 16k context, your single-GPU prefill latency and memory explode far more than 8×. What does the senior answer name?
Checkpoint
You want a model that generates text left-to-right and can be trained efficiently on all positions at once. What makes that possible inside attention?
Could you derive the √d_k scaling, explain multi-head as a shape split, and write the exact attention shapes on a whiteboard?
Takeaways
- Attention = softmax(Q·Kᵀ/√d_k)·V: query·key scores → softmax over keys → weighted sum of values; one hop between any two tokens.
- √d_k is exact variance correction: q·k has variance d_k·σ²; large logits saturate softmax and kill gradients.
- Multi-head is a shape split, not more compute — same FLOPs as one d_model attention; the win is per-head resolution.
- Shapes: (B,T,d_model) → (B,h,T,d_head) → scores (B,h,T,T) → softmax over keys → @V → merge; the T×T matrix dominates.
- Two quadratics: compute O(T²·d) and HBM I/O O(T²·d + T·d²) — the I/O term is the long-context wall.
Next: embeddings & positional encoding — what embeddings encode, pooling, learned vs RoPE, and how position caps context length.
Sources
Free to read · better with Enzo
Learn it with Enzo
Save your progress, answer the checkpoints, and let Enzo quiz you on what you just read.