投机解码:无损加速大模型生成

大模型的自回归生成有一个根深蒂固的矛盾:模型越强、参数越多,每吐出一个 token 的成本就越高,但生成过程却被迫一个 token 一个 token 地串行推进。投机解码(speculative decoding)正是为化解这一矛盾而生。它用一个便宜的草稿模型先并行猜出一串 token,再交给昂贵的大模型一次性并行验收,在保证输出分布与原模型完全一致的前提下,把多次串行前向压缩成少数几次并行前向。

自回归生成的瓶颈

标准 Transformer 解码是逐 token 串行的。假设要生成 K 个 token,就要把整个模型(自注意力与多层前馈)连续跑 K 次,而每次前向只为得到下一个 token。问题在于,现代 LLM 的瓶颈往往不是算力(FLOPs),而是访存:每一步都要把数十亿甚至上千亿参数从显存(HBM)搬到计算单元。参数规模越大,单次前向的「搬运」开销越高,而这一步只产出 1 个 token,计算利用率极低。

这种「memory-bound」特性意味着:即便模型每次能同时算出多个位置的概率,自回归的因果约束仍强制我们一次只看一个 token。投机解码的核心洞察是,如果我们能廉价地预判接下来几个 token 大概是什么,就可以把多次串行前向合并成一次并行前向,从而摊薄访存成本。

投机解码的核心思想

投机解码引入两个角色:

  • 草稿模型(draft model):一个明显更小、更快的模型,例如层级裁剪版或蒸馏版。它基于当前前缀,一次并行猜出 k 个候选 token。
  • 目标模型(target model):即我们真正想要的大模型。它对草稿给出的 k 个位置做一次并行前向,得到每个位置的概率分布,并逐位置决定接受或拒绝。

直觉上,草稿模型虽弱,但它在「简单」位置上往往猜得不错;大模型只需做一次「判断题」式的验证,而不是亲自把每个 token 都生成一遍。验证是并行的,因此当草稿猜对的比例越高,省下的串行步数就越多。

下面用一段文字示意单轮投机解码的流程:

给定当前前缀
  草稿模型并行生成 k 个候选 token(记为 x1, x2, ..., xk)
  目标大模型对这 k 个位置做一次并行前向,得到每个位置的概率
  目标模型逐位置验收:按概率比随机接受或拒绝
  若第 i 个 token 被拒,则在该位置由目标模型从修正分布重采样一个 token,本轮结束
  若 k 个候选全部接受,目标模型再额外生成 1 个 token
最终得到的序列,与纯目标模型自回归逐 token 生成的结果同分布

验收准则:拒绝采样保证无损

投机解码最精妙的地方在于「无损」:最终输出与直接用目标模型采样(或贪心、温度采样等既定策略)在统计上完全一致,加速不以任何质量为代价。这靠的是基于概率比的拒绝采样(rejection sampling)。

设目标模型在第 i 个位置的概率为 p_i,草稿模型给出的概率为 q_i。对草稿提出的 token x,验收规则为:

  • 以概率 min(1, p_i(x) / q_i(x)) 接受该 token;
  • 若拒绝,则在第 i 个位置改从「修正分布」采样:其概率正比于 max(0, p_i(x) - q_i(x)),再做归一化。

这一构造的妙处在于,无论草稿在哪些位置猜对、在哪些位置猜错,目标模型在每个位置上边际采样到的 token,其分布恰好等于 p_i。接受保证了「草稿猜对的地方沿用草稿」,拒绝时的修正分布保证了「草稿猜错的地方由目标模型补回正确的分布」。因此整条序列的联合分布与纯目标模型自回归生成严格一致。

用伪代码可表示为:

def speculative_decode(prefix, draft, target, k):
    # draft: 小模型  target: 大模型  k: 投机长度
    guesses = draft.generate(prefix, n=k)          # 草稿模型并行猜 k 个 token
    target_logits = target.forward(prefix + guesses)  # 一次并行前向拿各位置概率
    accepted = []
    for i in range(k):
        p = target_prob(target_logits[i], guesses[i])   # 目标模型概率
        q = draft_prob(draft, prefix + accepted, guesses[i])  # 草稿模型概率
        ratio = p / q
        if random_uniform() < min(1.0, ratio):
            accepted.append(guesses[i])                 # 接受该 token
        else:
            adjusted = normalize(maximum(p - q, 0))     # 拒绝:修正分布重采样
            token = sample(adjusted)
            return prefix + accepted + [token]
    extra = target.sample_one(prefix + accepted)         # 全部接受则额外生成 1 个
    return prefix + accepted + [extra]

自投机与 Medusa

标准投机解码需要一个与目标模型任务相关、且明显更快的独立草稿模型,但这在工程上并不总是现成可得。于是出现了「免独立小模型」的变体:

  • 自投机(self-speculative / blockwise):由同一大模型自身承担草案。思路之一是 blockwise parallel decoding,即让大模型先用一个更轻的「草拟模式」(例如跳过部分层、或使用更早的表示)并行猜出若干 token,再由完整模型验收。这样无需维护第二个模型。
  • Medusa:在大模型主干之上额外挂若干个「解码头」(Medusa heads),每个头负责预测未来第 1、2、…、m 个位置的 token。配合树状注意力,多个候选续写可被一次性并行验证,进一步提升接受长度与加速比。

这些变体的共同点是:草案来源由「外部小模型」变成「大模型自身或其附属结构」,从而降低部署复杂度,但核心验收逻辑仍是同一套拒绝采样。

与朴素并行解码、KV 缓存的关系

需要区分投机解码与「朴素并行解码」。朴素并行解码直接让模型一次预测未来多个位置,但自回归的因果性决定这些预测彼此并不独立,且无法保证分布正确。投机解码的不同之处在于「先猜后验」:草稿可以 любой 来源,关键是目标模型用拒绝采样把分布校正回 p,因此做到无损。

KV 缓存则是投机解码的好搭档。验证阶段目标模型对 k 个位置做并行前向,这些位置共享同一前缀的 KV 缓存,不必重复计算前缀的注意力,从而让并行验证高效落地。换句话说,投机解码借用了 KV 缓存来低成本地展开 k 个候选位置的验证。

适用场景与前提

投机解码的收益并非无条件,需要满足若干前提:

  • 草稿与目标模型任务相关:草稿必须在目标模型要服务的分布上「猜得准」。若二者领域差异大,草稿接受率骤降,加速比趋近于 1,甚至因多跑了草稿前向而变慢。
  • 加速比受接受长度限制:每轮平均接受的 token 数(含可能的额外 1 个)决定提速倍数。理想情况下约为 k 量级,但难预测、长尾或高熵位置会降低接受率。
  • 对长尾与困难位置收益下降:在需要「深思」的推理、计数、强约束生成中,草稿往往频繁猜错,加速比收窄。
  • 小批量或单请求场景收益更明显:在已做大规模连续批处理的系统里,单请求的解码瓶颈被部分掩盖,投机解码的边际收益相对降低,但二者并不冲突。

与其他加速手段的互补

投机解码是一种「算法层」的并行化加速,与多种系统层、数值层手段正交且可叠加:

  • 量化(INT8/FP8 等):降低大模型与草稿模型的访存与算力开销,和投机解码分别作用于不同维度,可同时启用。
  • 连续批处理(如 vLLM):提升 GPU 在请求间的利用率,与投机解码「单请求内并行验证」互补,共同压低端到端延迟。
  • 更短的草稿与更激进的 k:构成可调的超参,需结合实际负载在「加速比」与「每轮开销」间权衡。

在工程落地时,通常先把量化、批处理等基础加速做扎实,再叠加投机解码以进一步压缩解码步数。

小结

投机解码用一个便宜的草稿模型并行猜测若干 token,再交给大模型一次性并行验收,借助基于概率比的拒绝采样保证输出分布与原模型严格一致,因此是一种真正无损的加速。其收益取决于草稿的接受长度与任务相关度,且能与量化、连续批处理等手段互补叠加。自投机与 Medusa 等变体进一步免去独立小模型,让这一思路更易部署。

参考与延伸阅读

  • arXiv:2211.17192 —— Fast Inference from Transformers via Speculative Decoding(Leviathan, Kalman, Matias,ICML 2023)。投机解码原作,提出拒绝采样验收准则并证明其无损性。已核验。
  • arXiv:1811.03115 —— Blockwise Parallel Decoding for Deep Autoregressive Models(Stern, Shazeer, Uszkoreit,NeurIPS 2018)。自投机/分块并行解码的思想源头。已核验。注:常见资料中误写的 2212.15714 在 arXiv 上并不存在,正确编号为 1811.03115。
  • arXiv:2401.10774 —— Medusa: Simple LLM Inference Acceleration Framework with Multiple Decoding Heads(Cai et al.,2024)。在大模型上挂载多解码头做树状并行验证的自投机方案。已核验。
本文累计阅读