Lesson 6 of 7 · 49 min
Attention & serving at scale
The production techniques that make transformers servable: MQA/GQA (cut KV bytes at the source), FlashAttention 1/2/3 (IO-aware exact attention), sliding-window/sparse attention, PagedAttention + continuous batching, and quantization. Real measured impact, the tradeoffs, and how vLLM/TensorRT-LLM/SGLang differ.
Where the textbook transformer meets the GPU budget
vLLM: Easy, Fast, and Cheap LLM Serving (PagedAttention)Woosuk Kwon, UC BerkeleyMQA & GQA — cut the cache at the source
H query heads but shares a single K/V head — so the cache shrinks by H×. Grouped-Query Attention (GQA) interpolates: G groups, each group of H/G query heads sharing one K/V head (MQA is G=1, MHA is G=H). Llama 2 70B and Mistral 7B both use 8 KV heads, cutting cache ~4–8× vs MHA at near-MHA quality. GQA’s killer trick: you can uptrain an existing MHA checkpoint into GQA by mean-pooling the per-group K/V weights and continuing for only ~5% of original pre-training compute.1MHA -> MQA -> GQA (the KV-head lever; cache scales with H_kv)23 variant query heads KV heads cache vs MHA quality used in4 ------- ----------- -------- ------------ ----------- -------------------5 MHA H H 1x (baseline) baseline orig Transformer, BERT6 MQA H 1 1/H slight drop PaLM, Falcon7 GQA H G (e.g 8) G/H ~= MHA Llama 2 70B, Mistral8 MLA latent compressed lowest ~= MHA DeepSeek-V2910 GQA-8 on a 64-query-head model = 8x smaller KV cache at near-MHA quality.11 Uptrain MHA -> GQA for ~5% of pretraining compute. Decide BEFORE fine-tuning.FlashAttention — IO-awareness, not smarts
T×T score matrix in HBM, applies softmax, reads it back. FlashAttention tiles Q/K/V into blocks that fit in fast SRAM, computes softmax incrementally (online softmax with rescaling), and never writes the full T×T matrix to HBM. The output is bit-for-bit the standard result. It reduces HBM I/O from O(T²·d) to O(T²·d²/M) where M is SRAM size — and on the backward pass it recomputes the matrix from cheap stored stats instead of storing it, trading FLOPs for memory.1FLASHATTENTION FAMILY (exact attention; the win is HBM I/O, not approximation)23 version key idea measured impact4 ------- ----------------------------------------- ----------------------------------5 FA-1 tile Q,K,V in SRAM; never write T x T HBM 3x on GPT-2 (1k); 2.4x LRA (1k-4k)6 FA-2 better work partitioning across warps 2x over FA-1; 50-73% of A100 peak7 FA-3 warp specialization + TMA + FP8 (Hopper) 740 TFLOPS FP16; ~1.2 PFLOPS FP889 A100: HBM ~2 TB/s vs SRAM ~19 TB/s -> moving the T x T matrix IS the bottleneck.10 Treat naive 3-loop attention as a bug, not a baseline.Sliding-window & sparse attention
O(T²) attention, restrict which pairs attend. Sliding-window attention (SWA) lets each token attend only to its W nearest neighbors — linear, not quadratic. The clever part is stacking: Mistral 7B uses W=4096 across 32 layers, so through layered composition a token at the top can reach back ~32×4096 ≈ 131k tokens of effective context, while a rotating buffer halves cache memory above 8k. Longformer adds a few global tokens that attend to everything (carrying long-range state); BigBird combines window + global + random for linear complexity (~8× less time/memory than vanilla at 4k). Mistral 7B is competitive with Llama 2 13B across benchmarks on this design.PagedAttention & continuous batching — the serving breakthrough
max_seq_len per request, leaving most slots empty — typical fragmentation waste is 60–80%. PagedAttention (vLLM) borrows OS virtual-memory paging: store K/V in fixed-size physical blocks (e.g. 16 tokens) and keep a per-request block table mapping logical → physical blocks. Waste drops to <4%, and prefixes can be shared safely across requests (system prompts, parallel samples). The measured win: 2–4× throughput over FasterTransformer/Orca at matched latency, and up to 24× over HuggingFace Transformers on high-batch offline inference.Key idea
Quantization — and how it interacts with internals
1QUANTIZATION: PICK BY HARDWARE + WORKLOAD23 method precision helps cost best on4 ----------- --------- ---------- ------------ ----------------------5 GPTQ W3/W4 memory visible @3-bit weight-only baseline6 AWQ W4 memory small single-GPU budget infer7 SmoothQuant W8A8 compute+mem small Ampere (INT8 GEMM)8 FP8 (TE) W/A/KV FP8 both small Hopper/Blackwell only910 W4 wins MEMORY but barely speeds compute (activations still FP16).11 W8A8 wins SPEED on Ampere. FP8 wins BOTH on Hopper -- and only there.12 KV-cache quant (INT8/FP8) ~2x cache, but watch perplexity drift past ~64k.The engines: vLLM, TensorRT-LLM, SGLang
1H100, Llama-3.3-70B FP8, concurrency 100 (Spheron benchmark, Mar 2026)23 engine throughput (tok/s) p95 TTFT (ms) pick when4 ------------- ------------------ ------------- ------------------------------5 TensorRT-LLM 2,780 1,280 max throughput on NVIDIA fleet6 SGLang 2,460 1,380 prefix-heavy RAG / agents7 vLLM 2,400 1,450 broadest support, fastest iter89 TensorRT-LLM ~8-16% ahead on throughput, ~12-15% on TTFT -- for a heavier deploy.10 SGLang's real edge (prefix reuse) does NOT show in this raw-throughput table.Common mistake
“vLLM is 4× faster than the baseline because of PagedAttention.”
Interview prep
- 01“MQA vs GQA vs MHA?” → KV-head lever: MHA H heads, MQA 1, GQA G; GQA-8 ≈ MHA quality at ~MQA speed; uptrain MHA→GQA for ~5% compute.
- 02“How does FlashAttention save memory?” → tile Q/K/V in SRAM, never materialize the T×T matrix in HBM, recompute in backward; exact, not approximate; 2–3×.
- 03“Sliding-window vs full attention?” → stacked SWA (Mistral W=4096×32 layers ≈ 131k reach) is linear and production-friendly; isolated windows lose retrieval.
- 04“What is PagedAttention and why does vLLM use it?” → OS-style paging of the KV cache → waste 60–80%→<4%, safe prefix sharing, 2–4× throughput.
- 05“PagedAttention vs continuous batching?” → different levers; continuous (iteration-level) batching alone was Orca’s 36.9×; they compose.
- 06“W4 vs W8A8 vs FP8?” → W4 saves memory only; W8A8 speeds compute on Ampere; FP8 does both on Hopper — choose by hardware, ablate at longest context.
- 07“Which serving engine?” → vLLM (general/iteration speed), TensorRT-LLM (max throughput on NVIDIA, heavy compile), SGLang (prefix-heavy RAG via RadixAttention).
- 08“KV cache quantization risks?” → ~2× cache saving but per-tensor scaling needed; RoPE+naive INT8 K is unstable; FP8 K with per-token scaling, watch >64k drift.
Common mistake
The red-flag answer: “to scale, I’d use a managed endpoint and a bigger GPU.”
Checkpoint
A 70B MHA model is too KV-cache-heavy to serve at your target context and concurrency. You have ~5% of original pre-training compute available. Best move?
Checkpoint
At 8k+ context your attention is the latency bottleneck on an A100, and you’re running a hand-written three-loop attention kernel. Highest-impact change, and why it’s safe?
Checkpoint
You deploy a single-layer sliding-window model and users report it can answer about recent text but fails whenever the needed fact is far earlier in a long document. Best explanation?
Checkpoint
Your serving GPU has plenty of KV-cache memory free, yet throughput is poor and many requests sit idle waiting for a batch to finish. Which lever helps most?
Checkpoint
On H100s you want both lower memory and higher matmul throughput for a Llama-class model, with minimal accuracy work. Best quantization choice?
Could you choose GQA vs MQA, justify FlashAttention by IO, design a PagedAttention + continuous-batching stack, and pick quantization by hardware?
Takeaways
- KV-head count is the top serving lever: GQA-8 ≈ MHA quality at ~MQA speed; uptrain MHA→GQA for ~5% compute, before fine-tuning.
- FlashAttention is exact, not approximate — tile in SRAM, never materialize T×T in HBM; 2–3× (FA-1), 50–73% A100 peak (FA-2), FP8 on Hopper (FA-3).
- Stacked sliding-window (Mistral W=4096×32 ≈ 131k reach) is linear; isolated windows lose retrieval — combine with RAG for million-token contexts.
- PagedAttention (waste 60–80%→<4%) + continuous batching (Orca 36.9×) compose — don’t credit one for the combined lift.
- Quantization is hardware-bound: W4 saves memory, W8A8 speeds Ampere, FP8 does both on Hopper; always ablate at your longest context.
Finally: the capstone — annotate a forward pass end to end, then size a real serving deployment.
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.