AI 工程基础体系 · 第 82/100 篇。内容覆盖机器学习、深度学习与生成式 AI;模型、数据、评测、权限和成本会作为同一生产系统处理。

推测解码:草稿模型、接受率、正确性与加速边界

推测解码(Speculative Decoding)是一类用于自回归语言模型推理的加速方法。它不改变最终目标模型的生成分布,而是先让一个较小、较快的草稿模型提出多个候选 token,再让较大的目标模型一次性验证这一段候选。

它利用了 Transformer 推理中的一个结构性事实:

  • 生成 token 时,目标模型通常必须逐步运行;
  • 但在已经知道一段候选 token 的情况下,目标模型可以利用因果注意力,一次前向计算这段候选对应的多个位置;
  • 如果候选大多正确,就能用一次目标模型前向替代多次目标模型单步前向。

因此,推测解码的核心不是“让小模型替代大模型”,而是:

用小模型承担候选生成,用大模型承担最终分布验证;在满足特定采样规则时,输出分布仍然与直接使用目标模型一致。

这里的“正确性”有两种不同含义:

  1. 分布正确性:随机采样时,最终 token 分布与目标模型单独采样完全一致;
  2. 解码结果正确性:贪心解码时,输出与目标模型逐 token 取最大概率的结果一致。

这两者需要使用不同的接受规则,不能混为一谈。


1. 前置:自回归模型为什么会成为推理瓶颈

设输入上下文为 x<tx_{<t},语言模型在第 tt 步输出词表上的概率分布:

p(xtx<t)p(x_t \mid x_{<t})

其中:

  • x<tx_{<t} 是已经生成的 token;
  • xtx_t 是待生成的下一个 token;
  • pp 是目标模型给出的概率分布。

生成下一个 token 后,上下文变为:

xt=(x<t,xt)x_{\leq t} = (x_{<t}, x_t)

然后模型再计算:

p(xt+1xt)p(x_{t+1} \mid x_{\leq t})

这个过程具有严格的自回归依赖。第 t+1t+1 个位置的输入包含第 tt 个位置刚生成的 token,因此普通解码通常是:

目标模型前向 -> 生成 token 1
目标模型前向 -> 生成 token 2
目标模型前向 -> 生成 token 3
...

Transformer 的因果注意力并不意味着每个生成步骤都无法并行。对于一段已经确定的输入序列,模型可以并行计算多个位置的 logits;真正限制普通生成的是,后一个位置的输入 token 在前一个步骤结束前尚未确定。

推测解码把“尚未确定的未来 token”暂时交给草稿模型预测:

草稿模型连续生成多个候选 token
目标模型一次前向验证这些候选
保留前缀中被接受的部分
从第一个拒绝位置继续

这里的并行化只发生在目标模型的验证阶段,并没有消除自回归依赖,而是用一个便宜模型近似预测这段依赖。


2. 草稿模型与目标模型

2.1 目标模型

目标模型(target model)是生产系统真正希望使用的语言模型。它通常参数量更大、质量更高、推理成本更高。

记目标模型在给定前缀 ss 时的下一个 token 分布为:

p(s)p(\cdot \mid s)

推测解码必须以目标模型的分布为最终依据。草稿模型不能直接决定最终输出。

2.2 草稿模型

草稿模型(draft model)是用来快速提出候选 token 的模型,记其分布为:

q(s)q(\cdot \mid s)

草稿模型通常具备以下特征:

  • 参数量明显小于目标模型;
  • 单 token 延迟较低;
  • 与目标模型使用相同或兼容的 tokenizer;
  • 输出分布尽量接近目标模型;
  • 可以在目标模型出现之前连续生成多个 token。

草稿模型可以是:

  • 目标模型的较小版本;
  • 目标模型的蒸馏模型;
  • 共享词表和训练语料的轻量模型;
  • 某些实现中的专用 draft head 或辅助头。

“更小”并不自动意味着“更适合”。如果草稿模型虽然小,但分布与目标模型差异很大,候选接受率会下降,草稿模型自身的计算也可能抵消收益。

2.3 两个模型必须对齐的内容

若目标是保持严格的采样分布,至少需要明确以下配置:

  • tokenizer 和词表;
  • BOS、EOS、特殊 token 的处理;
  • temperature;
  • top-k、top-p 等截断规则;
  • repetition penalty 等 logits 变换;
  • 禁止词或词级约束;
  • 随机采样策略;
  • 上下文截断和位置编码方式。

严格来说,ppqq 不是“原始 logits 转成 softmax 后的分布”就结束了,而是经过生产解码规则处理后的实际采样分布。如果目标模型实际使用了 temperature 和 top-p,那么接受率公式中的 pp 应该是处理后的目标分布,而不是未经处理的原始分布。


3. 一轮推测解码的数据流

设当前已确认前缀为:

s = [我, 喜欢]

草稿长度设为 γ=4\gamma=4。草稿模型连续生成:

候选 d1 = 看
候选 d2 = 电影
候选 d3 = ,
候选 d4 = 周末

目标模型随后接收:

[我, 喜欢, 看, 电影, ,, 周末]

通过一次带因果 mask 的前向,得到对应位置的目标分布:

p1 = p(x | [我, 喜欢])
p2 = p(x | [我, 喜欢, 看])
p3 = p(x | [我, 喜欢, 看, 电影])
p4 = p(x | [我, 喜欢, 看, 电影, ,])
p5 = p(x | [我, 喜欢, 看, 电影, ,, 周末])

其中:

  • p1p_1 用于验证 d1d_1
  • p2p_2 用于验证 d2d_2
  • ...
  • p4p_4 用于验证 d4d_4
  • p5p_5 是所有草稿 token 都接受后,用于生成额外的 bonus token。

目标模型不能用 p2p_2 验证任意 token。它必须在草稿前缀已经被接受的条件下使用对应位置的分布。若第一个草稿 token 被拒绝,那么后面的草稿 token 都建立在错误前缀上,必须丢弃。

典型状态变化如下:

flowchart LR
    A[已确认前缀] --> B[草稿模型连续生成 γ 个候选]
    B --> C[目标模型一次前向计算 γ+1 个位置]
    C --> D{逐 token 验证}
    D -->|全部接受| E[提交 γ 个候选并采样 bonus token]
    D -->|第 j 个拒绝| F[提交前 j-1 个候选]
    F --> G[从第 j 个目标分布采样替代 token]
    G --> H[丢弃第 j 个及之后的草稿状态]
    E --> I[进入下一轮]
    H --> I

关键路径是:验证必须从左到右进行。目标模型虽然并行计算了多个位置,但候选 token 的接受结果仍然具有前缀依赖。


4. 接受率:从候选概率到接受概率

4.1 采样模式下的接受概率

在某个固定前缀 ss 下:

  • 草稿模型从 q(xs)q(x\mid s) 采样候选 token xx
  • 目标模型给出 p(xs)p(x\mid s)

对候选 xx,定义接受概率:

a(xs)=min(1,p(xs)q(xs))a(x\mid s)=\min\left(1,\frac{p(x\mid s)}{q(x\mid s)}\right)

q(xs)>0q(x\mid s)>0 时,这个公式直观地表示:

  • 如果目标模型认为 xx 的概率不低于草稿模型,即 p(x)q(x)p(x)\ge q(x),则无条件接受;
  • 如果目标模型认为 xx 被草稿模型高估,则按比例接受。

如果 q(x)=0q(x)=0,草稿模型不会采样到 xx,因此不会发生除零。实现中仍应显式处理零概率,而不能直接对张量做无保护除法。

4.2 为什么接受率是 xmin(p(x),q(x))\sum_x \min(p(x),q(x))

候选 xx 被提出的概率是 q(x)q(x),提出后被接受的概率是:

min(1,p(x)q(x))\min\left(1,\frac{p(x)}{q(x)}\right)

所以候选 xx 被接受的联合概率为:

q(x)min(1,p(x)q(x))=min(q(x),p(x))q(x)\min\left(1,\frac{p(x)}{q(x)}\right) = \min(q(x),p(x))

对所有 token 求和,得到单个位置的期望接受率:

A=xmin(p(x),q(x))A=\sum_x \min(p(x),q(x))

它还可以写成:

A=1TV(p,q)A=1-\mathrm{TV}(p,q)

其中 TV\mathrm{TV} 是总变差距离:

TV(p,q)=12xp(x)q(x)\mathrm{TV}(p,q)=\frac12\sum_x |p(x)-q(x)|

因此,接受率越高,说明草稿分布和目标分布越接近。接受率不是一个单纯的“模型准确率”,而是两个完整概率分布之间的重叠程度。

4.3 完整数值算例

假设某个位置的词表只有三个 token:

token 目标分布 pp 草稿分布 qq
A 0.50 0.25
B 0.30 0.50
C 0.20 0.25

草稿模型采样结果为 B。因为:

a(B)=min(1,0.30/0.50)=0.6a(B)=\min(1,0.30/0.50)=0.6

所以 B 有 60% 的概率被接受,40% 的概率被拒绝。

如果草稿采样结果为 A:

a(A)=min(1,0.50/0.25)=1a(A)=\min(1,0.50/0.25)=1

A 必然接受。

如果采样结果为 C:

a(C)=min(1,0.20/0.25)=0.8a(C)=\min(1,0.20/0.25)=0.8

C 有 80% 的概率接受。

整个位置的期望接受率为:

A=min(0.50,0.25)+min(0.30,0.50)+min(0.20,0.25)A=\min(0.50,0.25)+\min(0.30,0.50)+\min(0.20,0.25)

A=0.25+0.30+0.20=0.75A=0.25+0.30+0.20=0.75

这说明:虽然草稿模型最偏好 B,而目标模型最偏好 A,但两个分布仍有 75% 的重叠质量。


5. 拒绝后为什么要从残差分布采样

仅仅“拒绝草稿 token,然后从目标模型重新采样”并不能保证结果分布正确。正确的拒绝处理需要构造残差分布:

r(xs)=max(0,p(xs)q(xs))ymax(0,p(ys)q(ys))r(x\mid s)= \frac{\max(0,p(x\mid s)-q(x\mid s))} {\sum_y \max(0,p(y\mid s)-q(y\mid s))}

如果候选被拒绝,就从 rr 中采样替代 token。

在上面的数值例子中:

pq=[0.25,0.20,0.05]p-q = [0.25,-0.20,-0.05]

取正部:

[pq]+=[0.25,0,0][p-q]_+=[0.25,0,0]

归一化后:

r=[1,0,0]r=[1,0,0]

因此,当 B 被拒绝时,替代 token 必然是 A。

5.1 正确性的推导

对任意 token xx,最终输出它的概率有两部分:

  1. 草稿模型提出 xx,且 xx 被接受;
  2. 草稿候选被拒绝,随后从残差分布采样到 xx

第一部分概率是:

q(x)min(1,p(x)q(x))=min(p(x),q(x))q(x)\min\left(1,\frac{p(x)}{q(x)}\right) = \min(p(x),q(x))

草稿被拒绝的总概率是:

1ymin(p(y),q(y))1-\sum_y\min(p(y),q(y))

残差分布的分母为:

ymax(0,p(y)q(y))\sum_y\max(0,p(y)-q(y))

注意恒等式:

ymax(0,p(y)q(y))=1ymin(p(y),q(y))\sum_y\max(0,p(y)-q(y)) = 1-\sum_y\min(p(y),q(y))

因此第二部分概率为:

(1ymin(p(y),q(y)))max(0,p(x)q(x))1ymin(p(y),q(y))\left(1-\sum_y\min(p(y),q(y))\right) \frac{\max(0,p(x)-q(x))} {1-\sum_y\min(p(y),q(y))}

化简得到:

max(0,p(x)q(x))\max(0,p(x)-q(x))

两部分相加:

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

所以最终输出分布正好是 p(x)p(x)

这就是采样模式下推测解码能够保持目标模型分布的原因。接受概率和残差分布必须配套使用;只采用其中一个都不足以构成严格的分布保持算法。


6. 多 token 草稿的验证过程

设草稿模型产生:

d1,d2,,dγd_1,d_2,\ldots,d_\gamma

目标模型在相应前缀下给出:

p1,p2,,pγ,pγ+1p_1,p_2,\ldots,p_\gamma,p_{\gamma+1}

逐位置验证:

ai=min(1,pi(di)qi(di))a_i=\min\left(1,\frac{p_i(d_i)}{q_i(d_i)}\right)

其中:

qi()=q(s,d1,,di1)q_i(\cdot)=q(\cdot\mid s,d_1,\ldots,d_{i-1})

pi()=p(s,d1,,di1)p_i(\cdot)=p(\cdot\mid s,d_1,\ldots,d_{i-1})

具体流程是:

  1. q1q_1 采样 d1d_1
  2. q2q_2 采样 d2d_2,条件中包含 d1d_1
  3. 继续生成直到 dγd_\gamma
  4. 目标模型一次性计算 p1,,pγ+1p_1,\ldots,p_{\gamma+1}
  5. 验证 d1d_1
  6. d1d_1 被拒绝,丢弃 d2,,dγd_2,\ldots,d_\gamma,从 p1q1p_1-q_1 的残差分布采样;
  7. d1d_1 被接受,再验证 d2d_2
  8. 第一个拒绝位置使用该位置的残差分布;
  9. 如果全部 γ\gamma 个候选都接受,则从 pγ+1p_{\gamma+1} 采样一个额外 token。

“bonus token”并不是凭空增加的 token。它是目标模型已经在同一次验证前向中计算出的下一个位置分布,因此可以避免再启动一次目标模型单步前向。

6.1 多 token 正确性的条件

多 token 情况下,单个位置的接受公式仍然成立,但必须满足:

  • 草稿 token 的条件分布确实是 qiq_i
  • 目标分布 pip_i 使用包含已验证草稿前缀的正确条件;
  • 一旦第 ii 个 token 被拒绝,后续草稿 token 不得继续提交;
  • 替代 token 必须从第 ii 个位置的残差分布采样;
  • 如果全部接受,bonus token 必须从目标模型的下一位置分布采样。

后续草稿 token 不能在拒绝后继续复用,因为它们的条件上下文已经发生变化。


7. 贪心解码与随机采样不是同一种接受规则

7.1 贪心模式

贪心解码的目标是:

xt=argmaxxp(xx<t)x_t=\arg\max_x p(x\mid x_{<t})

此时没有随机采样分布保持问题。可以使用如下规则:

  • 若草稿 token did_i 等于目标模型的最大概率 token,则接受;
  • 否则拒绝草稿 token,直接使用目标模型的最大概率 token。

例如:

目标模型分布:A=0.45, B=0.40, C=0.15
草稿 token:B

虽然 B 可能具有较高概率,但目标模型的贪心结果是 A,因此必须拒绝 B,并输出 A。

贪心模式下不能简单使用采样模式的概率接受规则来理解结果。一个概率为 0.8 的草稿 token,并不代表贪心时有 80% 的机会被接受;贪心验证通常比较目标模型和草稿模型的 argmax 是否一致。

7.2 温度采样

当使用 temperature TT 时,实际分布通常是:

pT(x)=softmax(z(x)T)p_T(x)=\mathrm{softmax}\left(\frac{z(x)}{T}\right)

其中 z(x)z(x) 是模型 logits。

如果线上目标模型使用 T=0.7T=0.7,草稿模型使用 T=1.0T=1.0,则接受公式不能直接比较两套未处理 logits,也不能把目标模型的原始 softmax 当作 pp。应当让 ppqq 分别表示经过各自实际解码处理后的分布。

为了严格保持目标采样分布,草稿模型可以使用与目标模型不同的 qq,但残差校正必须基于这两个实际分布计算。

7.3 Top-k、top-p 与支持集问题

top-k 和 top-p 会把一部分 token 的概率设为零。此时:

  • q(x)=0q(x)=0,草稿不会提出 xx
  • p(x)>0p(x)>0q(x)=0q(x)=0,该概率质量会进入残差分布;
  • 若目标模型和草稿模型使用不同的截断策略,接受率可能显著变化。

例如:

目标模型允许 {A, B, C}
草稿模型只允许 {A, B}

那么 C 无法被草稿直接提出,但它仍可能在拒绝校正阶段出现。若实现错误地把目标模型和草稿模型的零概率位置都裁掉,就会改变目标分布。


8. 可执行的单步采样示例

下面的 Python 示例实现一个简化版的单位置推测采样。它不运行神经网络,而是直接使用两个离散分布,因此适合验证接受概率和残差校正逻辑。

import numpy as np

def sample_from_probs(probs, rng):
    probs = np.asarray(probs, dtype=np.float64)
    probs = probs / probs.sum()
    return rng.choice(len(probs), p=probs)

def speculative_sample_one(p, q, rng):
    """
    p: 目标模型分布
    q: 草稿模型分布
    返回:
      token: 最终输出 token 的索引
      accepted: 草稿 token 是否被接受
      draft_token: 草稿模型提出的 token
    """
    p = np.asarray(p, dtype=np.float64)
    q = np.asarray(q, dtype=np.float64)

    if p.ndim != 1 or q.ndim != 1 or len(p) != len(q):
        raise ValueError("p 和 q 必须是一维且长度相同")
    if np.any(p < 0) or np.any(q < 0):
        raise ValueError("概率不能为负数")
    if not np.isclose(p.sum(), 1.0) or not np.isclose(q.sum(), 1.0):
        raise ValueError("p 和 q 必须归一化")

    draft_token = sample_from_probs(q, rng)

    # q[draft_token] 不会为 0,因为该 token 是由 q 采样得到的。
    accept_prob = min(1.0, p[draft_token] / q[draft_token])

    if rng.random() < accept_prob:
        return draft_token, True, draft_token

    residual = np.maximum(p - q, 0.0)
    residual_sum = residual.sum()

    if residual_sum <= 0:
        raise RuntimeError("残差分布为空,概率实现可能存在数值错误")

    residual /= residual_sum
    replacement = sample_from_probs(residual, rng)
    return replacement, False, draft_token


if __name__ == "__main__":
    # token 0=A, 1=B, 2=C
    p = np.array([0.50, 0.30, 0.20])
    q = np.array([0.25, 0.50, 0.25])

    rng = np.random.default_rng(7)

    counts = np.zeros(3, dtype=np.int64)
    accepted = 0
    trials = 200_000

    for _ in range(trials):
        token, was_accepted, _ = speculative_sample_one(p, q, rng)
        counts[token] += 1
        accepted += int(was_accepted)

    print("经验输出分布:", counts / trials)
    print("经验接受率:", accepted / trials)
    print("理论接受率:", np.minimum(p, q).sum())

预期结果应接近:

经验输出分布: [0.50, 0.30, 0.20]
经验接受率: 0.75
理论接受率: 0.75

由于使用随机数,经验值不会精确等于理论值,但试验次数增加时应逐渐接近。

代码中的关键步骤分别对应如下逻辑:

  1. 使用 qq 采样草稿 token;
  2. 根据 p(x)/q(x)p(x)/q(x) 计算接受概率;
  3. 接受时直接返回草稿 token;
  4. 拒绝时构造 max(pq,0)\max(p-q,0)
  5. 将残差归一化后采样替代 token。

若把拒绝分支改成直接从 pp 采样,输出分布通常会偏离目标分布,因为草稿已经贡献了一部分概率质量,不能在拒绝后再次无条件注入完整的 pp


9. Transformer 前向与 KV Cache 如何支持验证

推测解码依赖 Transformer 的因果注意力和 KV Cache。

9.1 普通解码中的 KV Cache

在普通逐 token 解码中,模型会缓存历史 token 的 key 和 value:

第 1 步:计算 token 1,缓存 K1/V1
第 2 步:只计算 token 2,读取 K1/V1,缓存 K2/V2
第 3 步:只计算 token 3,读取 K1/V1/K2/V2,缓存 K3/V3

这样可以避免每一步重复计算历史 token 的投影。

9.2 推测解码中的目标模型 Cache

一轮推测解码中,目标模型可以接收:

已确认前缀 + 多个草稿 token

并在一次前向中计算这些位置的 logits 和隐藏状态。实现通常需要处理两类状态:

  • 已确认前缀的 target KV cache;
  • 草稿候选对应的临时 target KV cache。

若第 jj 个候选被拒绝:

  • d1,,dj1d_1,\ldots,d_{j-1} 可以提交;
  • dj,,dγd_j,\ldots,d_\gamma 不能提交;
  • 拒绝位置的替代 token 会成为新前缀的一部分;
  • 从拒绝位置之后的临时 cache 必须丢弃或重算。

原因是 Transformer 的 KV cache 与实际输入 token 一一对应。如果草稿 token djd_j 被替换为另一个 token,那么用 djd_j 计算出的后续 key/value 已经不再适用于新序列。

草稿模型也需要维护自己的 cache,并在拒绝后回滚到已接受前缀,再追加替代 token。不同框架可能采用:

  • cache 截断;
  • 预分配 cache 后修改有效长度;
  • 重新计算拒绝点后的短前缀;
  • 专用的 assisted generation cache 管理。

这些属于实现细节,不改变采样正确性的数学条件。


10. 一轮到底能生成多少 token

设草稿长度为 γ\gamma

如果第一个拒绝发生在位置 jj

d1, ..., d(j-1) 被接受
dj 被拒绝并替换

这一轮通常提交 jj 个 token:前 j1j-1 个草稿 token 加上一个替代 token。

如果所有 γ\gamma 个草稿 token 都被接受,则还可以从目标模型的下一个位置采样一个 bonus token,因此提交 γ+1\gamma+1 个 token。

设事件 EiE_i 表示前 ii 个草稿 token 全部被接受,则一轮输出 token 数 NN 的期望为:

E[N]=1+Pr(E1)+Pr(E2)++Pr(Eγ)\mathbb{E}[N] = 1+\Pr(E_1)+\Pr(E_2)+\cdots+\Pr(E_\gamma)

其中:

  • 第一个 token 无论候选是否被接受,都会产生一个最终 token,所以有常数项 1;
  • 输出至少 2 个 token,需要第一个草稿 token 被接受;
  • 输出至少 3 个 token,需要前两个都被接受;
  • 以此类推。

如果近似认为每个位置的条件接受率都为 aa,且各位置接受事件近似独立,则:

E[N]1+a+a2++aγ\mathbb{E}[N]\approx 1+a+a^2+\cdots+a^\gamma

a=0.8,γ=4a=0.8,\gamma=4 时:

E[N]1+0.8+0.64+0.512+0.4096=3.3616\mathbb{E}[N]\approx 1+0.8+0.64+0.512+0.4096=3.3616

这只是近似。真实系统中不同位置的接受率通常不同,而且位置之间存在条件依赖。


11. 加速的成本模型

推测解码是否加速,取决于一轮生成的总耗时与实际输出 token 数,而不是只看接受率。

设:

  • Cd(γ)C_d(\gamma):草稿模型生成 γ\gamma 个候选的耗时;
  • Ct(γ+1)C_t(\gamma+1):目标模型验证 γ\gamma 个候选并计算 bonus 位置的耗时;
  • CoC_o:采样、cache 管理、同步和内存访问等额外开销;
  • L=E[N]L=\mathbb{E}[N]:一轮平均输出 token 数;
  • Ct,1C_{t,1}:普通解码时目标模型生成一个 token 的平均耗时。

则推测解码的平均每 token 成本近似为:

Cspec/tokenCd(γ)+Ct(γ+1)+CoLC_{\text{spec/token}} \approx \frac{C_d(\gamma)+C_t(\gamma+1)+C_o}{L}

相对于普通目标模型解码的理想加速比近似为:

SLCt,1Cd(γ)+Ct(γ+1)+CoS \approx \frac{L\cdot C_{t,1}} {C_d(\gamma)+C_t(\gamma+1)+C_o}

只有当:

LCt,1>Cd(γ)+Ct(γ+1)+CoL\cdot C_{t,1} > C_d(\gamma)+C_t(\gamma+1)+C_o

时,推测解码才真正加速。

11.1 为什么接受率高也可能不加速

可能出现以下情况:

  • 草稿模型不够小,Cd(γ)C_d(\gamma) 很大;
  • 目标模型验证长序列的成本接近多次单 token 解码;
  • GPU 上目标模型本来就高度吞吐优化,验证阶段没有获得足够并行收益;
  • batch 较大,目标模型的算力已经被其他请求占满;
  • cache 回滚和动态 shape 带来额外开销;
  • 每轮都需要设备同步,破坏流水线;
  • 输出很短,启动和调度开销占主导。

因此,“接受率 80%”不能直接推导出“速度提升 80%”。接受率只描述候选被保留的概率,不描述草稿成本、目标验证成本和系统调度成本。

11.2 草稿长度不是越大越好

增大 γ\gamma 有两个相反效果:

  • 候选块更长,目标模型一次验证后可能提交更多 token;
  • 草稿模型需要生成更多 token,且更长的候选块更容易在某处被拒绝;
  • 目标模型验证的序列更长,显存和计算成本增加;
  • 一旦早期 token 被拒绝,后面的草稿计算全部浪费。

γ\gamma 过大时,额外候选带来的期望收益可能低于其草稿和验证成本。

实际系统通常需要测量不同 γ\gamma 下的:

  • 每轮平均输出 token 数;
  • 各位置接受率;
  • 草稿模型耗时;
  • 目标模型验证耗时;
  • cache 内存;
  • P50、P95、P99 延迟;
  • 不同输入长度和 batch 下的结果。

固定使用某个草稿长度,不一定适用于所有请求。


12. 加速边界:哪些场景适合,哪些场景不适合

12.1 适合的场景

推测解码通常更适合:

  • 目标模型单 token 成本高;
  • 草稿模型显著更快;
  • 两个模型在任务和领域上接近;
  • 输出具有较强的局部可预测性;
  • batch 较小,重点是降低单请求延迟;
  • 生成长度足够长,可以摊薄初始化开销。

代码生成、结构化文本、格式固定的回答有时更容易获得较高接受率,因为目标模型对局部后缀的预测更稳定。但这不是规范保证,只能通过实际评测验证。

12.2 不适合的场景

以下情况可能收益很低甚至变慢:

  • 草稿模型和目标模型能力差距过大;
  • 草稿模型训练领域与线上流量不匹配;
  • 目标模型使用复杂的动态约束,草稿模型无法同步;
  • 每个位置的候选都高度不确定;
  • batch 推理已经接近 GPU 饱和;
  • 服务端频繁取消请求,草稿阶段的计算容易被浪费;
  • 输出很短或首 token 延迟极其重要;
  • 目标模型部署在高吞吐推理引擎中,单步解码已经经过强优化。

推测解码本质上增加了一个模型和一套状态管理。如果目标模型本身不是瓶颈,增加草稿模型可能只增加系统复杂度。


13. 常见错误与失败表现

13.1 误解一:草稿模型生成什么,目标模型只检查最后一个 token

这是错误的。目标模型必须验证候选块中的每个位置,因为每个位置有不同的条件前缀:

p1 依赖原始前缀
p2 依赖原始前缀 + d1
p3 依赖原始前缀 + d1 + d2

只检查最后一个 token,无法保证前面 token 的目标分布正确。

13.2 误解二:第一个 token 被拒绝后,还可以保留后面的候选

不能保留。若 d1d_1 被替换,原本基于 d1d_1 生成的 d2d_2 已经不再对应当前上下文。

这会同时造成:

  • 结果分布错误;
  • target KV cache 与实际 token 不一致;
  • 后续 logits 条件错误。

13.3 误解三:拒绝后直接从完整目标分布采样

如前面的推导所示,这会重复计算草稿已经覆盖的概率质量,导致最终分布偏移。严格采样需要使用:

[pq]+[pq]+\frac{[p-q]_+}{\sum [p-q]_+}

而不是直接使用 pp

13.4 误解四:比较两个模型的 logits 就能决定接受

logits 不能直接比较,因为:

  • 两个模型的 logit 偏置和尺度可能不同;
  • 需要先转换为与实际采样规则一致的概率;
  • top-k、top-p、temperature 会改变有效分布;
  • 接受规则使用的是 p(x)/q(x)p(x)/q(x),不是简单的 logit 大小差。

在数值实现中可以用 log-probability 计算比值:

p(x)q(x)=exp(logp(x)logq(x))\frac{p(x)}{q(x)} = \exp(\log p(x)-\log q(x))

这样通常比直接对极小概率做除法更稳定。

13.5 误解五:接受率高就代表模型质量高

接受率衡量的是 ppqq 的分布重叠。一个草稿模型可能在目标任务上质量一般,但由于输出分布较平滑,仍有较高重叠;也可能生成文本看起来合理,但在目标模型的高概率区域覆盖不足,接受率并不高。

接受率应与以下指标一起解释:

  • 生成质量;
  • 输出分布一致性;
  • 各位置接受率;
  • 每轮输出 token 数;
  • 真实端到端延迟。

14. 正确性验证方法

14.1 分布级验证

对于固定 prompt 和固定解码配置,可以分别运行:

  1. 目标模型直接采样;
  2. 推测解码采样。

收集大量输出后,对以下对象做统计比较:

  • 第一个 token 的频率;
  • 固定前缀后的下一个 token 分布;
  • 完整序列中 n-gram 分布;
  • 输出长度分布;
  • EOS 出现概率。

对于单位置测试,应检查经验频率是否接近 pp。可以使用总变差距离:

TV^(p^,p)=12xp^(x)p(x)\widehat{\mathrm{TV}}(\hat p,p) = \frac12\sum_x |\hat p(x)-p(x)|

样本量增加时,该距离应按统计误差下降,而不是出现系统性偏差。

14.2 贪心一致性验证

对于贪心模式,运行:

  • 目标模型逐 token argmax;
  • 推测解码验证并提交。

两者必须逐 token 一致。测试时应覆盖:

  • 普通文本;
  • EOS;
  • 长上下文;
  • batch padding;
  • temperature 为 0 的框架特殊路径;
  • repetition penalty 和禁止 token。

14.3 运行时不变量

生产实现可以检查以下不变量:

  • target cache 中的有效 token 数等于已确认序列长度;
  • draft cache 与已确认前缀长度一致;
  • 第一个拒绝位置之后的草稿 token 没有进入最终输出;
  • EOS 出现后不再继续提交 token;
  • target 和 draft 使用相同 tokenizer 映射;
  • 采样前的 logits 处理配置符合预期;
  • 目标模型验证的每个位置对应正确的 attention mask。

这类检查通常可以在小流量或离线测试中启用,发现问题后再关闭高成本断言。


15. 生产系统中的故障路径

推测解码增加了目标模型之外的依赖,因此需要考虑故障处理。

草稿模型不可用

如果草稿模型加载失败、显存不足或推理异常,最安全的降级方式是:

关闭推测解码
恢复目标模型普通自回归解码

不能把草稿 token 当作目标模型输出继续返回。

目标模型验证失败

目标模型是最终权威。如果目标模型前向失败:

  • 当前轮不能提交未经验证的草稿 token;
  • 请求应重试或返回明确错误;
  • cache 状态必须丢弃或回滚到上一个已确认边界。

Cache 回滚失败

如果 cache 长度与已确认 token 数不一致,继续生成可能产生静默错误。应当:

  1. 标记当前请求状态无效;
  2. 清理 target 和 draft cache;
  3. 从完整已确认序列重建 cache;
  4. 若重建失败,降级或终止请求。

配置动态变化

如果请求中途改变了 temperature、top-p、工具约束或禁止词集合,旧的草稿候选不能继续按旧分布验证。解码分布发生变化后,应丢弃当前草稿块,并在新配置下重新生成。


16. 框架集成时应检查什么

主流 Transformer 框架通常会提供生成接口和 assisted generation 一类能力,但具体参数名、模型兼容条件和 cache 实现会随版本变化。不能把某个版本中的实验性参数当作所有版本的稳定 API。

以 Hugging Face Transformers 为例,集成前应确认:

  • 当前版本是否支持目标模型架构;
  • draft model 是否要求相同 tokenizer;
  • generate 的采样参数是否同时作用于两个模型;
  • cache 实现是否支持动态长度和回滚;
  • 是否支持 batch 请求;
  • logits processor 和 stopping criteria 是否在验证路径中一致;
  • 编译、量化和 flash attention 后端是否仍能正确处理候选块。

一个可靠的集成验证顺序是:

  1. 先用小模型和短 prompt 验证输出一致性;
  2. 再验证贪心模式的逐 token 一致;
  3. 再验证随机采样的经验分布;
  4. 再测单请求延迟;
  5. 最后测试 batch、长上下文、取消请求和显存压力。

不能只比较生成文本是否“看起来一样”。随机采样本来就可能生成不同文本,分布正确性必须通过统计检验验证。


17. 与 Transformer 架构的关系

推测解码并不是对注意力机制的替代,而是利用了 Transformer 的两个性质:

因果 mask 保证位置语义正确

对于候选块中的第 ii 个位置,因果注意力只允许它看到:

原始前缀 + 前面已经出现的候选 token

因此,目标模型可以一次前向计算多个位置,同时保持每个位置的自回归条件。

KV Cache 降低已确认前缀的重复计算

已确认前缀的 key/value 可以复用,目标模型主要处理新增候选块。若没有 KV Cache,验证候选块时会重复计算大量历史上下文,推测解码的收益会显著下降。

Attention Is All You Need 介绍了 Transformer 以注意力机制进行序列建模的基本结构;推测解码则是在自回归生成阶段进一步利用“已知候选序列可以并行计算”的系统优化。


18. 一个简化的成本反例

假设目标模型普通单 token 解码成本为 10 ms。

某轮推测解码配置为:

  • 草稿模型生成 4 个 token:8 ms;
  • 目标模型验证 4 个候选和 bonus 位置:12 ms;
  • 其他开销:2 ms;
  • 平均每轮输出:3.2 个 token。

则推测解码平均每 token 成本为:

8+12+23.2=6.875 ms\frac{8+12+2}{3.2}=6.875\text{ ms}

理论加速约为:

10/6.8751.4510/6.875\approx1.45

但如果草稿模型变慢到 15 ms:

15+12+23.2=9.06 ms\frac{15+12+2}{3.2}=9.06\text{ ms}

加速只剩:

10/9.061.1010/9.06\approx1.10

如果接受率下降,使每轮平均输出变成 2.0 个 token:

15+12+22.0=14.5 ms\frac{15+12+2}{2.0}=14.5\text{ ms}

此时推测解码反而比普通解码更慢。

这个反例说明,性能必须由完整成本模型决定。接受率、草稿长度、草稿模型延迟、目标验证延迟和运行时开销共同决定边界。


19. 最终边界

推测解码可以概括为四个相互约束的对象:

  1. 草稿模型负责快速提出候选;
  2. 接受概率决定候选是否可以直接保留;
  3. 残差分布负责在拒绝时补回目标模型缺失的概率质量;
  4. 目标模型验证保证最终输出仍由目标分布控制。

在随机采样下,严格正确性的关键等式是:

q(x)min(1,p(x)q(x))+max(0,p(x)q(x))=p(x)q(x)\min\left(1,\frac{p(x)}{q(x)}\right) + \max(0,p(x)-q(x)) = p(x)

在贪心解码下,关键是目标模型的 argmax 一致性,而不是概率接受率。

在性能方面,关键条件是:

每轮节省的目标模型单步计算>草稿生成成本+验证和状态管理开销\text{每轮节省的目标模型单步计算} > \text{草稿生成成本}+\text{验证和状态管理开销}

因此,推测解码不是“模型越大越应该使用”的通用开关,也不是只要增加草稿长度就必然加速的技巧。它是一种建立在概率校正、因果 Transformer 并行验证、KV Cache 管理和端到端成本模型之上的推理优化。只有当草稿分布足够接近目标分布、草稿模型足够便宜,并且目标模型验证确实能够高效并行化时,加速才会稳定出现。


系列导航与关联阅读

官方资料

本文依据研究论文、标准组织与主流框架官方文档重新梳理;正文、示例与工程清单由 WR BLOG 编写。