Lesson 5 of 7 · 48 min
Inference lifecycle: prefill, decode & the KV cache
What actually happens when you call the model: the autoregressive loop, the prefill/decode split and why their costs are inverted, KV-cache mechanics and the exact memory formula, why decode is bandwidth-bound, and how temperature/top-k/top-p shape the output — with real per-token byte counts.
Why the first token is slow and the rest fly
T tokens is processed in a single parallel forward pass — all positions at once, under the causal mask — computing Q/K/V at every layer and producing logits (you only need the last position’s logits to start generating). This is where the KV cache is built: the K and V tensors for every input token at every layer, stored so they don’t have to be recomputed. Decode: you append the sampled token and run one new position through the model; its query attends over the cached K/V plus its own new K/V, you sample the next token, repeat until a stop token or length cap.1ONE CALL = PREFILL + DECODE23 prefill ── all T prompt tokens in ONE parallel pass ──► build KV cache (per layer)4 COMPUTE-BOUND, parallel; cost ~ proportional to T; this is most of TTFT56 decode ── tok ── tok ── tok ── ... ── <stop>7 each step: 1 new query attends over ALL cached K/V, sample, append8 MEMORY-BANDWIDTH-BOUND, sequential; per-step cost ~ reading the KV cache910 TTFT ~= prefill time (shrink/cache the PROMPT to improve it)11 total ~= TTFT + out_tokens x inter-token-latency (shrink the OUTPUT to improve it)1213 Real systems spend most wall-clock in DECODE -> the KV cache is the cost center.</stop>t would recompute K and V for all t−1 previous tokens at every layer, every step — turning generation into an O(T²) cumulative recompute. Caching K/V makes each decode step O(T) work (one new query against cached keys) instead. The measured payoff is large: Hugging Face’s benchmark on a T4 GPU went from 61s to 11.7s — a 5.21× speedup — just by enabling use_cache=True. Interview angle. “Why does the KV cache only exist at inference, not training?” → training runs one parallel forward pass over the whole sequence (teacher forcing), so there’s no token-by-token reuse to cache; the cache is purely an autoregressive-decoding optimization.
Mastering LLM Inference Optimization: From Theory to Cost-Effective DeploymentAI Engineer (Mark Moyou, NVIDIA)The KV-cache memory formula — exact, with real numbers
KV_bytes = 2 · L · B · T · H_kv · D_h · dtype_bytes. The leading 2 is for K and V; L = layers, B = batch, T = sequence length, H_kv = KV heads (H for MHA, 1 for MQA, G for GQA — Lesson 6), D_h = head dim, and dtype_bytes = 2 for FP16/BF16, 1 for INT8. It grows linearly in T and B, but the per-token constant is heavy.1KV CACHE = 2 * L * B * T * H_kv * D_h * dtype_bytes (factor 2 = K and V)23 Llama 7B (L=32, H_kv=32, D_h=128, FP16):4 per token = 2*32*32*128*2 = 524,288 bytes = 0.5 MB5 T=4096, B=1 -> 2 GB T=4096, B=8 -> 16 GB67 GPT-2 124M (L=12, H_kv=12, D_h=64, FP16):8 per token = 2*12*12*64*2 = 36,864 bytes ~= 36 KB9 T=1024, B=1 -> ~36 MB1011 Llama 2 70B (L=80, H_kv=8 via GQA, D_h=128, FP16):12 per token = 2*80*8*128*2 = 327,680 bytes ~= 0.31 MB13 T=128k, B=1 -> ~39 GB (KV cache ALONE -> the floor for 128k serving)1415 Note 70B has FEWER KV bytes/token than 7B: GQA (8 vs 64 heads), Lesson 6.0.31 MB/token × 32,000 × 4 ≈ 40 GB — narrate the formula, plug the numbers, and note GQA is why it’s not catastrophic.At long context, the question stops being “how big is the model” and becomes “how big is the cache.” The KV cache — linear in context and batch, heavy per token — is the memory and bandwidth term that decides what you can actually serve.
Why decode is bandwidth-bound, not compute-bound
Key idea
Speculative decoding: beating the bandwidth wall
k tokens, and the big target model verifies them all in a single parallel forward pass (verification is parallel, just like prefill). Accepted tokens are kept; the first rejection truncates and the target’s own distribution takes over. Because the costly target pass now confirms several tokens at once, you get a 2–3× throughput win with identical output distribution to plain sampling — it’s exact, not lossy.Key idea
Sampling: temperature, top-k, top-p
probs = softmax(logits/T) — T<1 sharpens toward the argmax (more deterministic, more repetitive), T>1 flattens (more varied), T=0 is greedy. Top-k keeps only the k highest-probability tokens before sampling (k=1 = greedy). Top-p (nucleus) keeps the smallest set of tokens whose cumulative probability exceeds p — adaptive: a small nucleus when one token dominates, a large one when the distribution is flat. They compose: apply temperature, then top-k/top-p, then renormalize and sample.- 01Extraction / classification / structured output → temperature 0–0.2 (you want the single most likely answer).
- 02Balanced assistant / Q&A → ~0.7 (some variety, still grounded); leave top-p near default.
- 03Brainstorming / creative drafts → 0.9–1.2 (explore the distribution).
- 04Tune temperature OR top-p, not both at once — they interact and make behavior hard to reason about.
seed helps reproducibility but is best-effort, not a contract. Downstream code must never assert output == expected_string; design for variation with schemas and eval gates.Common mistake
“The KV cache speeds things up, so bigger context is basically free at inference.”
Interview prep
- 01“Walk me through an LLM call.” → tokenize → prefill (parallel, builds KV cache, compute-bound, sets TTFT) → decode (one token/pass, reads KV, bandwidth-bound) → stop.
- 02“Why is the first token slow, then streaming fast?” → prefill processes the whole prompt at once; each decode step is a cheap incremental pass reusing the cache.
- 03“Estimate KV cache for a 70B at 32k, batch 4.” → 2·L·B·T·H_kv·D_h·dtype = ~0.31 MB/token × 32k × 4 ≈ 40 GB; GQA keeps it sane.
- 04“Why is decode memory-bandwidth-bound?” → one tiny matmul per step, but you must stream the whole K/V history from HBM; bandwidth (~2 TB/s on A100) is the wall.
- 05“Why does the KV cache exist only at inference?” → training is one parallel teacher-forced pass; there’s no token-by-token reuse to cache.
- 06“How big is the KV cache vs the weights?” → at long context it rivals or exceeds the weights; it’s the dominant memory cost of serving, not the params.
- 07“Is temperature 0 deterministic?” → no — greedy ≠ deterministic; batching, GPU FP non-associativity, MoE routing, and silent updates all vary it.
- 08“Temperature vs top-k vs top-p?” → temperature rescales logits; top-k keeps k tokens; top-p keeps a cumulative-mass nucleus; tune one, compose carefully.
(B, H_kv, max_seq_len, D_h) per layer for K and V — and pre-allocating to max_seq_len × batch is exactly the waste PagedAttention fixes, Lesson 6); “how does batching interact with the cache?” (each sequence carries its own cache; naive static batching wastes slots, which is why continuous batching + paging matter); “KV cache vs beam search memory?” (beam search multiplies the cache by the beam width — often why beam search is too expensive to serve). The applied-scientist bar compresses this to “KV caching + speculative decoding (2–3×) + INT8/FP8 tradeoffs” — be ready to bind each to a number.Common mistake
The red-flag answer: “the KV cache just stores the past so the model sees it.”
Checkpoint
Two requests to the same model: (A) 4,000-token prompt, 50-token answer; (B) 400-token prompt, 2,000-token answer. Which feels slower end-to-end, and why?
Checkpoint
You must serve a Llama-2-70B (GQA, 8 KV heads, 80 layers, D_h=128, FP16) at 32k context, batch 4. Roughly how much GPU memory does the KV cache alone need?
Checkpoint
During decode your GPU shows low compute utilization but is clearly the bottleneck. What’s the most accurate explanation?
Checkpoint
A teammate sets temperature=0 and writes `assert output == golden_string` in a CI test. It passes locally, flakes in CI. Best fix and reasoning?
Checkpoint
You enable beam search (width 4) for higher-quality generation and immediately hit OOM at a context that was fine with greedy decoding. Why?
Could you walk the prefill/decode split, derive the KV-cache memory for any model, and explain sampling without claiming false determinism?
Takeaways
- A call is prefill (parallel, compute-bound, builds the KV cache, sets TTFT) + decode (sequential, bandwidth-bound, sets total time).
- KV cache = 2·L·B·T·H_kv·D_h·dtype_bytes — linear in context and batch; 0.5 MB/token for Llama 7B, ~39 GB for a 70B at 128k.
- The cache avoids O(T²) recompute (5.21× on a T4) but costs memory; at long context it’s the dominant serving cost, not the weights.
- Decode is memory-bandwidth-bound (stream the cache from HBM); prefill is compute-bound — different levers for each.
- Sampling is three layers on one softmax (temperature, top-k, top-p); temperature 0 is greedy, not deterministic.
Next: attention & serving at scale — MQA/GQA, FlashAttention, sliding-window, PagedAttention, quantization, and the engines that ship them.
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.