投机解码深度实战:从草稿-验证采样到 EAGLE 树状草稿,把 LLM 解码延迟砍一半
如果你做过 LLM 线上服务调优,一定撞过这堵墙:首 token 出来之后,后面每个 token 都要等,而 GPU 大部分时间在空转。
这不是玄学。自回归解码每生成一个 token 就要把整个模型权重从 HBM 搬进 SM 一次,算算术强度(arithmetic intensity)低得可怜。以 70B 模型 FP8 推理为例,权重约 70GB,H100 的 HBM 带宽 3.35TB/s,单步解码的理论下限就是 70GB / 3.35TB/s ≈ 21ms,也就是单用户上限约 47 token/s——不管你有多少 FLOPS 都没用,这是内存墙。
增大 batch 能把权重搬运摊薄,吞吐上去了,但每个请求的延迟也上去了,而且高 batch 下 GPU 已经饱和,再优化收益递减。要在低延迟(小 batch)场景下提速,只剩一条路:让一次前向解码出多个 token。
这就是投机解码(Speculative Decoding)。
一、第一性原理:草稿 + 并行验证
思路非常朴素:
- 用一个又快又小的草稿模型(draft model)自回归生成 γ 个候选 token。它慢不到哪去,因为小模型一次前向也就几毫秒。
- 把 prefix + γ 个候选一并喂给目标模型(target model),一次前向拿到 γ+1 个位置的真实概率分布。注意这一步是并行的,代价约等于单次解码。
- 逐位置做接受-拒绝采样:草稿和目标分布越接近,接受率越高。
- 若全部接受,还能从目标模型第 γ+1 个位置免费多采样一个 token(bonus token)。
关键在于:这套算法在数学上保证输出分布与目标模型逐 token 采样严格同分布。也就是说它不是近似加速,是无损加速——这一点是它区别于量化、剪枝、稀疏化的根本优势,也是它能被 vLLM、SGLang、TensorRT-LLM 全线支持的原因。
接受率的数学结论很干净。设 α 为草稿与目标分布的平均不匹配率,则单轮期望接受 token 数为:
E[accepted] = (1 - α^(γ+1)) / (1 - α)
当 γ=5、α=0.4 时约 1.6 个;α 降到 0.2 时约 3.2 个。所以整个游戏的核心不是把 γ 调大,而是把 α 做小——即让草稿模型尽可能"像"目标模型。
二、核心算法:30 行讲透接受-拒绝采样
下面是一个可直接跑的最小实现(来自 Leviathan et al. 的经典算法,我做了工程化改写):
import torch
from torch import Tensor
@torch.no_grad()
def speculative_sample(
prefix: Tensor, # (1, T) 已生成 token
draft_model, target_model,
gamma: int = 5,
max_new_tokens: int = 256,
temperature: float = 0.0,
):
x = prefix
n = 0
stats = {"draft": 0, "accepted": 0, "rounds": 0}
while n < max_new_tokens:
# ---------- 1. 草稿阶段:小模型自回归 gamma 步 ----------
draft_ids, draft_probs = [], []
cur = x
for _ in range(gamma):
logits = draft_model(cur).logits[:, -1, :]
p = logits.softmax(-1) if temperature == 0 else (logits / temperature).softmax(-1)
nxt = p.argmax(-1, keepdim=True) if temperature == 0 else torch.multinomial(p, 1)
draft_ids.append(nxt)
draft_probs.append(p)
cur = torch.cat([cur, nxt], dim=-1)
draft_ids = torch.cat(draft_ids, dim=-1) # (1, gamma)
draft_probs = torch.stack(draft_probs, dim=1) # (1, gamma, V)
stats["draft"] += gamma
# ---------- 2. 验证阶段:目标模型单次前向 ----------
tgt = target_model(torch.cat([x, draft_ids], dim=-1)).logits
tgt = tgt[:, -(gamma + 1):, :] # (1, gamma+1, V)
q = tgt.softmax(-1) if temperature == 0 else (tgt / temperature).softmax(-1)
# ---------- 3. 接受-拒绝采样 ----------
accepted = 0
rejected = False
for i in range(gamma):
xi = draft_ids[:, i]
p_i = draft_probs[:, i, :]
q_i = q[:, i, :]
ratio = q_i.gather(-1, xi[:, None]).squeeze(-1) / p_i.gather(-1, xi[:, None]).squeeze(-1).clamp_min(1e-12)
if torch.rand((), device=x.device) < torch.clamp(ratio, max=1.0):
accepted += 1
else:
# 关键:拒绝后不是丢弃,而是从残差分布 (q - p)+ 重采样
resid = (q_i - p_i).clamp_min(0.0)
resid = resid / resid.sum(-1, keepdim=True)
draft_ids[:, i] = torch.multinomial(resid, 1).squeeze(-1)
accepted += 1
rejected = True
break
keep = draft_ids[:, :accepted]
x = torch.cat([x, keep], dim=-1)
n += accepted
stats["accepted"] += accepted
stats["rounds"] += 1
if not rejected:
# 全部接受 → 从目标分布的最后一个位置补采一个 bonus token
bonus = q[:, gamma, :]
nxt = bonus.argmax(-1, keepdim=True) if temperature == 0 else torch.multinomial(bonus, 1)
x = torch.cat([x, nxt], dim=-1)
n += 1
stats["accepted"] += 1
stats["accept_rate"] = stats["accepted"] / max(stats["draft"], 1)
return x, stats
三个容易被写错的地方,值得单独说:
- 残差重采样不能省。很多手写实现拒绝后直接 fallback 成"用目标模型重新采样",这在数学上就偏了,破坏了无损性质。必须是
(q - p)⁺ / Σ(q - p)⁺。 - 验证时要算 γ+1 个位置。第 γ+1 个位置是 bonus token 的来源,少了它等于白扔一步。
temperature=0(贪心)时行为退化。贪心下接受条件变成"argmax 一致",接受率通常更高,投机解码在低温度场景收益最明显——这解释了为什么它在代码补全、结构化输出、Agent 工具调用(普遍 temperature≈0)上特别有效。
三、草稿从哪来:三种路线的取舍
1)独立小模型(Draft Model) 最经典的方案,如用 Qwen2.5-0.5B 给 32B 打草稿。优点是实现简单、框架支持最好;缺点是必须同词表,且要额外部署一份权重,显存和运维成本实打实增加。小模型和你业务分布不 match 时,接受率会崩到 20% 以下,加速比直接变负。
2)自草稿(Self-Drafting):Medusa / MTP / EAGLE 这是 2025-2026 年的主流。不再用独立模型,而是在目标模型上挂轻量头:
- Medusa:在最后一层隐藏状态上并联多个 MLP head,分别预测第 +2、+3… 个 token。零额外权重搬运,但接受率中等。
- MTP(Multi-Token Prediction):DeepSeek 系列原生自带,训练时就让模型学会预测后续 token,推理时天然高接受率,是目前性价比最高的路线之一。
- EAGLE / EAGLE3:把目标模型倒数第二层的特征和 token 嵌入拼起来,用一个极小的自回归头做草稿。因为它复用了目标模型的语义特征,接受率显著高于 Medusa,实测 0.6~0.8 是常见区间。EAGLE3 进一步去掉了特征层面的不确定性累积。
3)无模型草稿:n-gram / 检索式 从 prompt 自身或历史语料里用 n-gram 匹配做草稿(如 REST、PLL)。零成本、零显存、零训练,在摘要、抽取、代码补全、RAG 重述这类"输出高度重复输入"的场景里,接受率能意外地高。我的建议是:这类场景优先试 n-gram,成本几乎为零。
选型结论:有原生 MTP 的权重就用 MTP,没有就上 EAGLE3,业务输出高度重复输入就先试 n-gram,只有在以上都不成立时才考虑独立草稿模型。
四、从链式到树状:把接受长度再翻一倍
链式草稿有个结构性问题:γ 个候选串成一条链,中间任何一个位置被拒,后面全废。γ=5 时如果第 2 个被拒,只拿到 2 个 token,白算 3 个。
树状草稿(Tree-based Drafting)的解法是:草稿阶段用 top-k 分支生成一棵候选树,验证阶段用 tree attention mask 让一次前向并行验证整棵树的所有路径。
# 构造树注意力掩码:节点只能看到自己的祖先
def build_tree_mask(parents: list[int], n: int) -> Tensor:
# parents[i] = -1 表示根节点(接在 prefix 之后)
mask = torch.zeros(n, n, dtype=torch.bool)
anc = [set() for _ in range(n)]
for i, p in enumerate(parents):
if p >= 0:
mask[i] = mask[p].clone()
mask[i, p] = True
anc[i] = anc[p] | {p}
mask = mask | torch.eye(n, dtype=torch.bool)
return mask, anc # mask[i,j]=True 表示 i 可 attend j
# 典型 EAGLE 配置:深度 5、每层保留 top-8、总候选 token 数 64
# 实测期望接受长度 3~4,比同 gamma 的链式高 40%+
代价是验证阶段的计算量随候选总数上升,而收益随最长路径上升。工程上的甜点区是总候选 token 数 40~80、深度 4~6,再大验证开销就吃掉了收益。这个数字在 vLLM 和 SGLang 里都作为默认超参存在,不是巧合。
五、生产落地清单
vLLM 示例(EAGLE3):
VLLM_USE_V1=1 vllm serve meta-llama/Llama-3.1-70B-Instruct \
--tensor-parallel-size 4 \
--speculative-config '{
"method": "eagle3",
"model": "yuhuili/EAGLE3-LLaMA3.1-Instruct-70B",
"num_speculative_tokens": 5
}'
SGLang 示例:
python -m sglang.launch_server --model meta-llama/Llama-3.1-70B-Instruct \
--speculative-algo EAGLE3 \
--speculative-draft-model-path yuhuili/EAGLE3-LLaMA3.1-Instruct-70B \
--speculative-num-steps 5 --speculative-eagle-topk 8 \
--speculative-num-draft-tokens 64
DeepSeek 系列直接用原生 MTP 即可,无需外部草稿权重。具体方法名以你所用框架版本为准。
什么时候不该开投机解码:
- 高 QPS 场景。batch 已经 128+,GPU 计算饱和,加投机只会增加调度开销,吞吐反而下降。投机解码是延迟优化手段,不是吞吐优化手段——这句话值得贴墙上。
- 超长输出。输出越长接受率越容易累积性衰减,测出来加速比可能只有 1.1x。
- 显存已经吃紧。草稿权重 + 树验证的 KV cache 都要额外显存,OOM 得不偿失。
- 追求可复现。投机引入了额外随机源,即使 seed 固定,跨版本的 token 序列也可能漂移。
怎么测才不骗自己: 一定要分开看 TPOT(每输出 token 时间)和 E2E 延迟,并记录真实接受率。框架日志里的 acceptance_rate 低于 0.4 时,请先去解决草稿质量问题(换 EAGLE、调 top-k、检查词表是否一致),而不是继续调 γ。γ 从 5 加到 8 在低接受率下的收益远不如把 α 从 0.5 降到 0.3。
六、我的判断
投机解码在 2026 年已经从"论文里的花活"变成了推理引擎的默认配置项。真正的工程门槛不在算法——算法就那 30 行——而在于三件事:草稿方案与业务分布的匹配度、树形超参的调优、以及接受率的持续可观测。
把接受率做成大盘指标、按业务场景分组监控,比纠结用哪个框架重要得多。当你的 Agent 工具调用、代码补全、结构化输出这类低温度场景占到七成流量时,投机解码是投入产出比最高的那一项优化——不用改模型、不用重新训练、不损失任何输出质量,改几行配置就有 1.5~2.5x 的 TPOT 下降。这种便宜,线上服务里不多见。

发表评论 取消回复