Skip to content

推测解码工程实践

推测解码(Speculative Decoding)是一种无损推理加速技术,通过小模型(草稿模型)快速生成候选 Token,再由大模型并行验证,在保证输出质量不变的前提下显著提升推理速度。

推测解码流程

基本原理

推测解码的工作流程:

  1. 草稿阶段:小模型自回归生成 K 个候选 Token
  2. 验证阶段:大模型一次前向传播验证所有候选 Token
  3. 接受/拒绝:根据大模型的概率分布决定接受哪些 Token
  4. 修正:拒绝位置之后重新采样
python
# 推测解码伪代码
def speculative_decode(draft_model, target_model, prompt, K=5):
    tokens = prompt
    while not finished:
        # 1. 草稿模型生成 K 个候选
        draft_tokens = []
        current = tokens
        for _ in range(K):
            next_token = draft_model.generate_one(current)
            draft_tokens.append(next_token)
            current = current + [next_token]

        # 2. 目标模型一次前向验证
        target_probs = target_model.forward(tokens + draft_tokens)

        # 3. 接受/拒绝判断
        accepted = 0
        for i, draft_token in enumerate(draft_tokens):
            if accept_criterion(target_probs[i], draft_token):
                accepted += 1
            else:
                # 从目标分布重新采样
                resampled = sample_from(target_probs[i])
                tokens = tokens + draft_tokens[:accepted] + [resampled]
                break
        else:
            tokens = tokens + draft_tokens

    return tokens

加速比分析

推测解码的理论加速比取决于:

  • 接受率 α:草稿 Token 被接受的概率
  • 草稿步数 K:每次推测的 Token 数
  • 速度比 β:大模型/小模型单步时间比

理论加速比 ≈ (1 + α·K) / (1 + K/β)

草稿模型接受率 αK=5 加速比K=10 加速比
7B → 70B0.72.3x2.8x
1.5B → 7B0.61.9x2.2x
70M → 7B0.41.4x1.5x

草稿模型选择

草稿模型应与目标模型同系列(如 Llama-2-7B 作为 Llama-2-70B 的草稿),这样分布更接近,接受率更高。跨系列模型组合的接受率通常较低。

vLLM 中的推测解码

python
from vllm import LLM, SamplingParams

llm = LLM(
    model="meta-llama/Llama-2-70b-hf",
    speculative_model="meta-llama/Llama-2-7b-hf",  # 草稿模型
    num_speculative_tokens=5,
    tensor_parallel_size=4,
)

params = SamplingParams(max_tokens=256, temperature=0.0)
outputs = llm.generate(["解释量子纠缠"], params)

限制条件

推测解码在 temperature=0(贪心解码)时效果最好,因为接受率最高。高温度采样会降低接受率,加速效果减弱。此外,推测解码目前仅支持自注意力模型。

相关资源

最近更新