CS336 // FIELD MAP
← field map
LECTURE 10 · SCALE IT AND SERVE ITPercy Liang · 2025-05-01 · 82 min

Inference

Stanford CS336 · Spring 2025 · lecture 10 of 17

Transcript: cleaned auto-captions with timestamps

TL;DR — Training is a single pass over a dense block of tokens, so it is compute-limited and life is good. Generation emits one token at a time, so it is memory-limited: you drag every parameter and every byte of KV cache across the bus to produce a single token. Percy derives that in two SymPy blocks — MLP arithmetic intensity is B·T, attention intensity is ST/(S+T), and during generation the batch dimension cancels out of the attention term entirely, which is why batching cannot rescue you. Everything else in the lecture — GQA, MLA, cross-layer attention, sliding windows, SSMs, diffusion LMs, int8, pruning, speculative decoding, PagedAttention — is a different way of moving fewer bytes for the same accuracy. The one thing to remember: for inference, memory traffic is the currency, not FLOPs.

This lecture is the course briefly turning around to look at the other end of the pipeline. Lectures 1–8 built a model and made the training loop fast; lecture 9 started asking how to spend a training budget well. Lecture 10 asks a question none of that answers: once the weights are frozen, what does it actually cost to use them — and why is that cost governed by a completely different bottleneck? It is also, as Percy admits about ninety seconds in, a lecture that is secretly about architecture. Most of the interesting inference wins of the last three years were not systems work at all; they were people changing the model so that serving it moves less memory.

Outline, with timestamps

One ratio explains the entire lecture

The tool Percy uses for everything here is arithmetic intensity: FLOPs performed divided by bytes moved between HBM and the compute units. It is a property of a computation. The accelerator has a matching number — peak FLOP/s divided by memory bandwidth — and the comparison between the two tells you which resource you are actually spending. An H100 does about 989 TFLOP/s of bf16 and moves about 3.35 TB/s, so its crossover sits at roughly 295 FLOPs per byte (10:52). Above that line the silicon is the constraint and you are doing as well as the hardware allows. Below it, the multiply units idle while the memory system fetches operands, and adding FLOPs to your chip buys you nothing.

Work the simplest case by hand. Multiply X (B×D) by W (D×F): read 2BD + 2DF bytes, write 2BF, do 2BDF FLOPs. When D and F dwarf B — thousands against dozens — the ratio collapses to exactly B. The intensity of a matmul is its batch dimension, so an H100 needs B > 295 to saturate. At B = 1 — a matrix–vector product — the intensity is 1: you stream a whole D×F weight matrix across the bus to do 2DF FLOPs on it. That is ~300× off peak, and it is precisely the shape of autoregressive generation.

That is the whole lecture in miniature: training sees every token at once and builds fat tensors, while inference must emit token t before computing t+1, so its tensors stay thin. Percy is explicit that the arithmetic is stylized — perfect overlap, no launch overhead, matmuls only — but the scaling is right.

The KV cache, and the two regimes it creates

Done naively, generation is absurd: re-encode the whole prefix per token, and emitting T tokens costs O(T³) FLOPs (14:19). But under causal masking the keys and values of earlier positions never change, so you compute them once and keep them. That is the KV cache, and it splits inference into two phases of opposite character: prefill ingests the prompt in parallel — the same shape as a training forward pass, compute-limited, fast — and generation extends the cache one token at a time.

Now redo the accounting for a Transformer block, with S tokens conditioned on and T tokens being queried (prefill sets T = S; generation sets T = 1). For the MLP — up, gate, down projections — the FLOPs come to 6BTDF and the traffic to 4BTD + 4BTF + 6DF, and under the same large-D,F limit the intensity is B·T (20:41). Prefill is fine: even a single long prompt clears 295. Generation gives T = 1, so the intensity is just B — the number of concurrent requests. Your efficiency now depends on your traffic, which is an uncomfortable place for an engineer to be.

Attention is worse. With FlashAttention, FLOPs are 4BSTD and bytes are 4BSD + 4BTD, so the ratio simplifies to ST/(S+T). In prefill that is S/2 — good. In generation it is S/(S+1), less than one regardless of context length or user count. Note what is missing: B appears in numerator and denominator and cancels. The reason is structural. MLP weights are shared — read them once, push a hundred sequences through. KV caches are not, so batching multiplies the bytes you must fetch in exact proportion to the work it adds.

"For the attention, the KV cache is every sequence its own unique snowflake."— Percy Liang, 26:44

So: prefill is compute-limited, generation is memory-limited, concurrency partly rescues the MLP half and cannot rescue the attention half at all — not by batching, not by better kernels. The only lever left is to make the cache smaller, which is what the rest of the lecture is about.

Napkin math: Llama 2 13B on one H100

To make it concrete, the executable lecture instantiates a Llama 2 13B configuration — D = 5120, F = 13824, L = 40, N = K = 40 heads of H = 128, V = 32000, S = 1024 — and computes three quantities symbolically (28:31). Parameters at bf16 are about 26 GB. The KV cache costs S · (K·H) · L · 2 · 2 bytes per sequence — two for key-and-value, two for bf16 — which is 0.84 GB per sequence at this context length. Since latency is set by memory traffic, model it as (parameters + B × cache) ÷ bandwidth, and throughput as B ÷ latency.

configmemorylatency / tokenthroughput
B = 1, MHA (K = 40)26.9 GB8.0 ms~125 tok/s
B = 64, MHA79.7 GB23.8 ms~2,690 tok/s
B = 256, MHA240.8 GB — does not fit71.9 ms~3,560 tok/s
B = 64, GQA (K = 8)36.8 GB11.0 ms~5,830 tok/s
B = 256, GQA (K = 8)69.0 GB20.6 ms~12,430 tok/s

Three things fall out. The latency/throughput trade-off is arithmetic, not a tuning preference — B raises both, so you pick a point on a curve rather than escaping it. Throughput saturates: 64 → 256 costs 3× the latency for 32% more tokens per second, because by then the cache dominates the parameter read and B cancels itself out. And the wall you hit first is capacity — 240 GB does not fit an 80 GB card, so the interesting batch sizes are unreachable until the cache shrinks.

Two footnotes worth not skipping. Inference has an embarrassingly easy parallelism: launch M replicas and get M× throughput at identical latency with zero communication, since nothing needs updating (34:29) — sharding is for models that genuinely will not fit, not a default. And TTFT is a prefill quantity, so it wants the opposite batch policy from decoding: small batches for snappy first tokens, large ones for throughput after (35:36). That is why serving stacks schedule the two phases separately.

Shortcut one: make the cache smaller

"If that's one thing you take away from this lecture, it's all about the memory for speed."— Percy Liang, 38:32

The cache is B × S × L × K × H × 2 × 2 bytes. Every architectural trick in this section attacks one of those factors while trying not to lose accuracy.

Grouped-query attention attacks K (39:40). Multi-head attention gives every query head its own KV head; multi-query attention went to one shared KV head and lost expressivity; GQA sits between, with N query heads sharing K KV heads in groups and cutting the cache by N/K. The table above shows what a 1:5 ratio buys on the 13B config — cache down, latency down, throughput more than doubled, and the batch size that overflowed the card now fits, compounding into another 2× (42:24). Quality essentially holds, which is why Llama 3 adopted GQA across the board.

Multi-head latent attention, from DeepSeek-V2, attacks H instead: keep every head, but project each token's key/value into a shared low-dimensional latent and cache that — N·H = 16384 down to C = 512. The wrinkle Percy flags is that this is not compatible with RoPE, so 64 un-compressed dimensions ride along for rotary, 576 total (44:40). Note that he pulls up the wrong accuracy slide here and says he will dig it up later; the intended claim — MLA at least matches MHA while being far cheaper — is in the DeepSeek-V2 paper's Tables 8–9.

Cross-layer attention attacks L: GQA shares KV across heads, so share it across layers too (46:31). Perplexity ticks up, but less than the cache shrinks — a Pareto improvement. Local (sliding-window) attention attacks S and does something qualitatively different: once a token leaves the window you evict it, so the cache stops growing with sequence length at all — O(S) becomes O(1). It also removes exactly the capability attention was invented for, so pure local is not enough. The production answer is hybrid stacks: character.ai reports one global layer per six, with cross-layer sharing on top (49:24).

Carry this away: when you profile a serving setup and it looks slow, resist the urge to reach for kernels first. Compute the cache size and the arithmetic intensity for your actual context length and concurrency. If attention intensity is pinned near 1 — and in decoding it always is — no kernel will save you, and the question becomes which factor of B·S·L·K·H you are willing to trade accuracy for. That is an architecture decision, and it has to be made before you train.

Shortcut two: leave the Transformer

All of the above is still a Transformer being squeezed. Percy's more interesting claim is that attention-plus-autoregression is fundamentally memory-limited, because the architecture was designed for training efficiency and nobody was thinking about serving (53:34). Two escape routes get a fast tour.

State-space models came out of signal processing, aimed at long context rather than fast decoding — but a fixed-size recurrent state gives an O(1) cache for free. S4 won on long-range synthetics and disappointed on language, and the diagnosis is worth internalizing: the gap localizes to associative recall, where you must retrieve the value for a key seen arbitrarily far back (55:20). Logically trivial; exactly what a fixed state cannot do and full attention can. Hyena, H3, then Mamba (input-dependent SSM parameters) closed most of it and matched Transformers around 1B, and Jamba scaled the idea to a 52B MoE — while still interleaving some attention layers.

Linear attention is the currently-live version of the same bet. Drop the exponential kernel for a feature map: if exp(q·k) is replaced by φ(q)·φ(k), associativity lets you accumulate a running state instead of a growing set of keys, so the layer becomes an RNN, linear in sequence length, with a constant cache (57:39). BASED combines linear with local attention and studies the resulting recall-vs-cache-size frontier directly — which is the honest framing, since if you store less you will fail some retrieval tasks. MiniMax-01 pushed the recipe to a 456B-parameter MoE. Percy's read: linear plus local plus a sprinkling of full attention now produces serious models, so "is attention all you need?" gets a hedged yes-and-no — a little of it, in a few layers, apparently is.

Diffusion language models attack autoregression itself: generate the whole sequence in parallel from noise, then refine over a fixed number of steps (61:38). Every step is a dense parallel pass, so the thin-matrix problem never arises and utilization is high by construction. Inception Labs' demos post tokens-per-second far outside anything autoregressive, hybrid Mamba stacks included. Percy hedges — the coding numbers look good, generality is unproven, little is published — and argues from headroom: with that much of a speed lead you can spend compute buying accuracy back.

Shortcut three: cheaper numbers, smaller models

Quantization is the most direct expression of "memory is the currency": fewer bytes per parameter is fewer bytes moved, with accuracy as the price (65:07). fp32 is a training format and has no place in a serving path; bf16 is the inference default; fp8 (e4m3 on H100) and int8 halve it again; int4 halves it once more. Quantization-aware training means retraining, so post-training quantization — calibrate scale and zero-point on sample data — is the common path.

The failure mode is outliers. Absmax scaling divides by the largest magnitude in the tensor, so one pathological activation crushes everything else into a couple of levels — and such outliers appear reliably in larger networks. LLM.int8() splits the matmul: outlier dimensions in fp16, the vast majority in int8 (67:35). Percy notes the honest caveat — that paper's motivation was fitting models into memory at all, and the mixed path is actually slower than plain fp16. AWQ refines it: pick the small fraction of weights to protect in high precision based on the activations they see, and int3 becomes reachable at ~4× less memory and a ~3× speedup.

Pruning is the same move at coarser grain — rip out whole layers, heads or hidden dimensions, then repair. NVIDIA's recipe: score importance on ~1,024 calibration samples, delete the low scorers, distill the original into the pruned model (69:17). Distillation does the real work, and it is cheap precisely because the pruned model is not a fresh initialization but a damaged copy retaining the original's structure — 8B and 4B models out of a 15B one at up to 40× fewer tokens than training from scratch. Percy generalizes this into a recipe: define a faster architecture, initialize it from the slow model however you can, distill to repair. Every lossy trick above can be run that way.

Two lossless wins: speculative decoding and the scheduler

Everything so far is lossy, and you are left wondering what you gave up. Speculative decoding is the exception, and it falls straight out of the prefill/generation asymmetry: checking a proposed sequence is a parallel prefill (cheap), producing one is sequential decoding (expensive). So let a small draft model p run ahead K tokens, then score all K in one forward pass of the target q (72:13).

The elegance is in the acceptance rule. Accept draft token x with probability min(1, q(x)/p(x)) — an importance weight correcting the proposal back to the target, capped at 1. On rejection, do not retry: sample from the normalized residual max(q − p, 0). That one modification to textbook rejection sampling guarantees at least one token per round instead of an unbounded loop, and the output is a provably exact draw from q. Percy sketches the two-symbol case: if the draft oversamples A, then P[A] = p(A)·(q(A)/p(A)) = q(A), and B's mass accumulates from both branches to q(B). No approximation anywhere — the ~2× speedup is bought entirely with math, and two Google teams derived it independently within months.

Speed tracks the acceptance rate, so you want the draft as close to the target as possible: 70B with an 8B draft, or 8B with 1B, ideally distilled from the target. The draft is a wide-open design space where everything else in the lecture applies — Medusa adds heads so it emits several tokens per pass, EAGLE feeds the target's own hidden features in so it is not standalone at all (76:41).

The last ten minutes cover what only shows up once real users arrive (77:17). Training gets a dense rectangle of tokens; serving gets a ragged one — requests arrive at different times, share prefixes, finish at different lengths. Static batching wastes on both ends. Continuous batching (iteration-level scheduling, from Orca) returns control to the scheduler after every decode step, so finished sequences leave and new ones join immediately. Selective batching handles the ragged shapes by splitting the work: attention runs per sequence, since each owns its cache, while the MLPs — the bulk of the FLOPs — take every sequence concatenated along one flattened token axis, because they do not interact.

Then PagedAttention, the idea behind vLLM and a straight lift from operating systems. A contiguous cache slab per request wastes memory twice: internally, since you reserve for the maximum length and usually generate far less, and externally, as gaps between slabs. So page it — fixed blocks, placed anywhere, tracked by a block table. Fragmentation collapses, and prefix sharing becomes reference counting plus copy-on-write at block granularity: exactly what you want when a thousand requests carry the same system prompt.

Percy's closing move reframes the whole lecture. Optimizing inference for a fixed model is the narrow problem, and it is not the one that matters.

"Who cares about that particular model? You care about delivering good accuracy given your resource budget."— Percy Liang, 82:15

What you build with this

No assignment starts here. Lecture 10 sits inside the window of Assignment 3: Scaling (handout), which opened with lecture 9's scaling laws and continues through lecture 11; Assignment 4: Data begins at lecture 11. Percy opens by calling this a "brief respite from scaling laws," and it is a genuine detour in the course's spine.

What it does connect to is earlier work you have already done. The parameter and FLOP counting is assignment 1's accounting exercise, reused with memory traffic in the denominator instead of compute in the numerator — if you built that counter, extending it to emit a KV-cache size and an arithmetic intensity is a short afternoon. The benchmarking and profiling discipline from Assignment 2: Systems (handout) is what you need to check any of these predictions against a real GPU, and the FlashAttention work there is assumed by the attention accounting in this lecture. Looking forward, the material returns in the alignment lectures, where RL post-training is bounded by rollout generation — that is inference, in the training loop, at scale.

Supporting materials, verified

Exercises

  1. Rebuild the napkin model code — Reproduce the lecture's table from scratch, then extend it. (1) Write a function taking D, F, L, N, K, H, V, S, B and bandwidth, returning parameter bytes, KV-cache bytes per sequence, per-token latency and throughput. (2) Check it against the lecture: Llama 2 13B at B=1 should give ~26.9 GB, ~8 ms, ~125 tok/s. (3) Add a memory-capacity constraint and solve for the largest B that fits 80 GB, as a function of K. (4) Plot throughput against B for K ∈ {40, 8, 1} and mark where each curve is cut off by capacity. A good answer shows the throughput knee, explains it algebraically (B appearing in both numerator and denominator once the cache dominates), and observes that GQA's real win is often the batch size it unlocks rather than the direct traffic saving.
  2. Measure the intensity you actually get code — Predictions are worth what you can verify. (1) Take any HuggingFace model that fits your GPU and time single-token decode at B ∈ {1, 4, 16, 64} with a fixed 1k-token prefix. (2) Compute achieved FLOP/s and achieved bytes/s and locate each point relative to your card's roofline crossover. (3) Repeat at context lengths 512, 2048, 8192 and watch the parameter-read term give way to the cache term. (4) Separately time prefill and confirm it lands on the compute-limited side. A good answer reports where measurement diverges from the model and names the cause — kernel-launch overhead at small B, imperfect overlap, or the non-matmul work the lecture drops.
  3. Implement speculative decoding and prove it is exact code — (1) Pick a target/draft pair from one family (e.g. a 1B draft against a 7B target). (2) Implement the loop: draft K tokens, score them in one target pass, accept with min(1, q/p), on rejection sample from normalized max(q−p, 0), and always emit at least one token. (3) Empirically verify exactness on a tiny vocabulary — build a toy p and q over 5 symbols, draw 10⁶ samples through your procedure, and check the histogram matches q within sampling error. (4) Sweep K ∈ {1..8} and report acceptance rate and wall-clock speedup. A good answer shows the histogram test passing and explains why the residual-sampling step, not the accept step, is what makes the distribution exact.
  4. Design a serving architecture under a stated budget — No code. You must serve a 30B-class model at 100 concurrent users, 32k context, with TTFT under 500 ms, on 8×H100. Using only this lecture's arithmetic, write a one-page proposal: how many replicas versus how much sharding; which KV-cache reduction (GQA ratio, MLA, hybrid local/global) and what it costs you; whether to quantize and to what; whether speculative decoding is worth the extra weights in memory. A good answer computes the cache size at 32k context first — that number decides most of the rest — states the accuracy risk of each lossy choice explicitly, and identifies which single assumption, if wrong, breaks the whole plan.
Next: L11 Scaling laws 2 · Back to the map.