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 is the entire reason transformers replaced RNNs: it gives every token a direct, O(1)-length path to every other token, so information no longer has to crawl through a recurrence. But “explain self-attention” is also the single most common opener in an LLM internals round — and the gap between a strong and a weak answer is razor-sharp. This lesson derives the mechanism from the math up, with the exact tensor shapes you’ll be asked to put on a whiteboard, so you can answer the Q and volunteer the follow-ups.
The full equation, from Vaswani et al. (2017): 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.
Why is this such a leap over the RNN it replaced? An RNN compresses all prior context into a single fixed-size hidden state passed step by step, so information from token 1 reaches token 500 only after 499 sequential, lossy hops — the bottleneck that capped sequence modeling for years. Attention deletes the bottleneck: token 500’s query can read token 1’s value directly, in one hop, with no compression. That is the “O(1)-length path between any two positions” framing — and it’s also why the whole sequence can be processed in parallel during training (every position computes its attention simultaneously), which is what made training at scale tractable. The cost of that parallel, all-pairs reach is the quadratic we pay for later in this lesson.
Interview angle. The weak answer stops at “there are three matrices Q, K, V.” The strong answer says the same three things in order: (1) it’s content-based mixing that maps a sequence to a same-length sequence; (2) per-pair query·key scores → softmax over keys → weighted sum of values; (3) every token reaches every other in one hop, killing the recurrence bottleneck; (4) decoder self-attention masks the future. Then immediately volunteer the masking variant before they ask — that signals you actually understand it rather than reciting.
Let’s build GPT: from scratch, in code, spelled outAndrej Karpathy

Why divide by √d_k — the exact variance argument

This is the question where interviewers separate “normalization trick” (weak) from a real derivation (strong). Assume each component of 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.
code
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.
Interview angle. Follow-ups probe whether you actually derived it: “What distribution of q·k are you assuming?” (zero-mean, variance d_k·σ²). “Why √d_k and not d_k?” (you’re correcting the standard deviation of the sum, which scales as √d_k). “Does it still matter after LayerNorm?” (yes — LayerNorm controls the input scale, but the dot-product-of-d_k-terms growth happens inside attention regardless). If you can write the one-line variance equation, you’ve passed; if you say “it keeps numbers from blowing up,” you’ve flagged yourself as surface-level.

Multi-head: a shape split, not more compute

Multi-head attention is widely misexplained as “running several independent attentions for more capacity.” It is not more compute — it is the same compute reorganized. 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.
So what do you gain? Resolution. A single softmax distribution must average all the positions it attends to into one resolved output per token — fine-grained, multi-relation structure gets blurred. Splitting into 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.

The exact tensor shapes through one block

Whiteboard these until they’re automatic — “implement MHA / a transformer layer” is a standard applied-scientist coding round, and the bugs they’re watching for are shape bugs: wrong softmax axis, missing transpose, forgetting .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).
code
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).
python
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.
The causal mask is the difference between a decoder (GPT) and an encoder (BERT). Setting 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

Attention has two quadratic costs, and conflating them is the most common production-analysis bug. Compute is 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).
code
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.

What attention heads actually learn

A useful sanity check on the resolution argument: when researchers probe trained models, individual heads turn out to specialize on interpretable relations — one head tracks the previous token, another links verbs to their objects, another resolves coreference, another attends to the delimiter or the first token (a learned “no-op” sink). This is empirical evidence for why multi-head works: the heads differentiate during training because identical heads waste capacity. It also grounds a senior intuition — attention is not a black box; you can attribute behavior to specific heads, which is the basis of mechanistic interpretability and model steering.
This specialization is also what efficient-serving variants trade away. MQA collapses all key/value heads to one, and GQA to a handful (Lesson 6): you keep the 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.

Interview prep

Nearly every internals round opens with attention. The rubric rewards mechanism-first answers that bind the math to a number or a shape, and that volunteer the variant (masking, multi-head, the quadratic) before being asked. Be able to write softmax(Q·Kᵀ/√d_k)·V and the variance argument on a whiteboard.
  1. 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.
  2. 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.”)
  3. 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.
  4. 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.
  5. 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).
  6. 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.
  7. 07“Self vs cross attention?” → self: Q,K,V from one sequence; cross: queries from the decoder, keys/values from the encoder.
  8. 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.
Going deeper, the follow-ups that reward depth: “initialize W_Q=W_K=W_V to identity — what happens?” (attention becomes a function of raw token similarity; with no scaling the logits can still saturate); “the model compiles but won’t learn — where do you look?” (softmax on the wrong axis, a double-softmax before cross-entropy, a broadcasting bug in the mask, an off-by-one causal mask — the classic human-debugger round); “how does the causal mask interact with KV-cache reuse?” (at decode you append one new query and attend over all cached keys/values, so the mask is implicit — the new token only sees the past, Lesson 5). Always lead with the mechanism, then the shape or the number.
paperAttention Is All You Need (the transformer paper)Vaswani et al. (arXiv)reponanoGPT — model.py (causal self-attention in ~20 lines)Andrej Karpathydocsd2l.ai — Attention Scoring Functions (the √d_k derivation)Dive into Deep LearningvideoStanford CS25 V1 — Transformers United (overview)Stanford Online

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?

AThe learning rate is too low — raise itBWithout √d_k, q·k has variance d_k·σ² so logits are large, softmax saturates toward one-hot, and its gradient ≈ 0 — restore the divisorCThe value matrix V needs its own normalization
Sign up free to answer and see why

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?

AMultiple heads do more total computation, so they’re strictly more powerfulBIt parallelizes better across GPU coresCResolution: one softmax averages all attended positions into a single blurred output, while separate heads specialize on different relations at the same cost
Sign up free to answer and see why

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?

Asoftmax is applied over the wrong dimension (e.g. over queries/rows instead of over keys/last dim)BThe dataset is too smallCYou forgot weight decay
Sign up free to answer and see why

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?

AThe model has more parameters at longer contextBAttention is quadratic: the T×T matrix and its HBM I/O grow ~64× from 8× more tokens — compute O(T²·d) and memory O(T²·d + T·d²)CTokenization overhead scales with the square of input length
Sign up free to answer and see why

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?

ABidirectional attention with a masked-language-modeling objectiveBA causal mask (set scores for future positions to −∞ before softmax) so each position attends only to itself and earlier tokensCRemoving positional information so order doesn’t matter
Sign up free to answer and see why

Could you derive the √d_k scaling, explain multi-head as a shape split, and write the exact attention shapes on a whiteboard?

New to itGetting thereConfident

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.