Skip to content
kefan.life
Go back

无损投机采样的概率论视角

投机采样(Speculative Sampling)用小模型猜、大模型验,能在不改变输出分布的前提下加速推理。

其中,“不改变输出分布”是整个方法的硬约束:外界无论用什么统计检验,都无法区分”大模型自己逐词生成”和”小模型猜测 + 大模型校验”两种模式的输出。

候选经过接受、拒绝和补采后,每个 token 的最终概率必须等于大模型给出的概率。

算法过程

问题与目标

设词表中任意候选 token xx,在当前已接受前缀条件下,大模型实际用于采样的概率为 p(x)p(x),小模型实际用于采样的概率为 q(x)q(x)。

在某个位置上,小模型以概率 q(x)q(x) 采样出候选 xx,大模型决定是否接受。如果接受,xx 直接输出;如果拒绝,从调整后的分布中重新采样。

目标:让 xx 最终被输出的概率严格等于 p(x)p(x)。

执行流程

  1. Draft 阶段:小模型(draft model)自回归生成 KK 个候选 token:x1,x2,…,xKx_1, x_2, \dots, x_K。
  2. Verify 阶段:大模型(target model)对前缀 + 全部 KK 个候选做一次并行 forward,拿到每个位置上大模型自己的概率分布 pip_i。
  3. 逐位置接受/拒绝:从位置 1 开始依次判定。一旦在位置 ii 触发拒绝,丢弃 xix_i 及其后所有 draft token,在位置 ii 从调整后的分布中重新采样一个 token,本轮结束。
  4. 全部通过的 bonus:如果 KK 个 token 全部被接受,大模型的 forward 已经顺带算出了位置 K+1K+1 的分布,直接从中采样一个额外 token。

因此每轮产出最多 K+1K+1 个 token(全部接受 + bonus),最少 1 个 token(第一个就被拒绝,重采样补一个)。

为什么拒绝后必须立刻停止?一个 token 能被拒绝,前提是 p(xi)<q(xi)p(x_i) < q(x_i)(见后文,p≥qp \ge q 时接受概率为 1,不会拒绝)。而调整后的分布中 p≤qp \le q 的 token 权重为 0,不会被抽中,因此重新采样的结果必然不同于被拒绝的候选 token,导致大模型与小模型对于序列的预测发生分歧。而序列前缀一旦改变,后续的预测全都作废,此为 Causal Attention 的性质。

接受/拒绝规则

小模型在当前位置采样出一个候选 xx 后,以如下概率接受它:

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

两种情况:

实际判定:生成均匀随机数 r∼U[0,1)r \sim U[0,1),当 r<A(x)r < A(x) 时接受,否则拒绝。

概率推导

无损投机采样通过拒绝与补采,将分布 qq 调整为 pp。由于两个分布的概率总和都为 1,且 qq 相对 pp 多出的概率总量,恰好等于 qq 相对 pp 缺少的概率总量,因此可以通过重新分配精确补齐。

联合概率与接受概率

固定当前前缀,draft 以概率 q(x)q(x) 提出一个候选 xx,再以概率 A(x)A(x) 接受它。因此,xx 被提出并直接接受的联合概率为:

q(x)⋅A(x)=q(x)⋅min⁡(1, p(x)q(x))=min⁡(q(x), p(x))q(x) \cdot A(x) = q(x) \cdot \min\left(1,\ \frac{p(x)}{q(x)}\right) = \min(q(x),\ p(x))

候选可能是词表中的任意 token。把各 token 被提出并接受的联合概率相加,就得到这次抽样与验证直接接受候选的总概率:

β=∑xmin⁡(q(x), p(x))\beta = \sum_x \min(q(x),\ p(x))

注意区别 A(x)A(x) 是给定候选 xx 的接受概率,β\beta 则考虑了 draft 可能抽到的所有候选。综上,1−β1-\beta 等于触发补采的概率。

对比直接接受贡献的概率与目标 p(x)p(x):

所以,补采只需填补 p(x)>q(x)p(x)>q(x) 的 token 的概率缺口。

拒绝后的补采

情况二的候选以 p(x)/q(x)p(x)/q(x) 的概率接受,以 1−p(x)/q(x)1-p(x)/q(x) 的概率拒绝。按这个比例随机拒绝,就把 draft 提出它的概率 q(x)q(x) 降到直接输出它的概率 p(x)p(x)。

候选被拒绝后,当前位置仍需输出一个 token,因此系统从调整分布 p′p' 中重新采样,补偿情况一留下的缺口。

对每个 token,目标概率 p(x)p(x) 减去直接接受已贡献的 min⁡(q(x),p(x))\min(q(x),p(x)),就是需要补偿的缺口。词表中目标概率的总和为 1,直接接受的总概率为 β\beta,所以缺口总量为:

1−β=∑x[p(x)−min⁡(q(x),p(x))]=∑xmax⁡(0,p(x)−q(x))\begin{aligned} 1-\beta &= \sum_x \left[p(x)-\min(q(x),p(x))\right] \\ &= \sum_x \max(0,p(x)-q(x)) \end{aligned}

这也正是触发补采的概率。按缺口比例采样,将各 token 的缺口除以 1−β1-\beta,便得到总和为 1 的补采分布:

p′(x)=max⁡(0, p(x)−q(x))1−βp'(x) = \frac{\max(0,\ p(x) - q(x))}{1 - \beta}

当 p=qp=q 时,候选全部接受,无需补采。

无损证明

对于词表中的任意 token xx,最终输出它有两条路径:draft 提出它并通过验证,联合概率为 min⁡(q(x),p(x))\min(q(x),p(x));或者其他候选被拒绝后,补采抽到它,联合概率为 (1−β)p′(x)(1-\beta)p'(x)。两条路径互斥,概率相加得到:

Pfinal(x)=min⁡(q(x),p(x))+(1−β)⋅p′(x)P_{\text{final}}(x) = \min(q(x), p(x)) + (1 - \beta) \cdot p'(x)

当 β<1\beta<1 时,代入 p′p' 并化简:

Pfinal(x)=min⁡(q(x),p(x))+(1−β)max⁡(0,p(x)−q(x))1−β=min⁡(q(x),p(x))+max⁡(0,p(x)−q(x))=p(x)\begin{aligned} P_{\text{final}}(x) &= \min(q(x),p(x)) + (1-\beta)\frac{\max(0,p(x)-q(x))}{1-\beta} \\ &= \min(q(x),p(x)) + \max(0,p(x)-q(x)) \\ &= p(x) \end{aligned}

而当 β=1\beta=1 时,q=pq=p,全部接受同样得到 Pfinal(x)=p(x)P_{\text{final}}(x)=p(x)。因此,无论是否需要补采,每个 token 的最终输出概率都等于 p(x)p(x),输出分布与大模型一致。

每一步在相同前缀下的条件分布都与大模型一致,因此整个序列的联合分布也一致。

案例

假设词表为 A、B、C、D,本轮 draft 依次提出 D、A、C、B 四个候选。

生成候选与计算分布

draft 依次提出候选,并保存各位置的分布 qiq_i。随后 target 对已有前缀和四个候选做一次 forward,得到 p1,…,p4p_1,\dots,p_4 和全部通过时使用的 p5p_5。

假设前两个位置的候选 D、A 均被接受,接下来验证第三个位置。该位置的完整词表分布如下:

词表 tokenq3q_3(draft)p3p_3(target)关系
A0.20.5p>qp > q,draft 低估了 A
B0.40.1p<qp < q,draft 高估了 B
C0.30.1p<qp < q,draft 高估了 C
D0.10.3p>qp > q,draft 低估了 D

接受、拒绝与补采

从第一个候选开始,使用对应位置的概率,依次判断是否接受,注意各位置下词表的采样概率不同:

位置该位置的候选qiq_ipip_i接受概率 AiA_i随机数 rir_i结果
1D0.20.510.8接受
2A0.40.10.250.1接受
3C0.30.11/31/30.7拒绝
4B————随 C 一起丢弃

第二个候选 A 的接受概率为 0.25,本轮随机数为 0.1,因此通过。第三个候选 C 被拒绝后,保留 D、A,在第三个位置补采。

按前面的概率定义,该位置各 token 被提出并直接接受的联合概率为:

[min⁡(q3(A),p3(A))min⁡(q3(B),p3(B))min⁡(q3(C),p3(C))min⁡(q3(D),p3(D))]=[0.20.10.10.1]\begin{bmatrix} \min(q_3(\mathrm A),p_3(\mathrm A)) \\ \min(q_3(\mathrm B),p_3(\mathrm B)) \\ \min(q_3(\mathrm C),p_3(\mathrm C)) \\ \min(q_3(\mathrm D),p_3(\mathrm D)) \end{bmatrix} = \begin{bmatrix}0.2\\0.1\\0.1\\0.1\end{bmatrix}

四项相加得到 β=0.2+0.1+0.1+0.1=0.5\beta=0.2+0.1+0.1+0.1=0.5。拒绝概率为 1−β=0.51-\beta=0.5,也等于各 token 的缺口总量 0.3+0+0+0.2=0.50.3+0+0+0.2=0.5。

补采时,调整分布 p3′p'_3 为:

[p3′(A)p3′(B)p3′(C)p3′(D)]=10.5[0.3000.2]=[0.6000.4]\begin{bmatrix} p'_3(\mathrm A) \\ p'_3(\mathrm B) \\ p'_3(\mathrm C) \\ p'_3(\mathrm D) \end{bmatrix} = \frac{1}{0.5} \begin{bmatrix}0.3\\0\\0\\0.2\end{bmatrix} = \begin{bmatrix}0.6\\0\\0\\0.4\end{bmatrix}

随后从 p3′p'_3 采样。假设抽到 A,本轮输出 D、A、A 后结束。

核对输出分布

沿用前面计算的联合概率、1−β1-\beta 和 p3′p'_3,将直接接受与补采的概率贡献相加:

因此,最终输出分布与目标分布一致。

为什么必须用随机性

为什么不用确定性阈值,比如 p(x)<q(x)p(x) < q(x) 就直接拒绝?

因为确定性策略无法维持概率守恒。假设 p(x)=0.1p(x) = 0.1,q(x)=0.2q(x) = 0.2。一刀切拒绝的话,xx 在输出中的概率变成 0,大模型原本赋予的 10% 可能性被抹杀。这是有损解码,输出分布不再等于 pp。

而投机采样的做法是:小模型以 0.2 的频率提出 xx,大模型以 0.10.2=0.5\frac{0.1}{0.2} = 0.5 的概率放行。最终 xx 出现的联合概率 =0.2×0.5=0.1= 0.2 \times 0.5 = 0.1,精确等于 p(x)p(x)。

总结

无损来自概率补偿的精确性:拒绝的总概率恰好等于待补偿的缺口,按缺口比例补采,就能恢复目标分布。因此,小模型的准确性影响候选的接受概率,而输出分布的正确性由采样规则保证。


Share this post on:

Previous Post
FlashAttention 详解(V1 & V2)
Next Post
Roofline 分析:瓶颈的判定与局限