投机解码:从理论到实现

原文:Speculative Decoding: From Theory to Implementation
原文更新日期:2025 年 11 月 5 日
翻译:GPT 5.6 Terra

让我们来谈谈投机解码(Speculative Decoding)。它是现代 LLM 推理中最优雅的优化技术之一。若你想知道怎样在不牺牲输出质量的情况下,让语言模型的吞吐量提升 2 至 3 倍,那么这篇文章正适合你。

读完本文后,你将能深入理解这一算法,并具备从零实现它的能力。我们会讨论其动机、数学原理、实现细节,以及足以决定真实部署成败的细微边界情况。

问题:自回归解码受内存限制

大型语言模型生成文本时采用自回归方式:一次生成一个 token。每一步都会发生以下事情:

  • 从内存加载模型权重(数十亿参数)
  • 对所有此前的 token 计算 attention
  • 生成下一个 token 的 logits
  • 通过采样或 argmax 选出下一个 token
  • 重复上述过程

LLM

问题在于:现代 GPU 的计算能力极强,但内存访问相对较慢。每生成一个 token,你都需要把数 GB 的权重从 HBM(高带宽内存)搬运至计算单元,完成数万亿次运算,最后却只得到……一个 token。

这被称为受内存带宽限制(memory-bandwidth bound)。计算单元大部分时间都在闲置,等待数据到达。对于一个以 BF16 运行的 70B 参数模型,每个 token 都要加载约 140GB 权重。若 GPU 的内存带宽是 2TB/s,那么无论计算多快,理论上每个 token 至少需要 70ms。

那么,怎样才能更高效?

投机解码

答案就是投机解码。它的核心思想是:验证 K 个 token 所需的时间,大致与生成 1 个 token 相同。

用一个具体例子说明。假设你正在使用 GPT-2-XL,并已生成如下序列:

"The capital of France is"

标准做法:一次一个 token

通常,若要再生成 4 个 token,需要:

  1. 使用 "The capital of France is" 做一次前向传播,生成 "Paris"
  2. 使用 "The capital of France is Paris" 做一次前向传播,生成 ","
  3. 使用 "The capital of France is Paris," 做一次前向传播,生成 "which"
  4. 使用 "The capital of France is Paris, which" 做一次前向传播,生成 "is"

c62de3cd-ccdd-4bc9-8047-c090174d0ee0

  • 合计:通过 GPT-2-XL 进行 4 次前向传播。

每次前向传播都会从内存加载全部 15 亿参数,并对整个序列运行 attention。若每次耗时 100ms,总计就是 400ms。

推测做法:一次验证多个 token

现在假设有一个较小的 GPT-2 模型,可以预先猜出这 4 个 token。我们构造序列:

"The capital of France is Paris, which is"

然后用整个序列在 GPT-2-XL 上进行一次前向传播。该次传播会得到:

  • 位于 "is"(即 "France is" 之后)的 logits,用于验证 "Paris"

  • 位于 "Paris" 的 logits,用于验证 ","

  • 位于 "," 的 logits,用于验证 "which"

  • 位于 "which" 的 logits,用于验证 "is"

  • 位于 "is" 的 logits,免费得到一个 bonus token!

  • 合计:通过 GPT-2-XL 进行 1 次前向传播。

9cae210a-43df-4d92-920b-fe0ad827185f

为什么耗时几乎相同?

验证 K 个 token 时,发生的是:

  • 只加载一次模型权重(成本相同)
    无论处理的是 5 个 token 的 "The capital of France is",还是 9 个 token 的 "The capital of France is Paris, which is",加载的都是同一份 15 亿参数权重。这是昂贵部分,而且成本不变。

  • 对长度多出 K 个 token 的序列运行 attention(当 K 较小时,成本可忽略)
    对 9 个 token 而非 5 个 token 计算 attention 确实会增加少量工作,但这是计算受限而非内存受限的工作。GPU 在加载内存时有数千个核心闲置,额外计算基本可以在等待数据期间完成。对于 K=3 至 5 这样的小值,100ms 的前向传播可能只多出 5 至 10ms,几乎可以忽略。

    (而且,相信我,计算受限的问题是最理想的问题类型。)

  • 得到 K 个而非 1 个 logit 输出(在这一计算受限部分几乎免费)
    Transformer 并不只在最后一个位置输出 logits,而是在序列的每个位置都输出。平常我们会丢弃除最后一个以外的所有输出;投机解码则使用它们验证猜测。为 9 个位置而非 5 个位置计算 logits,只是多出一些并行进行的矩阵乘法,基本免费。

关键时刻:廉价猜测加快速验证

完整成本如下:

  • 4 次独立 GPT-2-XL 前向传播:400ms

  • 4 次 GPT-2 猜测:4 × 8ms = 32ms

  • 1 次 GPT-2-XL 验证传播:105ms(因多出 4 个 token 而略长)

  • 投机解码总成本:32ms + 105ms = 137ms

  • 加速比:400ms / 137ms ≈ 2.9x

这里假设所有猜测均正确。即使只有 50% 的猜测正确,结果仍然快得多。

关键洞见是:如果可以廉价地猜测 token(使用小模型),那么一次性验证它们几乎没有额外成本。你把一个顺序问题变成了部分并行的问题。

算法:生成草稿、验证、接受

下面更深入地了解投机解码算法,并尝试从零构建它。

阶段 1:生成草稿

使用一个小而快的“草稿模型”(draft model)以自回归方式生成 K 个 token。该模型应当:

  • 比目标模型小得多(参数量少 10 至 100 倍)
  • 快到生成 K 个 token 的成本低于一次目标模型前向传播
  • 与目标模型足够相似,从而能猜对部分 token

草稿生成阶段的实现如下:

def generate_draft_tokens(self, input_ids: torch.Tensor, num_tokens: int) -> Tuple[List[int], List[float]]:
    """
    Use the draft model to generate candidate tokens.

    Args:
        input_ids: Current token sequence
        num_tokens: Number of tokens to draft

    Returns:
        Tuple of (draft_tokens, draft_probabilities)
    """
    draft_tokens = []
    draft_probs = []

    current_ids = input_ids.clone()

    for _ in range(num_tokens):
        with torch.no_grad():
            outputs = self.draft_model(current_ids)
            logits = outputs.logits[0, -1, :]  # Last position
            probs = torch.softmax(logits, dim=0)

            # Sample next token
            next_token = torch.multinomial(probs, num_samples=1)
            token_id = next_token.item()

            draft_tokens.append(token_id)
            draft_probs.append(probs[token_id].item())

            # Append token for next iteration
            current_ids = torch.cat([current_ids, next_token.unsqueeze(0)], dim=1)

    return draft_tokens, draft_probs

关键细节是:我们同时保存 token 以及它们在草稿模型下的概率。接受阶段需要用到 draft_probs

阶段 2:验证

真正的关键在这里。将所有 K 个草稿 token 一次性送入目标模型,只进行一次前向传播。

# Create sequence with all draft tokens
draft_sequence = torch.cat([
    input_ids,
    torch.tensor([draft_tokens], device=self.device)
], dim=1)

# Single forward pass through target model
with torch.no_grad():
    outputs = self.target_model(draft_sequence)
    all_logits = outputs.logits[0]  # Shape: [seq_len, vocab_size]

注意,这里把完整序列(原始输入加上全部 K 个草稿 token)一次性输入目标模型。输出包含每个位置的 logits,也就是说,我们得到用于验证每个草稿 token 的概率分布。

这就是关键的效率增益:不再做 K 次独立前向传播(每个 token 一次),而是只做一次稍长的前向传播。

阶段 3:通过拒绝采样接受 token

接下来是微妙的部分:需要以一种保持目标模型精确概率分布的方式,逐个接受或拒绝草稿 token。

def verify_draft_tokens(self, input_ids: torch.Tensor, 
                       draft_tokens: List[int], 
                       draft_probs: List[float]) -> List[int]:
    """
    Verify draft tokens using the target model in a single forward pass.

    This is where the magic happens! We process all draft tokens at once
    and get probability distributions at each position.
    """
    # Create sequence with all draft tokens
    draft_sequence = torch.cat([
        input_ids,
        torch.tensor([draft_tokens], device=self.device)
    ], dim=1)

    # Single forward pass through target model
    with torch.no_grad():
        outputs = self.target_model(draft_sequence)
        all_logits = outputs.logits[0]  # Shape: [seq_len, vocab_size]

    # Verify each draft token
    accepted_tokens = []
    seq_len = input_ids.size(1)

    for i in range(len(draft_tokens)):
        # Get target model's probability distribution at this position
        position = seq_len - 1 + i
        target_probs = torch.softmax(all_logits[position], dim=0)
        target_prob = target_probs[draft_tokens[i]].item()
        draft_prob = draft_probs[i]

        # Acceptance criterion: p_target(token) / p_draft(token)
        acceptance_ratio = min(1.0, target_prob / draft_prob)

        if torch.rand(1).item() < acceptance_ratio:
            # Accept the draft token
            accepted_tokens.append(draft_tokens[i])
        else:
            # Reject and sample from adjusted distribution
            # Adjusted distribution: max(0, p_target - p_draft)
            adjusted_probs = torch.clamp(
                target_probs - torch.softmax(all_logits[position], dim=0), 
                min=0.0
            )

            if adjusted_probs.sum() > 0:
                adjusted_probs = adjusted_probs / adjusted_probs.sum()
                new_token = torch.multinomial(adjusted_probs, num_samples=1).item()
            else:
                # Fallback: sample from target distribution
                new_token = torch.multinomial(target_probs, num_samples=1).item()

            accepted_tokens.append(new_token)
            # Stop verifying remaining tokens
            break

    # Bonus token: if all drafts accepted, get one more from target model
    if len(accepted_tokens) == len(draft_tokens):
        position = seq_len - 1 + len(draft_tokens)
        bonus_probs = torch.softmax(all_logits[position], dim=0)
        bonus_token = torch.multinomial(bonus_probs, num_samples=1).item()
        accepted_tokens.append(bonus_token)

    return accepted_tokens

下面拆解接受循环中的逻辑。

数学原理:理解拒绝采样

在每个位置 i,有:

  • p_target:目标模型给该草稿 token 的概率
  • p_draft:草稿模型给该 token 的概率(此前已保存)

接受概率为:

acceptance_ratio = min(1.0, p_target / p_draft)

为什么使用这个公式?

情况 1:p_target ≥ p_draft

目标模型认为该 token 比草稿模型认为的更可能出现。因此以 1.0 的概率接受它。草稿模型只是偏保守,这没有问题。

情况 2:p_target < p_draft

草稿模型过度自信。以 p_target / p_draft 的概率接受它,以草稿模型的过度自信程度进行降权。

例如:

  • p_draft = 0.8p_target = 0.4,则以 0.4 / 0.8 = 0.5 的概率接受。
  • p_draft = 0.9p_target = 0.1,则以 0.1 / 0.9 ≈ 0.11 的概率接受。

调整后的分布

拒绝一个 token 后,不能直接从目标模型的分布采样,否则会造成偏差。应从一个调整后的分布采样:

p'(t) = max(0, p_target(t) - p_draft(t)) / Z

其中 Z 是归一化常数。该调整去除了草稿模型的“贡献”,只保留目标模型额外提供的部分。

直觉上,草稿模型已在其采样出的 token 上“用掉”了一部分概率质量;我们需要从剩余部分采样。

adjusted_probs = torch.clamp(
    target_probs - torch.softmax(all_logits[position], dim=0), 
    min=0.0
)

if adjusted_probs.sum() > 0:
    adjusted_probs = adjusted_probs / adjusted_probs.sum()
    new_token = torch.multinomial(adjusted_probs, num_samples=1).item()

该拒绝采样过程在数学上已证明:能够产生与标准自回归解码完全相同的分布。因此这不是近似方法,而是得到逐位一致的结果(在期望意义上)。

Bonus token

还有一个巧妙的优化:

# Bonus token: if all drafts accepted, get one more from target model
if len(accepted_tokens) == len(draft_tokens):
    position = seq_len - 1 + len(draft_tokens)
    bonus_probs = torch.softmax(all_logits[position], dim=0)
    bonus_token = torch.multinomial(bonus_probs, num_samples=1).item()
    accepted_tokens.append(bonus_token)

若所有 K 个草稿 token 均被接受,我们已从目标模型得到 K+1 个位置的 logits(原来的 K 个位置再加一个)。既然前向传播的成本已经支付,这个额外位置几乎免费,因此可以再从中采样一个 token。

也就是说,完美的草稿生成可以用一次目标模型前向传播的成本得到 K+1 个 token。

整合起来

完整的生成循环如下:

def generate(self, prompt: str, max_new_tokens: int = 50, 
             num_draft_tokens: int = 4, verbose: bool = True) -> str:
    """
    Generate text using speculative decoding.
    """
    input_ids = self.tokenizer.encode(prompt, return_tensors='pt').to(self.device)
    generated_tokens = 0
    iterations = 0
    total_accepted = 0

    while generated_tokens < max_new_tokens:
        iterations += 1

        # Step 1: Draft tokens
        draft_tokens, draft_probs = self.generate_draft_tokens(input_ids, num_draft_tokens)

        # Step 2: Verify drafts
        accepted_tokens = self.verify_draft_tokens(input_ids, draft_tokens, draft_probs)

        # Update statistics
        num_accepted = len(accepted_tokens)
        total_accepted += num_accepted
        generated_tokens += num_accepted

        if verbose:
            draft_text = self.tokenizer.decode(draft_tokens)
            accepted_text = self.tokenizer.decode(accepted_tokens)
            print(f"Iteration {iterations}:")
            print(f"  Drafted: {draft_text!r}")
            print(f"  Accepted: {accepted_text!r} ({num_accepted}/{len(draft_tokens)} tokens)")

        # Add accepted tokens to sequence
        input_ids = torch.cat([
            input_ids,
            torch.tensor([accepted_tokens], device=self.device)
        ], dim=1)

        if generated_tokens >= max_new_tokens:
            break

    result = self.tokenizer.decode(input_ids[0], skip_special_tokens=True)
    return result

每次迭代都会:

  • 用小模型生成 K 个草稿 token
  • 用大模型的一次前向传播验证全部 K 个 token
  • 接受 1 至 K+1 个 token
  • 重复,直至生成足够的 token

设置模型

要使用投机解码,需要来自同一模型家族的两个模型:

class SpeculativeDecoder:
    def __init__(self, draft_model_name: str, target_model_name: str):
        """
        Initialize the speculative decoder with two models.

        Args:
            draft_model_name: Name of the small, fast model (e.g., "gpt2")
            target_model_name: Name of the large, accurate model (e.g., "gpt2-xl")
        """
        print(f"Loading draft model: {draft_model_name}")
        self.draft_model = AutoModelForCausalLM.from_pretrained(draft_model_name)
        self.draft_model.eval()

        print(f"Loading target model: {target_model_name}")
        self.target_model = AutoModelForCausalLM.from_pretrained(target_model_name)
        self.target_model.eval()

        self.tokenizer = AutoTokenizer.from_pretrained(draft_model_name)
        self.tokenizer.pad_token = self.tokenizer.eos_token

        # Move to GPU if available
        self.device = "cuda" if torch.cuda.is_available() else "cpu"
        self.draft_model.to(self.device)
        self.target_model.to(self.device)

该示例使用 GPT-2(1.24 亿参数)作为草稿模型,使用 GPT-2-XL(15 亿参数)作为目标模型。草稿模型约小 12 倍,因此推理也约快 12 倍。

  • 重要:两个模型必须使用相同的 tokenizer 和词表。否则需要 token 映射逻辑,会增加复杂度。

性能分析

分析理论加速比。定义:

  • T_target:目标模型一次前向传播的耗时
  • T_draft:草稿模型一次前向传播的耗时
  • α:平均接受率,也称块效率(草稿 token 被接受的概率)
  • K:每次迭代的草稿 token 数

每次迭代耗时

T_iteration = K × T_draft + T_target

每次迭代的期望 token 数

这一项更棘手。若每个 token 独立地以概率 α 被接受,则期望接受的 token 数为:

E[tokens] = 1 + α + α² + α³ + ... + α^(K-1) + α^K
          = (1 - α^(K+1)) / (1 - α)

当 K 较大且 α 合理时,近似为:

E[tokens] ≈ 1 / (1 - α)

每个 token 的有效耗时

T_effective = (K × T_draft + T_target) / E[tokens]

示例计算

使用一组现实数值:

  • T_target = 100ms(GPU 上的 GPT-2-XL)
  • T_draft = 8ms(同一 GPU 上的 GPT-2,约快 12 倍)
  • K = 4(生成 4 个草稿 token)
  • α = 0.6(60% 接受率)
T_iteration = 4 × 8ms + 100ms = 132ms
E[tokens] = (1 - 0.6^5) / (1 - 0.6) ≈ 2.2 tokens
T_effective = 132ms / 2.2 ≈ 60ms per token
  • 加速比:100ms / 60ms ≈ 1.67x

这还没有使用 KV cache。正确管理 KV cache 后,加速比可达到 2 至 3 倍。

最优 K 值

K(草稿 token 数)的选择是一项权衡:

  • 太小:未能充分利用验证传播。
  • 太大:会把时间花在随后将被拒绝的 token 的草稿生成上。

最优 K 取决于:

  • 模型之间的速度比 T_target / T_draft
  • 接受率 α
  • 内存约束

实践中,K=3K=5 对大多数场景效果良好。

实际中的接受率

哪些因素决定接受率?主要有:

  • 模型相似度:来自同一模型家族的草稿模型和目标模型(如 GPT-2 → GPT-2-XL),比不匹配的模型具有更高接受率。

任务难度

  • 可预测文本(新闻、文档):α ≈ 0.7

  • 创意写作:α ≈ 0.4-0.5

  • 代码生成:α ≈ 0.5-0.6

  • 上下文长度:上下文越长,通常包含的信息越多,token 越容易预测,接受率也往往更高。

  • Temperature:较低的 temperature(更贪心)通常会带来更高的接受率,因为两个模型都会收敛于明显的选择。

运行代码时会看到类似输出:

Iteration 1:
  Drafted: ' a topic that'
  Accepted: ' a topic' (2/4 tokens)

Iteration 2:
  Drafted: ' that has been'
  Accepted: ' that has been' (4/4 tokens)

Iteration 3:
  Drafted: ' discussed extensively'
  Accepted: ' discussed' (1/4 tokens)

这种波动很正常:有的迭代会全部接受,有的只会接受一个 token。

实现注意事项

内存需求

必须同时在内存中保留两个模型。对于 GPT-2/GPT-2-XL:

  • GPT-2:约 500MB
  • GPT-2-XL:约 6GB
  • 合计:约 6.5GB

对于 LLaMA-7B(草稿模型)和 LLaMA-70B(目标)这类更大模型,需要约 14GB + 140GB ≈ 150GB+,通常需要多张 GPU。

KV Cache 支持

上述实现未使用 KV cache,这意味着每一步都会重新计算所有此前 token 的 attention。加入 KV cache 支持可以显著改善性能:

  • 草稿模型缓存:为草稿生成过程维护单独的 KV cache。
  • 目标模型缓存:只缓存已经验证的前缀。
  • 缓存失效:拒绝 token 时,将缓存截断至拒绝位置。

使用 KV cache 后,可以获得 2 至 3 倍加速,而不是约 1.5 至 2 倍。

Batch size 的考量

在 batch 推理中使用投机解码较为棘手。batch 中不同序列可能接受不同数量的 token,从而产生不规则张量(ragged tensors)。

可选方案:

  • 独立处理:逐条处理序列,较简单但效率较低。
  • Padding:补齐到最大接受长度,但会在 padding 上浪费计算。
  • 动态 batching:使用能高效处理可变长度序列的框架。

大多数生产系统为了简洁会使用独立处理。

贪心解码与采样

以上实现使用采样(torch.multinomial)。对于贪心解码,可以简化为:

# Greedy acceptance: just check if draft token matches argmax
target_token = torch.argmax(target_probs)
if draft_tokens[i] == target_token:
    accepted_tokens.append(draft_tokens[i])
else:
    accepted_tokens.append(target_token)
    break

贪心解码通常具有更高接受率,因为两个模型都会收敛到相同的明显选择。

投机解码擅长的场景

并非所有场景都能同等受益。以下是它表现出色的场景:

  • 大型目标模型加小型草稿模型(体量差距大于 10 倍):速度差越大,收益越高。
  • 中高接受率(大于 40%):若草稿模型太差,额外开销会占主导。
  • 单用户或小 batch 推理:比大 batch 更容易管理。
  • 中等长度上下文(小于 8K tokens):超长上下文会提高验证成本。
  • 可预测任务:问答、摘要、翻译通常比创意写作效果更好。

投机解码难以发挥作用的场景

  • 小型目标模型:若目标模型本来就很快(小于 20ms/token),绝对收益很小。
  • 低接受率(小于 30%):草稿生成开销会超过收益。
  • 极长上下文(大于 32K tokens):验证传播变得昂贵。
  • 高度创意的任务:低困惑度意味着更低可预测性,接受率更低。
  • 内存受限环境:同时保存两个模型是一种奢侈。

总结

投机解码很好地展示了如何将一个顺序问题转变为部分并行问题。关键洞见包括:

  • LLM 推理的瓶颈是内存带宽,而非计算能力。
  • 验证几乎免费,因为无论 K 是多少都只需要一次前向传播。
  • 拒绝采样保持分布的正确性,因此能得到精确结果。
  • 大多数 token 可预测,使草稿模型出人意料地有效。

实现时需要谨慎处理:

  • 为保持正确性所需的拒绝采样数学原理
  • 草稿模型和目标模型的概率跟踪
  • 拒绝 token 时的提前停止
  • 为提升效率提取 bonus token

最终结果是:在真实工作负载上实现 1.5 至 3 倍加速,且质量完全不下降。对于一种概念上只是“先猜再检查”的技术而言,这相当不错。

若要大规模部署 LLM,投机解码应当成为优化工具箱的一部分。它的数学细节很微妙,但回报是真实的。


完整可运行代码见:github.com/jaygala223/scratch

可自行尝试:

python3 speculative_decoding.py

观察你的 LLM 如何在不牺牲任何输出质量的前提下,更快地生成文本。

posted @ 2026-08-24 22:06  顾北清  阅读(11)  评论(0)    收藏  举报