本文目录导读:

推测性拒绝采样(Speculative Rejection Sampling)是一种用于加速大型语言模型(LLM)推理的技术,属于推测性解码(Speculative Decoding) 框架的一部分,它的核心思想是:用一个更小、更快的草稿模型(Draft Model)先生成多个候选词,然后用目标大模型(Target Model)并行验证这些候选词,只接受目标模型认为合理的部分,丢弃不合理的部分。
以下是对该技术的详细解释,包括其原理、步骤、优势以及与传统解码方式的对比。
核心问题:为什么需要推测性拒绝采样?
传统自回归解码(Autoregressive Decoding)是串行的:生成第 n+1 个 Token 必须依赖前 n 个 Token 的计算结果。
- 瓶颈:对于拥有数十亿参数的大型模型,每一步的串行计算都代价高昂(显存带宽限制),即使使用高效的批处理(Batching),生成每个 Token 的延迟仍然很高。
- 目标:希望利用更小的草稿模型快速生成一段文本,由大模型以并行方式一次性验证这段文本,从而在保持目标模型输出质量不变的前提下,大幅降低生成延迟。
技术原理与步骤(类比版)
可以把大模型想象成一位严格的教授,草稿模型是一位聪明的助教。
- 草稿生成(助教写草稿):
- 草稿模型(助教)快速、自回归地生成一段长度为
K的 Token 序列([A, B, C, D, E]),这一步很快,因为草稿模型很小。
- 草稿模型(助教)快速、自回归地生成一段长度为
- 并行验证(教授同时批改):
- 大模型(教授)一次性接收草稿序列及其所有前缀上下文,如果草稿序列是
[A, B, C, D, E],大模型会同时计算:- 给定上下文
(..., A),预测下一个 Token 的概率分布P_large(next | ... , A)。 - 给定上下文
(..., A, B),预测下一个 Token 的概率分布P_large(next | ... , A, B)。 - ...以此类推。
- 给定上下文
- 这一步是高度并行的,因为大模型可以一次性处理所有前缀,而不是逐步迭代。
- 大模型(教授)一次性接收草稿序列及其所有前缀上下文,如果草稿序列是
- 拒绝采样(教授挑剔地接受):
- 对于草稿模型生成的每个 Token(
B、C、D、E),大模型执行一个概率比较:- 条件:
P_large(草稿Token | 上下文)大于P_draft(草稿Token | 上下文),则无条件接受该 Token(因为大模型更看好这个 Token)。 - 否则:以概率
P_large(草稿Token | 上下文) / P_draft(草稿Token | 上下文)接受该 Token。
- 条件:
- 一旦某个 Token 被拒绝(例如在位置
D被拒绝),其之后的所有 Token(E)都会被丢弃。
- 对于草稿模型生成的每个 Token(
- 补救与继续:
- 在第一个被拒绝的 Token 位置(例如位置
D),大模型会从其自身的概率分布P_large(next | ... , A, B, C)中采样一个 Token(X),作为真正的输出。 - 关键点:这一步不产生额外的串行计算开销,因为大模型已经在并行验证阶段计算了这个位置的概率分布。
- 在第一个被拒绝的 Token 位置(例如位置
- 循环:
- Token
A、B、C、X被作为最终的输出序列,以X为最新上下文,重新启动下一轮推测(草稿模型从X开始重新生成草稿)。
- Token
为什么它能加速?
- 减少大模型的串行调用次数:传统方法中,生成
N个 Token 需要N次大模型串行调用的前向传播,猜测性拒绝采样- 在理想情况下(草稿模型质量很高,大部分 Token 被接受),生成
N个 Token 可能只需要大约N/(K+1)次大模型调用(因为每次验证K个草稿 Token 只需要一次大模型前向传播)。
- 在理想情况下(草稿模型质量很高,大部分 Token 被接受),生成
- 利用并行计算优势:现代 GPU 非常擅长在批处理中处理多个序列,并行验证步骤将
K个不同的前缀上下文打包成一个 batch,一次前向传播即可完成所有计算,远比K次独立的串行前向传播要快。
关键保证:输出分布不变(Important Property)
推测性拒绝采样有一个非常优雅的性质:它保证生成的最终 Token 序列的分布,与直接使用目标大模型(没有任何推测)的分布完全相同。
- 这不是近似加速,而是精确等价的加速。
- 原因在于其精确的拒绝采样过程(基于概率比值),在数学上确保了从目标分布中无偏差地采样,如果草稿模型给出的概率与目标模型不匹配,拒绝采样机制会正确地对采样分布进行校正。
对比:与普通拒绝采样
| 特性 | 普通拒绝采样 (Rejection Sampling) | 推测性拒绝采样 (Speculative Rejection Sampling) |
|---|---|---|
| 模型 | 单一模型 | 双模型(草稿模型 + 目标模型) |
| 过程 | 从提议分布采样,以概率 P_target / P_proposal 接受,如果拒绝,重试。 |
草稿模型生成序列,目标模型并行验证序列,如果拒绝一个 Token,丢弃其后所有 Token,并立即在该位置采样。 |
| 计算代价 | 每次尝试都需要完整的目标模型计算(串行)。 | 一次目标模型计算验证 K 个 Token(并行),显著降低了串行成本。 |
| 主要目的 | 从复杂分布中采样 | 加速 LLM 推理(延迟降低) |
实际应用与例子(语言模型加速)
- 场景:使用 GPT-4(目标模型,巨大,慢)与一个微调过的 7B 参数模型(草稿模型,快)。
- 过程:
- 草稿模型快速生成 5 个 Token:
"the quick brown fox jumps" - 大模型并行验证这 5 个前缀:
"the"→P_large很大,接受。"quick"→P_large略小,但以高概率接受。"brown"→P_large非常低(因为大模型认为此时应该生成"fast"),拒绝。
- 大模型从
"the quick"上下文中采样自己的 Token:"fast"。 - 输出当前步:
"the quick fast"。 - 草稿模型从
"the quick fast"开始下一轮生成。
- 草稿模型快速生成 5 个 Token:
- 结果:虽然放弃了
"brown",但大模型成功验证了前 2 个 Token("the quick"),只用了一次大模型调用就生成了 3 个 Token(比传统一步一个 Token 快 2 倍左右),如果草稿模型质量越高,被接受的 Token 序列越长,加速比就越接近K倍。
局限性
- 对草稿模型依赖性强:草稿模型与目标模型输出分布的一致性(Alignment)至关重要,如果草稿模型差(经常生成目标模型会拒绝的 Token),接受率会很低,导致加速效果不明显甚至更差(因为验证的开销可能超过收益)。
- 额外显存开销:需要加载两个模型到显存中。
- 草稿模型的质量瓶颈:草稿模型的 Token 生成仍需要串行(尽管它很快)。
推测性拒绝采样是一种无损、高效的 LLM 推理加速方法,它利用草稿模型的快速生成能力和目标模型的并行验证能力,用额外的草稿模型计算换取对大模型更少的串行调用,其核心价值在于以数学上精确的方式复用目标模型的并行计算,从而在不牺牲输出质量的前提下显著降低延迟。