推测性拒绝采样

wen IT资讯 21

本文目录导读:

推测性拒绝采样

  1. 核心问题:为什么需要推测性拒绝采样?
  2. 技术原理与步骤(类比版)
  3. 为什么它能加速?
  4. 关键保证:输出分布不变(Important Property)
  5. 对比:与普通拒绝采样
  6. 实际应用与例子(语言模型加速)
  7. 局限性

推测性拒绝采样(Speculative Rejection Sampling)是一种用于加速大型语言模型(LLM)推理的技术,属于推测性解码(Speculative Decoding) 框架的一部分,它的核心思想是:用一个更小、更快的草稿模型(Draft Model)先生成多个候选词,然后用目标大模型(Target Model)并行验证这些候选词,只接受目标模型认为合理的部分,丢弃不合理的部分。

以下是对该技术的详细解释,包括其原理、步骤、优势以及与传统解码方式的对比。


核心问题:为什么需要推测性拒绝采样?

传统自回归解码(Autoregressive Decoding)是串行的:生成第 n+1 个 Token 必须依赖前 n 个 Token 的计算结果。

  • 瓶颈:对于拥有数十亿参数的大型模型,每一步的串行计算都代价高昂(显存带宽限制),即使使用高效的批处理(Batching),生成每个 Token 的延迟仍然很高。
  • 目标:希望利用更小的草稿模型快速生成一段文本,由大模型以并行方式一次性验证这段文本,从而在保持目标模型输出质量不变的前提下,大幅降低生成延迟。

技术原理与步骤(类比版)

可以把大模型想象成一位严格的教授,草稿模型是一位聪明的助教

  1. 草稿生成(助教写草稿)
    • 草稿模型(助教)快速、自回归地生成一段长度为 K 的 Token 序列([A, B, C, D, E]),这一步很快,因为草稿模型很小。
  2. 并行验证(教授同时批改)
    • 大模型(教授)一次性接收草稿序列及其所有前缀上下文,如果草稿序列是 [A, B, C, D, E],大模型会同时计算:
      • 给定上下文 (..., A),预测下一个 Token 的概率分布 P_large(next | ... , A)
      • 给定上下文 (..., A, B),预测下一个 Token 的概率分布 P_large(next | ... , A, B)
      • ...以此类推。
    • 这一步是高度并行的,因为大模型可以一次性处理所有前缀,而不是逐步迭代。
  3. 拒绝采样(教授挑剔地接受)
    • 对于草稿模型生成的每个 TokenBCDE),大模型执行一个概率比较
      • 条件P_large(草稿Token | 上下文) 大于 P_draft(草稿Token | 上下文),则无条件接受该 Token(因为大模型更看好这个 Token)。
      • 否则:以概率 P_large(草稿Token | 上下文) / P_draft(草稿Token | 上下文) 接受该 Token。
    • 一旦某个 Token 被拒绝(例如在位置 D 被拒绝),其之后的所有 Token(E)都会被丢弃
  4. 补救与继续
    • 在第一个被拒绝的 Token 位置(例如位置 D),大模型会从其自身的概率分布 P_large(next | ... , A, B, C)采样一个 TokenX),作为真正的输出。
    • 关键点:这一步不产生额外的串行计算开销,因为大模型已经在并行验证阶段计算了这个位置的概率分布。
  5. 循环
    • Token ABCX 被作为最终的输出序列,以 X 为最新上下文,重新启动下一轮推测(草稿模型从 X 开始重新生成草稿)。

为什么它能加速?

  • 减少大模型的串行调用次数:传统方法中,生成 N 个 Token 需要 N 次大模型串行调用的前向传播,猜测性拒绝采样
    • 在理想情况下(草稿模型质量很高,大部分 Token 被接受),生成 N 个 Token 可能只需要大约 N/(K+1) 次大模型调用(因为每次验证 K 个草稿 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 参数模型(草稿模型,快)。
  • 过程
    1. 草稿模型快速生成 5 个 Token:"the quick brown fox jumps"
    2. 大模型并行验证这 5 个前缀:
      • "the"P_large 很大,接受。
      • "quick"P_large 略小,但以高概率接受。
      • "brown"P_large 非常低(因为大模型认为此时应该生成 "fast"),拒绝。
    3. 大模型从 "the quick" 上下文中采样自己的 Token:"fast"
    4. 输出当前步:"the quick fast"
    5. 草稿模型从 "the quick fast" 开始下一轮生成。
  • 结果:虽然放弃了 "brown",但大模型成功验证了前 2 个 Token("the quick"),只用了一次大模型调用就生成了 3 个 Token(比传统一步一个 Token 快 2 倍左右),如果草稿模型质量越高,被接受的 Token 序列越长,加速比就越接近 K 倍。

局限性

  1. 对草稿模型依赖性强:草稿模型与目标模型输出分布的一致性(Alignment)至关重要,如果草稿模型差(经常生成目标模型会拒绝的 Token),接受率会很低,导致加速效果不明显甚至更差(因为验证的开销可能超过收益)。
  2. 额外显存开销:需要加载两个模型到显存中。
  3. 草稿模型的质量瓶颈:草稿模型的 Token 生成仍需要串行(尽管它很快)。

推测性拒绝采样是一种无损、高效的 LLM 推理加速方法,它利用草稿模型的快速生成能力和目标模型的并行验证能力,用额外的草稿模型计算换取对大模型更少的串行调用,其核心价值在于以数学上精确的方式复用目标模型的并行计算,从而在不牺牲输出质量的前提下显著降低延迟。

抱歉,评论功能暂时关闭!