前置知识: 大语言模型

推测解码

00:00
4 min Intermediate

理解推测解码的原理和实现,通过小模型辅助加速大模型推理

推测解码

大模型生成一个 token 需要 200ms,但验证 5 个 token 只需要 220ms。推测解码利用这个不对称性——小模型猜 5 个 token,大模型一次验证,速度提升 2-3x。

类型 构建 语言 Python 前置条件 Phase 10 Lesson 12(推理优化) 预计时间 ~45 分钟

学习目标

  • 理解推测解码的数学保证输出分布与自回归完全一致
  • 实现推测解码验证和采样步骤
  • 理解接受和加速比的关系
  • 掌握 draft 模型的选择策略

算法流程

1. Draft 模型自回归生成 K 个候选 token: [t1, t2, ..., tK]
2. Target 模型并行验证:一次前向传播计算所有 K+1 个位置的 logits
3. 从左到右验证每个 token:
   - 如果 p_target(t) >= p_draft(t):接受
   - 否则:以概率 p_target(t)/p_draft(t) 接受,或从修正分布重新采样
4. 如果所有 K 个 token 都被接受,额外采样第 K+1 个 token
import torch
import torch.nn.functional as F


def speculative_decode(target_model, draft_model, tokenizer,
                       prompt, max_tokens=256, K=5, device='cuda'):
    """推测解码"""
    input_ids = tokenizer.encode(prompt, return_tensors='pt').to(device)
    generated = input_ids.clone()

    while generated.shape[1] < input_ids.shape[1] + max_tokens:
        # Step 1: Draft 模型生成 K 个候选 token
        draft_tokens = []
        draft_probs = []
        current = generated.clone()

        with torch.no_grad():
            for _ in range(K):
                outputs = draft_model(current)
                logits = outputs.logits[:, -1, :]
                probs = F.softmax(logits, dim=-1)
                token = torch.multinomial(probs, 1)
                draft_tokens.append(token)
                draft_probs.append(probs)
                current = torch.cat([current, token], dim=-1)

        # Step 2: Target 模型并行验证
        with torch.no_grad():
            outputs = target_model(current)
            target_logits = outputs.logits

        # Step 3: 逐个验证
        n_accepted = 0
        for i in range(K):
            pos = generated.shape[1] + i - 1
            target_probs = F.softmax(target_logits[:, pos, :], dim=-1)
            draft_prob = draft_probs[i].gather(-1, draft_tokens[i])

            token_id = draft_tokens[i].item()
            target_prob = target_probs[0, token_id].item()
            draft_p = draft_prob.item()

            # 接受条件
            if target_prob >= draft_p or random.random() < target_prob / max(draft_p, 1e-10):
                n_accepted += 1
                generated = torch.cat([generated, draft_tokens[i]], dim=-1)
            else:
                # 拒绝:从修正分布采样
                modified_probs = torch.clamp(target_probs - draft_probs[i], min=0)
                modified_probs = modified_probs / modified_probs.sum()
                new_token = torch.multinomial(modified_probs, 1)
                generated = torch.cat([generated, new_token], dim=-1)
                break

        # 如果全部接受,额外采样一个 token
        if n_accepted == K:
            last_pos = generated.shape[1] - 1
            last_probs = F.softmax(target_logits[:, last_pos, :], dim=-1)
            extra_token = torch.multinomial(last_probs, 1)
            generated = torch.cat([generated, extra_token], dim=-1)

    return tokenizer.decode(generated[0])

加速比分析

其中 是接受 是候选长 是 draft 模型与 target 模型的成本比。

接受K=5K=7K=10
0.62.0x2.2x2.4x
0.82.7x3.1x3.5x
0.93.0x3.6x4.2x

Draft 模型选择

策略示例接受Draft 成本
系列小模型LLaMA-7B → LLaMA-70B(0.7-0.9)
蒸馏模型Distill → Teacher
同模型早期共享 backbone最低
N-gram 模型统计模型低(0.3-0.5)极低

关键术语

术语通俗说法实际含义
推测解码”猜然后验证小模型生成候选 token,大模型并行验证的加速方法
Draft 模型”草稿模型”生成候选 token轻量级模型
接受”猜比例”大模型验证时接受 draft token比例
修正分布”纠正采样”拒绝 draft token 后,从修正概率分布重新采样

延伸阅读

知识检测

学习进度

-- 已学文档
--% 知识覆盖率

学习推荐

专注模式