/ CS336
[CS336-10] Inference
Notes for Stanford CS336 Spring 2025 lecture 10: Inference, KV Cache.
这篇是 Stanford CS336 Spring 2025 第 10 讲 Course Materials 的学习笔记。主题是 LLM inference。
1. Prefill 和 Decode 是两种 workload#
推理不是单一过程。一次请求通常分成 prefill 和 decode。
Prefill 阶段处理 prompt 中已有 token,可以并行计算整段上下文,形态接近训练 forward,通常更 compute-bound。
Decode 阶段每次只生成一个新 token,需要不断读取历史 KV cache,batch 和 sequence 动态变化,通常更 memory-bound。
这就是为什么训练优化和推理优化关注点不同。训练更关心大 batch matmul 吞吐,推理更关心延迟、KV cache、动态 batching 和 memory bandwidth。
2. 降低推理成本的两条路线#
第一条是 lossy shortcut:允许模型或表示发生变化,只要质量损失可接受。例如 quantization、pruning、distillation、使用更小模型、减少 KV heads、局部 attention 等。
第二条是 lossless shortcut:保持目标模型采样分布不变,但改变执行方式。例如 speculative decoding 使用 draft model 生成候选 token,再由大模型并行验证;如果设计正确,可以保持 exact sampling。
系统工程里经常需要同时使用两类方法:一个负责降低单 token 成本,一个负责提高整体吞吐。
3. KV Cache 是 decode 的核心瓶颈#
在自回归生成中,每一层 attention 都需要历史 K/V。缓存它们可以避免重复计算 prefill 结果,但也带来显存压力。
KV cache 大小随 batch size、sequence length、layers、KV heads、head dimension、dtype 增长。长上下文和高并发时,KV cache 可能比模型权重更像瓶颈。
因此 GQA、MLA、CLA、local attention 等方法都可以从减少或压缩 KV cache 的角度理解。
4. GQA、MLA、CLA 和 Local Attention#
Grouped-query attention 让多个 query heads 共享更少的 key/value heads,减少 KV cache。
Multi-head latent attention 进一步把 key/value 投影到更低维 latent 表示,希望在减少缓存的同时尽量保持质量。
Cross-layer attention 让多层共享或复用部分 attention 信息,减少每层独立存储和计算。
Local attention 限制 token 只看局部窗口,降低长上下文成本,但需要搭配少量 global 层或其他机制保留远距离信息。
这些方法的共同目标是:在尽量不伤害质量的前提下降低 decode memory pressure。
5. 动态 workload 和 PagedAttention#
在线推理请求长度不同、到达时间不同、生成长度不同。如果直接为每个请求分配连续大块 KV memory,会产生碎片和浪费。
PagedAttention 借鉴操作系统分页思想,把 KV cache 分成 blocks,通过 block table 管理逻辑序列和物理缓存。这样可以更好地支持动态 batching、request arrival/departure 和 memory reuse。
这说明推理系统不是单纯模型问题,也很像操作系统和内存管理问题。
6. takeaway#
第 10 讲的核心是:LLM inference 的瓶颈和训练不一样。
- prefill 更像训练,decode 更像 memory-bound 的逐 token 服务;
- KV cache 是推理系统的中心资源;
- GQA、MLA、local attention 等方法都在压缩或减少缓存;
- speculative decoding 用小模型帮助大模型加速,但保持输出分布;
- dynamic batching 和 paged KV cache 是服务化推理的关键。
如果训练系统关心“多久训完”,推理系统关心的是“每个请求多快、多稳、多便宜”。
参考#
- Stanford CS336 Spring 2025 Course Materials: lecture_10.py