Lesson 7 of 7 · 46 min
Capstone: annotate a forward pass
Assemble the whole track into one exercise: trace tensor shapes and cost through a mini-GPT from token ids to logits, then size a real serving deployment — KV-cache memory, what fits on which GPU, and the levers to hit a latency/throughput target. The canonical LLM-engineer internals interview, rehearsed end to end.
Capstone
Annotate a forward pass, then size a deployment
vocab=50,257, d_model=768, n_layer=12, n_head=12, d_head=64, FP16.
Let’s reproduce GPT-2 (124M) — the full mini-GPT in codeAndrej KarpathyPart 1 — trace the forward pass, shape by shape
1FORWARD PASS: "the cat sat" -> next-token logits (GPT-2 124M, FP16)23 text "the cat sat"4 | tokenizer (byte-level BPE, ~4 chars/token) [L1]5 v6 ids (B, T) = (1, 3) integers7 | token embedding table (50257 x 768), tied to unembed [L1, L3]8 v9 x (1, 3, 768) floats10 | + position info (RoPE rotates Q,K; learned-abs adds) [L3]11 v12 --- repeat for each of 12 blocks --------------------------------[L4]13 x -> RMSNorm/LN (1, 3, 768) pre-norm [L4]14 | Q,K,V projections each (1, 3, 768) [L2]15 | reshape to heads (1, 12, 3, 64) 12 x 64 = 76816 | Q @ K^T / sqrt(64) (1, 12, 3, 3) <- T x T [L2]17 | causal mask + softmax over keys, @ V [L2]18 | merge heads (1, 3, 768) -> W^O [L2]19 x = x + attn(...) (1, 3, 768) residual add [L4]20 x = x + ffn(norm(x)) FFN 768->3072->768 (SwiGLU) [L4]21 ----------------------------------------------------------------22 | final norm + unembed (768 -> 50257) [L1, L3]23 v24 logits (1, 3, 50257) -> take last row, sample [L5]T and thus everything downstream; the embedding table (L1/L3) is 38.6M params, tied to the unembed. Each block (L4) is pre-norm + attention + FFN on a clean residual stream; inside attention (L2) the (1,12,3,3) score matrix is the T×T object whose memory grows quadratically and which FlashAttention refuses to materialize. The FFN is 4× wide and holds most of the FLOPs. The final unembed produces (B,T,vocab) logits; you sample from the last position (L5). Interview angle. If asked “where does the cost live?” → attention’s T×T term grows with context, but the FFN dominates per-token FLOPs at short context — name both.T=3 (all prompt tokens at once). At each decode step the new token is a single position — T_new=1 — whose query is (B,1,d_head) per head, but it attends over all cached keys, so the score is (B, n_head, 1, T_cached), not a full T×T. That asymmetry is the whole story: prefill does the big parallel T×T work once (compute-bound, sets TTFT), then every decode step is a thin 1×T_cached read against the growing cache (bandwidth-bound, sets total time). Same diagram, two very different cost profiles — which is exactly the prefill/decode split from Lesson 5 made concrete on the shapes.1# A mini-GPT forward, annotated -- the shapes match the diagram above.2import torch, torch.nn.functional as F34def block(x, ln, qkv, proj, ffn, n_head): # x: (B, T, C)5 B, T, C = x.shape6 h = ln(x) # pre-norm [L4]7 q, k, v = qkv(h).chunk(3, dim=-1) # each (B, T, C) [L2]8 dh = C // n_head9 q, k, v = [t.view(B, T, n_head, dh).transpose(1, 2) for t in (q, k, v)] # (B,h,T,dh)10 att = (q @ k.transpose(-2, -1)) / dh**0.5 # (B, h, T, T) [L2]11 mask = torch.tril(torch.ones(T, T, device=x.device)).bool()12 att = att.masked_fill(~mask, float("-inf")).softmax(-1) # over keys13 y = (att @ v).transpose(1, 2).reshape(B, T, C) # merge heads14 x = x + proj(y) # residual [L4]15 x = x + ffn(x) # FFN residual (norm inside ffn) [L4]16 return x17# stack 12 of these, add embed/unembed, and you have a GPT. Shapes never change at C=768.Part 2 — the cost of one call (prefill + decode)
1ONE REQUEST: 1,200-token prompt -> 250-token answer (GPT-2-class)23 prefill: 1,200 tokens, ONE parallel pass -> sets TTFT (compute-bound) [L5]4 decode: 250 sequential steps, read KV cache -> sets most of total (bw-bound)[L5]56 KV cache built during prefill, per the formula: [L5]7 GPT-2 124M: 2 * 12 * 1 * T * 12 * 64 * 2 = 36,864 bytes/token (~36 KB)8 at T=1,450 (1200 + 250) -> ~52 MB (tiny -- a 124M model)910 total_latency ~= TTFT(prefill 1,200) + 250 * inter_token_latency11 cheapest feature: big CACHED prompt + SHORT output (output ~3-5x input cost) [L5]Part 3 — size a real serving deployment
d_head=128, FP16) for 32k context. The three line items: weights, KV cache, framework overhead. Weights ≈ 140 GB in FP16 (so this is a multi-GPU model regardless). KV cache per token ≈ 2·80·8·128·2 = 0.31 MB; at 32k context that’s ~10 GB per concurrent request; at batch 4, ~40 GB. Interview angle. The senior framing: KV-head count (GQA) is why a 70B has fewer KV bytes/token than a 7B, and the KV cache — not the weights — is what scales with your concurrency and context.1SERVING SIZING -- Llama-2-70B, GQA-8, FP16, 32k context23 weights .................. ~140 GB (FP16) -> multi-GPU no matter what4 KV / token ............... 2*80*8*128*2 = 0.31 MB [L5/L6]5 KV @ 32k, 1 request ...... 0.31 MB * 32,000 ~= 10 GB6 KV @ 32k, batch 4 ........ ~40 GB -> grows with concurrency78 LEVERS to fit / hit SLO: [L6]9 FP8 weights .............. 140 GB -> ~70 GB (Hopper; ~2x matmul too)10 FP8 / INT8 KV cache ...... ~0.31 -> ~0.15 MB/token (watch >64k drift)11 GQA (already on) ......... 8 vs 64 heads = 8x smaller cache than MHA12 PagedAttention ........... waste 60-80% -> <4%; pack concurrent requests13 continuous batching ...... keep the GPU saturated (Orca 36.9x)14 FlashAttention(-3) ....... exact, IO-aware; FP8 attn on HopperT, the position scheme set the context ceiling, the block wiring made it trainable, and the KV/attention choices set what you can serve.Key idea
T and the spelling/arithmetic ceiling; position fixes the context ceiling; pre-norm + RMSNorm fix trainability and per-token cost; KV-head count fixes serving bytes; sampling is the only knob users see. Pick a tokenizer and you’ve picked your worst-case tasks; pick an attention scheme and you’ve picked your serving-cost curve.Part 4 — what breaks at 3am
1PRODUCTION FAILURE MODES (and the lesson behind each)23 Symptom Root cause Caught by / fix4 ------------------------------ -------------------------------- -----------------------------5 OOM under load spike KV cache scales with concurrency admission control + paging [L5/6]6 p99 latency cliff at long ctx quadratic attention / cache reads cap context; FlashAttn; GQA [L2/6]7 "it got dumber" overnight silent provider/model swap pin version + eval gate [L5]8 garbage on rare inputs glitch token in the vocab filter / retrain tokenizer [L1]9 long-ctx retrieval rots naive RoPE scaling / KV quant needle-in-a-haystack gate [L3/6]10 cost spikes with no traffic rise output-token blowup / reasoning cap output; route models [L5]1112 Most 3am pages are a KV-cache or a quadratic-attention story in disguise.Part 5 — the decisions an interviewer will push on
- 01“Walk me through a forward pass.” → ids → embed (tied) → +position → 12× [pre-norm → attn (T×T) → FFN] → final norm → unembed → sample.
- 02“Where does the cost live?” → attention T×T grows with context (compute + HBM I/O); FFN dominates per-token FLOPs at short context — name both.
- 03“Size the KV cache for this model/context/batch.” → 2·L·B·T·H_kv·D_h·dtype; narrate, plug, note GQA shrinks H_kv.
- 04“What fits on one A100?” → weights + KV(B,T) + framework; for 7B FP16 at 4k, batch 8 ≈ 32 GB (near 40 GB); GQA/FP8 to push it.
- 05“Hit a 200ms p95 / max-throughput SLO — what do you change?” → TTFT: FlashAttention + continuous batching + cache prompt; throughput: GQA + FP8 + the right engine.
- 06“Why is a 70B cheaper to serve per-token than you’d think?” → GQA: 8 KV heads vs 64 → fewer KV bytes/token than a 7B MHA model.
- 07“Extend this to 128k context — what breaks?” → KV cache (~39 GB at 70B) and RoPE scaling; validate with needle-in-a-haystack, add paging/quantized cache.
- 08“temperature 0 in your pipeline — safe to assert exact output?” → no; greedy ≠ deterministic; assert schema/semantics and gate with evals.
Common mistake
The red-flag answer: treating model quality and serving cost as separate, late-stage concerns.
Checkpoint
Tracing a forward pass, an interviewer points at the (B, n_head, T, T) tensor and asks what it implies at scale. Best answer?
Checkpoint
You must serve Llama-2-70B (GQA-8, 80 layers, D_h=128, FP16) at 32k context, batch 4, and decide if KV cache or weights dominates the per-request scaling. Best framing?
Checkpoint
A Llama-7B FP16 service at 4k context runs fine at batch 1 on a 24 GB GPU but OOMs at batch 8. What changed, and the cheapest fix to raise the ceiling?
Checkpoint
An interviewer asks why open-weight models almost always ship GQA while a frontier lab might train MHA at the same size. Strongest synthesis answer?
Checkpoint
Your team plans to extend a 4k-context 70B to 128k and quantize the KV cache to fit. What single safeguard belongs in the rollout plan above all?
Could you annotate a full forward pass, size a real deployment’s KV cache and GPU fit, and run the synthesis interview end to end?
Takeaways
- A forward pass is ids → embed (tied) → +position → N× [pre-norm → attn (T×T) → FFN] → final norm → unembed → sample.
- Cost lives in two places: attention’s T×T (grows with context) and the FFN (most per-token FLOPs at short context) — name both.
- Serving sizing = weights (fixed) + KV cache (scales with context × batch) + framework; KV, not weights, scales with your traffic.
- The biggest serving levers are architectural: GQA (KV-head count) and continuous batching — both beat a model swap.
- Quality and cost are the same forward-pass decision: tokenizer → tasks + token cost; attention → quality + serving curve; validate changes with needle-in-a-haystack.
You’ve traced every stage of a transformer from token to logit to GPU. Revisit any lesson’s interview prep before a round — the mechanism-plus-a-number habit is the whole game.
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.