第 10 讲:推理(Inference)

目录 · ← l9 · l11 →

第 10 讲:推理(Inference)

日期:4 月 29 日(周三,Spring 2026) | 讲师:Percy Liang | 材料:lecture_10.py | 截止:作业 2 到期、作业 3 发布

概览

推理才是模型真正被使用的地方,而它的特性与训练截然不同:memory-bound(访存受限)动态(dynamic)(请求到达和结束的时间各不相同)。本讲推导 prefill 与 generation 两个阶段中 MLP 与 attention 层的算术强度,计算 Llama 2 13B 在 H100 上的理论延迟与吞吐,然后系统梳理加速手段:削减 KV cache(GQA、MLA、CLA、局部/滑窗 attention、DeepSeek v4 的 CSA/DSA/HCA)、量化(QAT/PTQ、AWQ)、剪枝 + 蒸馏、投机采样(无损!),以及面向动态负载的系统优化(continuous batching、PagedAttention)。

核心概念与定义

  • 为什么推理效率重要:训练是一次性成本,推理会被重复无数次(OpenAI 每天处理约 8.6T token)。Agent 让它更严重:内部轨迹可以无限增长。生成的 token 数 = 花掉的算力
  • 指标TTFT(time-to-first-token,首 token 延迟,由 prefill 决定)、latency(单条查询的秒/token,面向交互)、throughput(多查询的 token/秒,面向批处理)。
  • 两个阶段prefill(把整个 prompt 并行处理,像训练一样——compute-bound)与 generation/decode(一次一个 token——memory-bound)。关键不对称性:检查比生成快(prefill 能一次算完所有位置)。
  • KV cache:避免在每个生成步为整段历史重算 key/value。对每个序列(B)、token(S)、层(L)、头(K),存一个 H 维向量。朴素推理生成 T 个 token 需 O(T³) FLOPs;有 KV cache 后降为 O(T²)。
  • 算术强度核算(bf16,每值 2 字节):
    • MLP 每步:FLOPs = 6·B·T·D·F;字节 = 4·B·T·D + 4·B·T·F + 6·D·F → 强度 ≈ B·T。Prefill(B·T 大)compute-bound;generation(T=1)强度 ≈ B——需要很多并发请求(batching)才能维持 compute-bound。
    • Attention 每步:FLOPs = 4·B·S·T·D;字节 = 4·B·S·D + 4·B·T·D → 强度 = S·T/(S+T)。Prefill(T=S)为 S/2(不错);generation(T=1)< 1——靠 batching 也无法改善,因为每个序列有自己的 KV cache(Q、K、V 都依赖 B),而 MLP 权重是共享的。
    • 总结:prefill 是 compute-bound,generation 是 memory-bound(每步都要读全部参数 + KV cache)。
  • 延迟/吞吐模型:latency ≈ memory / bandwidth(读全部参数 + KV cache),throughput = B / latency。batch 越大:延迟越差(要读写的 KV cache 更大)、吞吐越好(分摊参数读取成本)——这是根本性权衡。此外:复制 M 份模型可让吞吐线性提升;TTFT 是 prefill 现象(TTFT 用小 batch,生成吞吐用大 batch)。
  • 削减 KV cache(memory-bound ⇒ cache 更小 ⇒ 更快):
    • GQA(grouped-query attention):N 个 query 头,但只有 K 个 key/value 头(K < N);MHA 为 K=N,MQA 为 K=1。把 KV cache 缩小 N/K 倍且几乎不损精度(Ainslie 2023)。Llama 2 13B 从 K:40 改到 K:8:单批延迟变差,但吞吐提升且能装进内存。
    • MLA(multi-head latent attention,DeepSeek v2):不存 K、V,而存压缩隐向量 c_t = W_c h_t(C 维);需要时上投影为 K = W_K c、V = W_V c。DeepSeek v2:N·H = 16384 → C = 512(+ 64 维 RoPE,共 576)。MLA 在更低成本下甚至略优于 MHA。与 RoPE 不兼容:需额外保留非旋转的 key 维度。
    • CLA(cross-layer attention):跨共享 KV(正如 GQA 跨头共享);改善”精度 vs KV 大小”的帕累托前沿。
    • 局部(滑窗)attention:只关注一个窗口(Longformer、Mistral);KV cache 与序列长度无关;有效上下文随层数线性增长;有时损精度 → 把局部与全局 attention 交错(混合层)。
    • DeepSeek v4 attention(1M 上下文):Compressed Sparse Attention(CSA,每 m 个 token 压成 1 个)、DeepSeek Sparse Attention(DSA,选 top-k)、Heavily Compressed Attention(HCA)。
    • 其他:线性注意力 / 状态空间模型(Mamba-2、GatedDeltaNet)、扩散语言模型。
  • 量化(quantization):比特更少 = 字节更少 = 更快(memory-bound)。fp32(训练)→ bf16(推理默认)→ fp8/int8 → int4/nvfp4。QAT(训练中量化,昂贵);PTQ(训练后量化,便宜:在样本数据上校准 scale/zero-point;GPTQ 用 Hessian 信息修正);AWQ(激活感知:依据激活幅度把 0.1–1% 的重要权重保留高精度;fp16→int3 得到 4 倍显存下降、3.2 倍加速)。
  • 剪枝 + 蒸馏:(1) 在约 1024 条校准样本上识别重要的 {层, head, 隐维度};(2) 去掉不重要的部分得到更小模型;(3) 用原模型向剪枝模型蒸馏(NVIDIA 的 pruning-KD 循环)。
  • 投机采样(speculative sampling,无损):廉价草稿模型 p 提议 γ 个 token;目标模型 q 并行(prefill 速度)为它们打分;按修正的拒绝采样接受/拒绝。两个关键性质:(1) 至少生成一个 token(否则拒绝采样会无限循环);(2) 保证是 q 的精确采样——用两个符号 {A,B} 举例证明:若 p(A)>q(A),残差 max(q−p,0) 修正后 P[采到 A]=q(A)、P[采到 B]=q(B)。扩展:Medusa(并行多头)、EAGLE(用目标模型特征做草稿)。
    • 类比:让一位手快的实习生先起草一段话;教授只读一遍(很快——检查比写作快),然后决定批准或修改;最终文本在统计上与教授独自写作完全相同。
  • Continuous batching(Orca):迭代级调度——新请求随到随加入批次,而不是等静态 batch 里所有序列都结束。Selective batching:attention 逐序列单独处理(长度参差),非 attention 运算把所有序列拼成一个 [Σs, H] 张量。
  • PagedAttention(vLLM):给 KV cache 做操作系统的”虚拟内存分页”——把每个序列的 KV 切成不连续的固定大小 block;消除内部/外部碎片;支持跨序列共享前缀 block(系统提示、同一 prompt 多采样)并用写时复制(copy-on-write)。vLLM 的其他优化:融合 block-attention kernel、FlashAttention/FlashDecoding、CUDA graphs。
    • 类比:静态 KV 分配像给每个进程预留”按最坏情况算”的连续内存(造成碎片);分页就是虚拟内存——按需映射 block、共享只读页、写时复制。

代码示例:attention 的算术强度(关键推导)

代码(Python):

# B 批、S 个历史 token、T 个待生成 token、D 模型维度;bf16 => 每值 2 字节
flops = 4*B*S*T*D                       # QK^T: 2*B*S*T*D  +  softmax@V: 2*B*S*T*D
bytes = 4*B*S*D + 4*B*T*D               # 读 Q,K,V;写 Y

intensity = (S*T) / (S + T)             # 化简后的算术强度

# Prefill:T = S  =>  强度 = S/2        (不错——compute-bound)
prefill_intensity = S / 2

# Generation:T = 1  =>  强度 = S / (S + 1) < 1   (很差——memory-bound)
generate_intensity = S / (S + 1)

代码做了什么: 统计 attention 矩阵乘的 FLOPs 与 HBM 字节数,得出算术强度为 S·T/(S+T)——prefill 时是 S/2,generation 时小于 1——并且与 B 无关

实现深挖:

  • 为什么与 B 无关:attention 中的 Q/K/V 都是逐序列的(B 同时放大分子与分母,相互抵消);而 MLP 的权重在 batch 间共享,所以 B 会提升强度。这正是”batching 救不了生成阶段的 attention”的原因——KV cache 是每个序列独有的。
  • 为什么强度 <1 是致命的:H100 的加速器强度约 295 FLOP/byte;生成阶段 attention 约每 1 个 FLOP 就要搬 1 字节——tensor core 几乎完全闲置。
  • 为什么这解释了 GQA/MLA 的必要性:每减少一个字节的 KV cache,都会直接降低生成延迟(latency ∝ 每步读取的显存)。

与作业的联系:作业 1 的资源核算要求给出每个组件的 FLOP/字节数;作业 2 的 FlashAttention 与 benchmark 工作针对的是同一组公式的训练侧。讲义的理论延迟/吞吐模型(compute_transformer_performance_statsnum_params = 2VD + 3LDF + 2L·(2DNH + 2DKH)kv_cache_size = 4·S·K·H·L 字节)正是你用来 sanity check 作业 2 实测数字的模板。

代码示例:投机采样(无损解码)

代码(Python):

def speculative_sample(draft_logits, target_logits, draft_next, rng):
    # draft_logits / target_logits: 下一个位置的 [vocab] 分布
    p = softmax(draft_logits)             # 草稿分布
    q = softmax(target_logits)            # 目标分布
    x = draft_next                        # 从草稿中抽出的候选 token
    u = rng.uniform(0, 1)
    if u < min(1, q[x] / p[x]):           # 以 q(x)/p(x) 的概率接受
        return x, True
    # 拒绝:从残差分布(归一化后)重采样
    residual = torch.clamp(q - p, min=0)
    x2 = sample(residual / residual.sum())
    return x2, False

代码做了什么: 从草稿模型抽取候选;以概率 q(x)/p(x) 接受(重要性加权);被拒时从归一化的残差 max(q−p, 0) 重采样。这是”被改造过的拒绝采样”,保证总能给出一个来自 q 的有效样本。

实现深挖:

  • 为什么它是精确的:接受(概率 min(1, q/p))与残差重采样的混合恰好复现 q。讲义的双符号证明:P[采到 A] = p(A)·(q(A)/p(A)) + p(B)·1·0 = q(A);P[采到 B] = p(B)·1 + p(A)·(1−q(A)/p(A))·1 = q(B)。
  • 为什么它快:草稿模型(如 8B)以 memory-bound 速度生成 γ 个 token;目标模型(如 70B)并行为这 γ 个 token 打分(prefill 式,compute-bound)——利用的正是”检查与生成之间的不对称性”。
  • 为什么”至少生成一个”:朴素拒绝采样可能永远拒绝;该修正保证前进,同时用残差修正保持分布精确。
  • 如何让草稿更好:向目标模型蒸馏草稿(提高接受率)、Medusa(并行草稿头)、EAGLE(用目标模型特征条件化草稿)。

与作业的联系:必修作业不实现它,但作业 5 的 RL 训练循环用的是同一套”通过高速推理服务器(vLLM)做 rollout”的机制——cs336_alignment/vllm_utils.py 中的 vLLM 接口正是本讲推理栈的工程落地。

代码示例:延迟/吞吐模型(Llama 2 13B on H100)

代码(Python):

def compute_transformer_performance_stats(config):
    # 参数量(embedding + 3 个 MLP 矩阵 + 每层 attention 的 QKV/O)
    num_params = 2*V*D + D*F*3*L + (2*D*N*H + 2*D*K*H)*L
    parameter_size = 2 * num_params                       # bf16

    # 每序列的 KV cache:S 个 token * K 头 * H 维 * L 层 * (K+V) * 2 字节
    kv_cache_size_per_seq = S * (K*H) * L * 2 * 2

    memory = B * kv_cache_size_per_seq + parameter_size   # 每步要读的总字节

    latency = memory / memory_bandwidth                   # 秒/token
    throughput = B / latency                              # token/秒
    return num_params, memory, latency, throughput

# Llama 2 13B 配置:S=1024, D=5120, F=13824, N=40, K=40, H=128, L=40, V=32000
# B=1:  latency ~ (26GB 参数 + 很小的 cache) / 3.35TB/s  (约 7.8ms/token)
# B=64:吞吐更好、延迟更差(KV cache 更大,要读更多)
# B=256:吞吐收益递减,而且装不进 80GB 的 H100!

代码做了什么: 计算参数量、显存(参数 + KV cache),以及在”显存带宽受限”下的延迟/吞吐,并代入 Llama 2 13B 在 batch 为 1/64/256 时的情形——展示延迟-吞吐权衡与显存上限。

实现深挖:

  • 为什么 latency = memory/bandwidth:生成是 memory-bound;每一步都必须从 HBM 读取全部参数加上整个 KV cache。这是理论下界(假设完美重叠)——真实系统只会更差。
  • 为什么 throughput = B/latency:每步并行生成 B 条序列,故 token/秒 = B × 每秒步数。
  • 为什么削减 KV cache 一举两得:显存更小 → 每步读取时间更短(延迟降低),同时能塞进更大的 batch(吞吐提升)。

与作业的联系TransformerPerformanceStats 这套符号化核算(讲义用 sympy)是可复用的模板,可用于作业 1 的资源核算,以及核对作业 2 的 profiling 数字(例如”为什么生成是 memory-bound?”——你的 attention kernel benchmark 应当反映这一点)。

关键要点

  1. 推理分两种:prefill(compute-bound,像训练)与 generation(memory-bound,一次一个 token)——而生成阶段的 attention 是batching 无法解决的 memory-bound(强度 <1,与 B 无关)。
  2. KV cache 是核心资源:latency ∝ 每步读取的(参数 + KV cache);用 GQA、MLA、CLA、局部 attention 或状态空间混合模型削减它。
  3. 存在无损加速:投机采样可证明是精确的(修正拒绝采样),利用”检查比生成快”。
  4. 有损加速:量化(QAT/PTQ/AWQ)与剪枝 + 蒸馏,用精度换显存与速度。
  5. 动态负载需要系统技巧:continuous batching(迭代级调度、selective batching)与 PagedAttention(分页、前缀共享、写时复制)——思路直接借自操作系统。

常见陷阱

  • 显存预算里忽略 KV cache:长上下文 + 大 batch 时,OOM 的往往是 KV cache(而不是参数),而且它随并发数增长。
  • memory-bound 时还用 MHA:高 batch 下 GQA/MLA 几乎是白送的收益;MHA 的精度优势很小。
  • 朴素量化:逐张量 scale 会丢掉离群通道的信息;应使用分块 scale(AWQ、MXFP8)并在真实激活范围上校准。
  • 以为投机解码会改变分布:它必须精确——若实现改动了采样,就不再是”目标模型”的输出。
  • 静态 batching:等所有请求结束会浪费 GPU(一条慢序列拖住所有人);要用 continuous batching。
  • KV 分配碎片化:为每个请求预留最大长度会造成内部/外部碎片和 HBM 浪费——要分页。
  • 忘记 TTFT 与吞吐的区分:为交互流量优化吞吐会伤害用户可见延迟;分阶段调 batch(prefill 小、generation 大)。

复习题

  1. 问: 为什么 batching 无法解决生成阶段 attention 的 memory-bound 问题?
    • 答: attention 的算术强度 S·T/(S+T) 中没有 B 项:Q、K、V 都是逐序列的,B 同时放大 FLOPs 与字节数并相互抵消。MLP 不同——权重在 B 间共享,因此 batching 把强度提升到约 B·T。生成阶段 attention 无论 batch 多大都保持 <1 的强度。
  2. 问: 投机采样如何保证样本精确来自目标模型 q?
    • 答: 它是被改造过的拒绝采样:以 min(1, q(x)/p(x)) 的概率接受草稿 x;被拒时从归一化的残差 max(q−p, 0) 重采样。混合代数(双符号情形已给出证明)恰好复现 q,而”至少接受一个 token”的修正保证过程终止。
  3. 问: Llama 2 13B 把 K 从 40 个头降到 8 个头(GQA)——什么变了、什么没变?
    • 答: KV cache 缩小 5 倍(每步延迟下降、可容纳更大 batch、吞吐上升)。query 头仍为 40(attention 的表达力基本保持);据 Ainslie 等 2023,精度损失很小甚至可忽略。