投机解码
廉价的草稿模型先猜 k 个 token,大模型一次前向就能验证全部 k 个。拒绝采样保证输出分布可证明地完全一致。
- 草稿模型
- n-gram 查找
- 验证
- 拒绝采样
- 接受率
- EAGLE/Medusa
为什么重要
问题
decode 受显存带宽限制:为了产出一个 token 要读 16 GB 权重,而算术单元几乎全程闲着。由此可以推出一个出人意料的结论。
检查五个候选 token 的代价,和生成一个 token 几乎完全一样。 贵的是那次权重读取,而它是共享的。多出四个位置的算术等于白送:它正好装进你已经付过钱的那部分闲置产能里。
那么,如果有个便宜的东西能猜出接下来几个 token,大模型就能在一次前向里把这些猜测全部检查一遍。唯一的问题是:这套检查是否保持输出分布不变,还是说你在悄悄提供一个更差的模型。
核心思路
解法
一个廉价的提议器生成 k 个候选 token。目标模型对全部 k+1 个位置跑一次前向。接着由一个拒绝采样检验接受最长的正确前缀,而这个检验的构造保证:被接受的 token 的分布与目标模型单独生成时完全一致。
def verify(draft_tokens, q_probs, p_probs, rng):
"""q = draft's distribution, p = target's. Output is distributed as p."""
accepted = []
for i, tok in enumerate(draft_tokens):
p, q = p_probs[i][tok], q_probs[i][tok]
if rng.random() < min(1.0, p / q): # accept
accepted.append(tok)
continue
# rejected: resample from the RESIDUAL distribution
residual = np.maximum(0.0, p_probs[i] - q_probs[i])
return accepted + [sample(residual / residual.sum(), rng)]
# every draft token accepted -> take the target's own next token free
return accepted + [sample(p_probs[len(draft_tokens)], rng)]构造上就是无损的
min(1, p/q) 加上从残差 max(0, p−q) 重采样,构成了一个修正版拒绝采样器:对每个 token,接受它的概率加上从残差重采样到它的概率,恰好等于 p。这是一个直接的恒等式,Leviathan 等与 Chen 等(2023)用寥寥几行就证明了它。投机解码不是「质量换速度」的交易。如果你的实现改变了输出,那就是有 bug。工作原理
工作原理
经济账
设每 token 接受率为 α,则每次目标前向输出的期望 token 数是几何级数:
E[tokens] = (1 - α^(k+1)) / (1 - α)
α = 0.9, k = 6 -> 5.2 tokens per target pass
α = 0.8, k = 4 -> 3.4
α = 0.5, k = 4 -> 1.9
α = 0.3, k = 4 -> 1.4 (barely worth the draft cost)一次拒绝会丢弃它之后的一切,因为那些提议都是以一个已经不存在的 token 为条件的。收益因此随 k 迅速递减。越过某个最优 k 之后,多出来的提议几乎轮不到,却照样要花草稿时间,投机得太远只会比普通解码更慢。
三种提议方式
- 草稿模型 —— 同系列的小模型(用 Llama-1B 为 Llama-70B 打草稿)。接受率最高,但要付出真实算力,还要管理第二个模型及其自己的 KV 缓存。
- n-gram / prompt 查找 —— 在 prompt 和近期输出中搜索当前后缀,把上次它后面跟的内容复制过来。完全不需要模型。对开放式写作接受率很差,对摘要、代码编辑和 RAG 则相当出色,因为在那些场景里,输出重复输入本来就是理所当然的。
- EAGLE / Medusa —— 在目标模型自身隐状态上训练的几个小型附加头,用来预测未来若干个位置。非常便宜、接受率高,代价是每个模型都要单独训练。EAGLE 系方法是 vLLM 中的旗舰投机选项。
它为什么和批处理相冲突
投机解码花的是闲置算力。批处理花的是同一份算力。batch 为 1 时富余很多,投机是一大笔收益;batch 为 128 时 GPU 已经算力饱和,投机基本只是在加活。
因此生产引擎会让投机长度动态化:batch 小的时候激进投机,并发上升就收敛,并根据实测接受率在线调整 k。
动手观察
动手试试
- token / 每次目标前向
- 2.79
- 理论值
- 3.05
- 加速比
- 1.89×
- 当前最优 k
- 4 (2.06×)
- 被接受的草稿 token
- 首个拒绝 —— 本轮结束
- 额外 token,同一次前向免费附赠
拒绝是前缀终止的:一旦某个 token 被拒,它之后的所有草稿都不再检查、直接丢弃,因为它们的条件已经不存在了。这就是每轮期望 token 数关于 α 是几何级数、而不是关于 k 线性的原因。
这条曲线存在内部极大值。越过它之后,多出来的草稿几乎永远轮不到,却照样要花草稿模型的时间,所以投机得太远反而是减速。用弱一点的草稿来源(试试 n-gram),最优值可能只有 k = 1 或 2。
- 草稿模型 — 同系列的小模型。接受率最高,但要付出真实算力,还要额外管理一套权重和 KV 缓存。
- n-gram 查找 — 直接从 prompt 里复制曾经出现过的后续。几乎零成本、不需要模型;在开放式写作上接受率很差,在摘要、代码编辑和 RAG 这类「输出重复输入」的场景里极好。
- EAGLE 头 — 在目标模型上加几个额外的头,用它自己的隐状态预测后续若干位置。非常便宜、接受率高,但必须针对每个模型单独训练。
切换草稿来源,会看到成本和接受率同时变化。最优解取决于你的负载里输出是否与输入相似。
当前设置下的理论加速比:2.06×。24 轮的实测值:1.89×。轮数越多两者越接近。这个方差是真实存在的,也是投机解码改善平均延迟、却略微恶化延迟方差的原因。
把接受率设成 0.35、k 设成 10。加速比掉到 1× 以下:你现在比普通解码还慢,付了十次草稿前向的钱,只换回大约一个 token。再把提议器换成 n-gram,成本降得足够低,即使接受率很差也依然有赚。prompt 查找解码值得上线,靠的正是这一点,接受率不高也无妨。
亲手实现
实现
出错的地方在记账。部分接受之后,你必须把每一个被拒绝位置的 KV 缓存回滚,而且两个模型都要回滚。
def step(self, req):
draft_tokens, q_probs = self.proposer.propose(req, k=self.k)
# one target pass over [context + all k drafts]
p_probs = self.target.forward(req.tokens + draft_tokens)
accepted = verify(draft_tokens, q_probs, p_probs, self.rng)
# CRITICAL: undo the KV written for tokens we did not accept
n_rejected = len(draft_tokens) - (len(accepted) - 1)
if n_rejected > 0:
self.target.kv.truncate(req, n_rejected)
self.proposer.kv.truncate(req, n_rejected)
req.tokens.extend(accepted)
self.stats.record(proposed=len(draft_tokens), accepted=len(accepted) - 1)
return accepted忘记回滚 KV 缓存
temperature=0 和固定种子,分别在开启和关闭投机的情况下跑引擎,断言输出逐 token 一致。$ python code/s14_speculative_decoding.py只需要 NumPy — 查看环境准备.
生产实践
在生产环境中
- vLLM —— 支持 n-gram、草稿模型和 EAGLE 三种提议器,投机长度会随 batch 大小自适应。
- Medusa —— 多个头产出候选树而不是一条链,用树形注意力掩码在一次前向里验证整棵树。每次前向的接受量更高。
- EAGLE-2/3 —— 目前的最好水平;动态形状的草稿树,接受率高到足以带来 3–4× 的端到端加速。
- 自投机技巧 —— 跳过目标模型的部分层来构成草稿。完全不需要额外权重,代价是接受率更低。
练习
- 1实现 prompt 查找解码,在两种负载上测量接受率:开放式创意写作 vs「总结这份文档」。差距应该非常悬殊。
- 2实现基于树的投机:提议一棵分叉的树而不是一条链,构造树形注意力掩码,并在一次前向里验证整棵树。在相同草稿成本下比较每次前向的期望 token 数。
- 3加上自适应 k:跟踪一个滚动的接受率估计,调整投机长度以最大化实测吞吐。展示它收敛到静态扫描找出的那个最优值。
继续学习
接下来
投机约束的是 token 何时被产出。S15 约束的是哪些 token 被允许出现:把 JSON schema 编译成状态机,把语法禁止的每一个 logit 都屏蔽掉,非法输出于是不再是「不太可能」,而是「不可能」。
自测
习题
先作答,再看解析。答错比答对更有价值,因为解析会指出你该回头重读哪一部分。
为什么验证 k 个草稿 token 的代价和生成一个 token 差不多?