Skip to content
Snula
Curriculum
RU Open

Curriculum

Week 11. Inference

Phase 3. Inference, GPUs, scaling · week 11 of 24

Learn in the app: tutor, coding problems →

Core: prefill and decode, the KV cache, continuous batching and PagedAttention, speculative decoding, the cost of a decode step · Depth: MLA, hybrid windows and MTP (05-ГЛУБИНА), track D · ≈ 13 h core / 24 h total

A model is trained once but serves millions of requests, so this is where most of the money goes. Nearly everything this week follows from one fact: to emit a single token, the model reads all of its weights and its entire cache from memory. Once you get this fact, you can derive batching, GQA and speculative decoding yourself.

Step 1. Prefill and decode

In plain terms. Llama-3-8B in bf16 weighs 16 GB. An H100 reads memory at 3.35 TB/s, so one pass over the weights takes at least ~4.8 ms. For a single user that caps you at about 200 tokens per second, no matter how fast the GPU multiplies. But if the same pass serves 64 users, the weights are read once and you get 64 tokens. There is 64 times more arithmetic, and the time is almost the same.

  • Prefill vs decode. Prefill is compute-bound, decode is memory-bound. Almost everything else follows from this. Prefill processes the whole prompt in one pass: each read of a weight serves thousands of tokens. Decode emits one token per sequence: each read of a weight serves B tokens. The intensity (FLOPs per byte) grows with the batch: details in week 13 and in problem D17

Step 2. KV cache

In plain terms. Without a cache, every step would have to recompute keys and values for the entire prefix. But the K and V of an old token do not change when new tokens arrive: you can compute them once and store them. A new token computes only its own q, k, v and attends to the stored ones. You pay in memory: for Llama-3-8B it is 128 KiB per context token, 1 GiB per 8k-token sequence.

  • KV cache: why, how big, how it grows with context: 2 · B · S · L · K · H · bytes (derived in week 9), linear in S and in B. On an 80 GB H100, after the 15 GiB of Llama-3-8B weights there is room for roughly 60 sequences of 8k. The batch is limited by the cache, not by compute
Prefill and decode: one wide pass and many narrow onesPrefill and decode: one wide pass and many narrow ones
Diagram 16. Prefill is one wide pass (compute-bound); decode is narrow steps of one token each (memory-bound). The KV cache grows by one position per step; its size is 2·B·S·L·K·H·bytes.
  • Shrinking the KV cache: GQA/MQA, MLA (DeepSeek), sliding window, cache quantization (storing K and V in 8 or 4 bits instead of 16; how to quantize and what it costs, in detail in week 14), offloading. GQA shrinks the cache by a factor of N/K, a window caps it at the window size, 8-bit quantization halves it

Depth (optional). MLA (multi-head latent attention, attention through a compressed latent; DeepSeek-V2) stores a short latent per token and layer instead of K and V: 576 numbers versus 2048 with GQA. Hybrid local and global attention grows the cache only in the global layers. The full treatment is in 05-ГЛУБИНА, section "Week 11. Shrinking the cache and drafting: MLA, hybrid windows, MTP".

  • Batching: static → dynamic → continuous batching; PagedAttention (vLLM). Static batching waits for the longest response in the batch. Continuous batching schedules per step: a finished request leaves, and a new one takes its place. PagedAttention stores the cache in blocks with a block table, like pages of virtual memory: no reservation for the maximum length, no fragmentation, and a shared prefix for many requests
  • Prefix caching (caching shared prefixes): if the start of a prompt has been seen before, its K and V are not recomputed. Prefixes are stored in a tree keyed by tokens; cold branches are evicted by LRU (least recently used first), and some go to CPU memory. The router must send requests with a shared prefix to the same replica, or the cache will not be found. Example: 20 conversations sharing a 1,000-token system prompt. Without the cache: prefill of 20,000 tokens and 2.6 GB of cache; with it: 1,000 tokens and 131 MB

Step 3. Speculative decoding

In plain terms. The target model p gives words A and B probabilities 0.6 and 0.4; the draft q gives 0.9 and 0.1. The draft proposes A 90% of the time; we accept with probability 0.6/0.9 = 2/3, so A comes out 60% of the time. The draft proposes B 10% of the time, and p(B) > q(B), so we always accept. Rejections: 0.9 · 1/3 = 30%, and on rejection we sample from the residual max(0, p − q) = (0, 0.3), that is, only B. B comes out 10% + 30% = 40% of the time. Exactly p, with no approximation.

  • Speculative decoding: a draft model + verification; why the distribution stays exact (work through the accept/reject step). Accepted from the draft: q(x)·min(1, p(x)/q(x)) = min(p, q). Rejection: the excess of q over p and of p over q sum to the same amount (both distributions sum to 1), the normalization cancels, and max(0, p − q) remains. Total: min(p, q) + max(0, p − q) = p. Variants: Medusa, EAGLE, n-gram, and the MTP (multi-token prediction) modules of DeepSeek-V3: depth, the same section of 05-ГЛУБИНА
  • Why verification is almost free: γ + 1 tokens in one pass read the weights once, just like a normal step. As long as B·(γ + 1) tokens stay below the critical batch (about 295 for an H100; details in week 13), the extra arithmetic hides behind the reads. At large batch there is no slack, and the gain disappears
Speculative decoding: the draft writes, the target model checksSpeculative decoding: the draft writes, the target model checks
Diagram 17. The draft writes γ tokens, and the target model checks them in one pass. On rejection, the token is drawn from the residual max(0, p − q); if everything is accepted, a bonus token is added; the result is distributed exactly as p.
  • Metrics: TTFT, TPOT, throughput vs latency. TTFT (time to first token) is queueing plus prefill; TPOT (time per output token) is one decode step. A larger batch raises throughput and worsens TPOT for each user

Step 4. What a decode step and a million tokens cost

In plain terms. Llama-3-8B on an H100, 8k context. Every step reads 16.06 GB of weights: 4.79 ms at 3.35 TB/s. The cache of one sequence (1.07 GB) takes another 0.32 ms to read. One user: a 5.1 ms step, 196 tokens per second. 32 users: the weights still take 4.79 ms, the cache 32 · 0.32 = 10.3 ms, the step is 15.1 ms, but each step yields 32 tokens: 2,126 tok/s. The batch grew 32×, the step 3×, the output 11×.

  • Decode step time, a lower bound: t ≈ B·KV/W + max(weight bytes / W, 2·N·B / P), where W is memory bandwidth, P is peak FLOP/s, KV is the cache bytes of one sequence and N is the parameter count. Why reading and compute are compared through a maximum is explained by the roofline model: details in week 13. The cache is always read, and there is no one to share it with. The weights are memory-bound while B is below the critical batch (about 295 tokens for an H100). In the example, 2·N·B/P at B = 32 is only 0.49 ms versus 4.79 ms for reading the weights. A real step is usually 1.2–2× slower than the estimate
  • The cache catches up with the weights at B = weight bytes / KV: at 8k that is 15 sequences (week 9). Beyond that, the step grows almost linearly in B, and output hits a ceiling of W / KV ≈ 3,120 tok/s, however much you grow the batch. int8 for weights and cache halves both parts: at B = 1 the step is 2.6 ms, and 80 GB holds 134 sequences of 8k instead of 59
  • The price of a token. GPUs are paid for by the hour, not by the token: $ per 1M = hourly price / (tok/s · 3600) · 10⁶. At a notional $3 per H100-hour: B = 1 gives $4.26 per million output tokens, B = 8 $0.77, B = 32 $0.39
  • The language of costs from microeconomics. Training is a fixed and sunk cost, inference a variable one: it grows with the number of tokens. Inside a step there is the same split: reading the weights is the fixed part (4.79 ms for the whole batch), the cache the variable part (0.32 ms per user). The average cost per token, 4.79/B + 0.32 ms, falls with the batch toward the marginal cost (the price of one more request, 0.32 ms): a floor of about $0.27 per million. That is why training a small model beyond Chinchilla pays off (week 14): variable costs go down
  • Latency versus throughput. A user's TPOT equals the step: 5.1 ms at B = 1, 15.1 ms at B = 32. Each B gives a pair (TPOT, price), and all these pairs lie on a Pareto front (the set of options where you can improve one only by worsening the other): cheaper means slower. The service requirement makes the choice. With TPOT of at most 10 ms we take B = 16 (9.9 ms, $0.52). For chat, B = 48 is enough: a 20.2 ms step, 50 tok/s per person, faster than they read, and $0.35. The same front helps choose between models: details in week 18
  • Laying out inference across GPUs (decode only with tensor parallelism, prefill on separate machines): details in week 14

Common mistakes

  • Forgetting the RoPE offset with a cache, so every new token gets position 0. Caught by test_rope_offset_matches_slicing, test_kv_cache_matches_full_forward ("if this test fails, the RoPE offset is usually to blame")
  • is_causal=True at T = 1: the SDPA mask is aligned to the top-left corner, so the new token sees only the first key. Hence is_causal=(T > 1) in Attention.forward; caught by test_kv_cache_matches_full_forward
  • Advancing cache.length inside each layer's update rather than once per pass: positions drift apart. Caught by test_kv_cache_matches_full_forward: the test calls advance itself after the step
  • On rejection, sampling from p instead of the residual: tokens the draft overrated come out too often. Caught by test_residual_zeroes_tokens_the_draft_overrated, test_output_distribution_equals_target_exactly
  • Estimating a decode step by FLOPs: at B = 32 arithmetic takes 0.49 ms out of 15 ms; the rest is reading weights and cache. Caught by test_decode_step_llama3_h100_batch_32_numbers

Code → nanolm/speculative.py: residual_distribution, accept_or_resample, speculative_generate. The KV cache already exists in modules.py: KVCache (update, advance, nbytes), offset in Attention.forward; generation with prefill and decode in NanoLM.generate. Metrics: SpecDecodeStats.acceptance_rate and tokens_per_target_forward; the baseline is baseline_generate. The task is in exercises_en/speculative.py: NANOLM_IMPL=exercises_en pytest tests/test_speculative.py -v. Step time and token price are in nanolm/budget.py: decode_step_time, cost_per_million, critical_batch, kv_bytes_per_token, hardware H100 and A100; tests in tests/test_budget.py. Exercise: plot the curve "TPOT versus $ per 1M" over B from 1 up to the memory ceiling for your model, and mark the point you would choose for a chat and for an agent with a strict TPOT. Code (depth) → nanolm/mla.py: MLA with the MLACache cache (per token only the latent c and the shared k_rope); the paths naive_attention and absorbed_attention must match. The task is in exercises_en/mla.py: NANOLM_IMPL=exercises_en pytest tests/test_mla.py -v. Read the MLA section in 05-ГЛУБИНА first.

The main test of the week is test_output_distribution_equals_target_exactly: 60,000 samples, and the empirical distribution must converge to the target. If it does, you have proved in code what you derived on paper earlier.

Back-of-the-envelope problems. Solve without a calculator, to the order of magnitude; answers below. A 13B-parameter model in bf16, 40 layers, GQA with 8 KV heads, H = 128, one H100 (80 GB, 3.35 TB/s of bandwidth, that is 3.35 GB/ms, 989 TFLOP/s); we leave 6 GB of memory for activations and fragmentation.

  1. How much does the KV cache weigh per token and per sequence of 4,096 tokens?
  2. What is the maximum batch at 4k context?
  3. What is the speed ceiling for a single user?
  4. How long is a decode step at B = 32 and 4k context, how many tokens per second does it give, and what limits it?
  5. At B = 32, which pays off more: int8 weights or an int8 cache? From what batch size does reading the cache take longer than reading the weights?

Answers. (1) 2·40·8·128·2 = 163,840 bytes, 160 KiB per token; × 4,096 ≈ 0.67 GB per sequence. (2) The weights take 26 GB, leaving 80 − 26 − 6 = 48 GB for the cache, 48 / 0.67 ≈ 71. (3) One pass over the weights takes 26 / 3.35 ≈ 7.8 ms, so no more than 129 tokens per second. (4) The cache is 32 · 0.67 ≈ 21.5 GB and takes 6.4 ms to read; with the weights the step is 14.2 ms, 32 / 0.0142 ≈ 2,260 tokens per second. The arithmetic takes 2·13·10⁹·32 / 989·10¹² ≈ 0.84 ms: the step is memory-bound. (5) int8 weights: (13 + 21.5) / 3.35 ≈ 10.3 ms; int8 cache: (26 + 10.7) / 3.35 ≈ 11.0 ms, so the weights win for now. The cache outweighs the weights when B · 0.67 > 26, that is from B ≈ 39; with a longer context the threshold is lower.

Math (Track D): D13: reservoir sampling and a proof by induction; D14: exactness of speculative decoding and the expected speedup, accounting for the cost of the draft.

Interview question of the week: "An 8B service responds slowly. What do you do?" A 3-minute structure: separate TTFT and TPOT → TTFT: queueing, prompt length, prefix caching → TPOT: decode is memory-bound, so fewer bytes (quantizing weights and cache, GQA) and more tokens per read (continuous batching, speculative decoding) → the batch ceiling is set by the cache (PagedAttention) → check quality after each measure.

Sources: Pope et al., Efficiently Scaling Transformer Inference (2022); Kwon et al., PagedAttention (2023); Yu et al., Orca (2022); Leviathan et al., Speculative Decoding (2022); Zheng et al., SGLang (2024), prefix caching.

Deeper: 05-ГЛУБИНА, sections "★★ Week 11. Proof that speculative decoding is correct" and "Week 11. Shrinking the cache and drafting: MLA, hybrid windows, MTP".

Week outcomes

  • I can explain through arithmetic intensity why prefill is compute-bound and decode is memory-bound.
  • I can prove on paper in 10 minutes that speculative decoding samples exactly from the target distribution.
  • I can implement a speculative_generate that passes test_output_distribution_equals_target_exactly.
  • I can compute the maximum batch for a given amount of memory for the KV cache.
  • I can estimate a decode step and the price of a million tokens on a napkin, and explain why the price falls with the batch.

Self-check

  1. Why does continuous batching beat static batching on a chat workload?
  2. What memory problem does PagedAttention solve?
  3. What does the speedup from speculative decoding depend on, and when is there none?
  4. By what factor does GQA with K groups shrink the cache compared to MHA, and what else shrinks it?
  5. Why does the average price per token fall with the batch, and where is its floor?

In the app each week has skills to rate yourself on, questions with answer checking, Python coding problems and a tutor grounded in the course.

Learn in the app: tutor, coding problems
← PreviousWeek 10. Training Next →Week 12. Sampling strategies

Snula
Snula: LLMs from scratch

  • Home
  • Curriculum
  • App
  • Privacy
  • Terms

The course text is licensed under CC BY-NC-SA 4.0, nanolm code under Apache-2.0.