大模型推理加速(四):推测解码为什么不改变分布,什么时候反而更慢

Jellow 编辑发布,AI 辅助生成并完成技术核对。 计算示例和实测结果会在文中分别说明。

自回归生成有一条难以绕开的依赖:确定第 t 个 token 后,才能知道第 t+1 个位置的输入。推测解码的做法听起来有些冒险——先让一个便宜的模型猜几个 token,再让目标模型一次检查。

如果只是“猜对就留,猜错就重来”,采样分布会被改变。真正关键的是接受概率和拒绝后的修正分布。它们保证输出仍服从目标模型,而不是悄悄变成草稿模型的偏好。

本文使用最基础的 speculative sampling 解释这件事。不同推理框架还有 n-gram、Medusa 等提议方式,具体接口和限制不能直接套用。

1. 一轮推测解码做了什么

记目标模型的条件分布为 p,草稿模型的条件分布为 q。一轮生成可以分成四步:

  1. 草稿模型自回归提出 k 个候选 token,并保存各步的 q 概率;
  2. 目标模型在已知候选前缀上一次前向,得到对应位置的 p 概率;
  3. 从第一个候选开始依次验收,候选 x 的接受概率是 min(1, p(x) / q(x))
  4. 第一次拒绝时,从修正分布 normalize(max(p - q, 0)) 采样一个 token,然后结束本轮。如果 k 个候选全部接受,再从目标模型验证结果中采样一个额外 token。

目标模型虽然不能提前知道最终会接受几个候选,但这一轮的候选前缀已经由草稿模型给出。它可以并行计算这些位置的 logits。接受判断仍要按顺序进行,因为后一个位置的 p、q 建立在前面候选都被保留的条件上。

这一算法由两组同期工作系统化提出,可参阅 Leviathan 等人的论文Chen 等人的论文

2. 为什么修正后仍然服从目标分布

先看只有一个位置、三个 token 的例子:

token       A     B     C
p          0.6   0.3   0.1
q          0.2   0.5   0.3
min(p,q)   0.2   0.3   0.1
max(p-q,0) 0.4   0.0   0.0

候选从 q 采样后被接受为 token x 的总概率是:

q(x) × min(1, p(x)/q(x)) = min(p(x), q(x))

三种 token 的接受概率之和为 0.6,所以拒绝概率是 0.4。拒绝后,max(p-q, 0) 归一化得到 [1, 0, 0],也就是一定选择 A。

最终分布为:

接受部分 [0.2, 0.3, 0.1]
+ 拒绝后 [0.4, 0.0, 0.0]
= 目标分布 [0.6, 0.3, 0.1]

一般情况下也有逐项恒等式:

min(p, q) + max(p - q, 0) = p

这就是修正分布不能省略的原因。工程实现还要处理 q(x)=0、数值误差、EOS 和截断等边界。

这里的“分布相同”不等于同一个随机种子会输出完全相同的文本。推测解码消耗随机数的顺序不同。temperature、top-p 等采样变换也必须体现在用于验收的实际 p、q 中;目标与草稿通常还需要兼容的 tokenizer 和词表。

3. 接受率怎样影响一轮产出

为了建立直觉,假设每个候选在给定前面都接受的条件下,接受概率恒为 α。第 i 个候选被接受,需要前 i 个全部通过,因此概率为 αⁱ。

忽略 EOS,一轮平均产出的 token 数是:

E[token/round] = 1 + α + α² + ... + αᵏ

末尾的 1 来自第一次拒绝后的修正 token,或者全部接受后的额外 token。

例如 k = 4、α = 0.8:

E = 1 + 0.8 + 0.64 + 0.512 + 0.4096
  = 3.3616 token/round

不能直接用 1 + k × α,因为后面的候选只有在前面全部接受时才有机会保留。实际每一步的接受概率会随上下文变化,这个等比模型只是估算。

4. 多产出 token 还不等于更快

设普通目标模型每生成一个 token 的时间为 Tbase,一轮推测解码需要:

Tround = Tdraft(k) + Tverify(k) + Toverhead

只有当 Tround / E[token/round] < Tbase 时,单请求平均每 token 时间才下降。

构造一个纯计算示例:Tbase 为 10 ms;草稿生成 4 个候选共 4 ms;目标验证用 14 ms;调度和采样额外 2 ms。使用上一节的平均产出:

20 ms / 3.3616 ≈ 5.95 ms/token
10 ms / 5.95 ms ≈ 1.68 倍

若接受率降到 0.3,同样 k = 4 时平均只产出约 1.43 个 token,20 ms / 1.43 已经慢于普通解码。上面的时间全部是假设值,只说明决策方法。

高接受率也不是充分条件。草稿模型太大、验证形状不适合硬件、跨设备传输昂贵,都会吃掉收益。在高并发服务中,目标模型原本就能通过 batching 保持忙碌;推测解码增加的计算可能降低总体吞吐,即使单请求延迟有所改善。

5. 草稿模型应该怎样选

好的草稿模型需要同时满足两件事:提出候选足够便宜,候选又与目标分布足够接近。只比较参数量无法判断结果。

可以按以下顺序做实验:

  1. 固定目标模型、tokenizer、采样参数和请求集合;
  2. 先测普通解码的 TTFT、TPOT、吞吐和显存;
  3. 对每个草稿方案记录候选长度、各位置接受率、每轮接受数分布;
  4. 分别测低并发延迟与高并发吞吐,不混成一个结论;
  5. 按任务拆分结果,例如代码、翻译、开放问答和长上下文。

接受率低时,先确认草稿与目标是否使用同一模板、相同采样变换和兼容 tokenizer,再调整候选长度。k 越大,单轮潜在产出越多,目标验证和草稿生成的成本也更高,没有对所有负载通用的最佳值。

Jay Mody 的实现讲解把自回归基线、验收步骤与复杂度放在同一篇文章中,适合对照伪代码理解。接入服务前,还可以把接受率与分布距离联系起来,并用边际成本选择候选长度。

6. 接受率其实对应两个分布的重叠

对单个位置,候选被接受的总概率是:

α = Σx min(p(x), q(x))

总变差距离定义为:

TV(p, q) = 1/2 × Σx |p(x) - q(x)|

利用概率和都为 1,可以得到:

α = 1 - TV(p, q)

因此“草稿模型接近目标模型”可以变成一个更具体的判断:经过实际 temperature、top-p 等变换后,两个条件分布重叠得越多,这一步的接受概率越高。

这个关系只描述当前条件前缀上的一个位置。生成任务里,每次接受都会改变后续条件分布,不能用少数 prompt 的平均 α 代表所有位置。更有解释力的监控应该按候选位置记录接受率:第一个候选经常通过、第四个候选经常失败,说明继续增大 k 的边际收益已经很低。

一般情况下,一轮产出的期望可以写成:

E[token/round]
= 1 + P(A1) + P(A1∩A2) + ... + P(A1∩...∩Ak)

只有在每一步条件接受概率都近似相同且为 α 时,才化成前面的等比数列。这也是为什么服务应记录“每轮接受 token 数分布”,而不只记录把所有位置混在一起的平均接受率。

7. 用边际收益选择候选长度 k

将 k 从 4 增加到 5,多出的潜在收益只有“前 5 个候选全部被接受”的那一项;成本却一定包含草稿模型再生成一步,并可能增加目标验证工作和临时张量。

可以直接测量相邻配置:

边际产出 = E[token/round | k+1] - E[token/round | k]
边际成本 = Tround(k+1) - Tround(k)

边际成本 / 边际产出 已高于普通目标模型每 token 的成本时,继续增加 k 没有意义。这个比较比“接受率超过某个固定百分比就启用”更可靠,因为不同硬件和实现中的验证成本并不相同。

假设各位置的条件接受率都是 0.8,从 k=4 增到 k=5,只增加 0.8⁵≈0.328 个期望 token。如果多一个候选让轮次增加 2 ms,那么边际成本约为 2/0.328≈6.1 ms/token。它是否值得,仍要和基线每 token 成本以及高并发吞吐损失比较。这些数字仍是公式示例。

8. 怎样验证“分布没变”

逐次比较同一随机种子的文本不是有效验证,因为随机数消耗顺序可以不同。更合理的测试分三层:

  1. 枚举小词表:像本文三 token 例子一样,精确求出接受与修正后的概率,逐项比较目标 p;
  2. 随机分布测试:生成许多归一化的 p、q,检查 min(p,q) 加拒绝质量后是否恢复 p,并覆盖零概率和浮点极值;
  3. 模型统计测试:在固定 prompt 集上大量采样,比较 token 频率或序列统计量,同时保留显著性阈值和样本量。

端到端实现还要专门测试:EOS 出现在草稿中间、目标和草稿词表映射、top-k/top-p 截断后 q 为零、批次中各请求拒绝位置不同,以及流式接口一次返回多个 token。数学公式正确,不代表这些状态机边界自动正确。

9. 部署时看这张决策表

现象 判断 下一步
低并发 TPOT 高,目标模型单步很贵 有潜在价值 测 k、各位置接受率与轮次成本
接受率高,但总体没变快 草稿或验证开销过高 分解 Tdraft、Tverify、调度和传输
k 越大,后部接受率快速下降 候选过长 按边际成本缩短 k 或动态调整
单请求变快,高并发 goodput 下降 额外计算挤占 batching 收益 按线上负载决定是否仅对部分请求启用
输出统计偏离目标模型 实现正确性问题 检查修正分布、采样变换、EOS 与词表
不同任务收益差异很大 p、q 接近程度依赖任务 按任务路由,避免使用全局开关

推测解码最值得启用的组合是:目标模型单步昂贵、草稿足够便宜、实际任务上的分布重叠高,而且目标验证一段候选的成本没有随 k 等比例增长。只缺其中一项,都可能让一个数学上正确的算法在系统上变慢。

继续阅读


本站的内容生成、审核与纠错说明

发表评论