编程 深度长文:LLM 推测解码(Speculative Decoding)工程化实战——从原理到3倍加速的完整实现

2026-07-26 17:15:22 +0800 CST views 7

深度长文:LLM 推测解码(Speculative Decoding)工程化实战——从原理到 3 倍加速的完整实现

一、背景:自回归解码的「慢」是一个系统性问题

2026 年的今天,大语言模型早已不再停留在聊天玩具的阶段。ChatGPT Work、Claude Cowork、GitHub Copilot Agent——AI 正在真实地嵌入开发者的日常流水线。但无论模型多么聪明,推理速度始终是一道绕不过去的坎。

你在终端里敲下一条指令,光标闪烁 2 秒才开始吐字,然后一个字一个字往外蹦——这种体验用久了会让人产生一种错觉:「AI 就是慢的,习惯了就好。」

但这个「慢」真的无可避免吗?

1.1 量化那个「慢」

我们来算一笔账。以 Llama-3-70B 为例,模型参数约 140GB(FP16),而一块 NVIDIA A100-80G 的显存带宽大约是 2TB/s。这意味着每生成一个 token,你至少需要把全部参数从显存搬运到计算单元一次——耗时约 70ms。这还没算上 KV Cache 的读写和 Attention 计算的开销。

单次 token 生成 = 参数加载          ~70ms
                + KV Cache 读取    ~15ms
                + Attention 计算   ~10ms
                + FFN 前向         ~20ms
                + KV Cache 写入    ~5ms
                = 约 120ms
每秒 ≈ 8 tokens

如果你的产品经理告诉你:「用户需要首 token 延迟 < 500ms,生成速度 > 30 tokens/s」,靠传统自回归解码,即使把模型量化到 INT4,70B 模型也顶多拉到 20 token/s——差距仍然巨大。

但这还不是最要命的。更大的痛点是:内存带宽是铁桶里的水龙头,瓶颈根本不在算力(FLOPs),而在数据搬运(Memory Bound)

这里有个非常反直觉的真相:运行一个 70B 模型的 GPU,其计算单元利用率可能还不到 20%。剩下的 80% 时间在干什么?在等数据。每次生成一个 token,都要把 140GB 的参数从 HBM(高带宽显存)搬到 SM(流式多处理器),这个过程就是所谓的「内存墙」(Memory Wall)。

1.2 为什么 KV Cache 帮不了这个忙

有人在想:既然 Attention 的 KV Cache 能复用,那能不能扩大它的范围,减少参数加载次数?

答案是不能。KV Cache 只缓存了 Attention 层的 Key 和 Value,占用的显存虽然大(4K 上下文中约 2-4GB),但占总推理时间的比例相对较小。参数量才是带宽瓶颈的主宰。不管上下文多长,每生成一个 token,你都需要把整个模型的所有参数加载一遍——因为每个 token 的计算路径都可能不同,FFN 层的权重无法缓存。

Transformer Decoder 单步计算流:
1. Input Embedding + Position Encoding
2. Self-Attention (利用 KV Cache,只算当前 token 的 Q)
3. Residual + LayerNorm
4. FFN (两个全连接层,权重要全部加载!)
5. Residual + LayerNorm
6. LM Head (大词汇表投影)

其中 Step 4 的参数量占整个模型的 2/3

1.3 现有方案的局限

工业界解决推理性能有几种常用手段,但各有代价:

  • 模型量化(INT4/INT8):参数体积直接减半或减到四分之一,但精度有损,对数学推理和代码生成任务的影响不可忽视
  • KV Cache 量化:只解决 Attention 部分的显存和带宽,对 FFN 部分无影响,加速幅度有限
  • 批量推理(Batching):通过增大 batch size 摊薄每 token 的固定开销,但单用户场景下 batch 无法做大
  • 结构剪枝/MoE 激活:直接减少参数量,但模型结构发生了变化,需要重新训练或微调

有没有一种方案,不修改模型、不损失精度、只靠优化计算流程就能拿到 2-3 倍加速

这就是 Speculative Decoding 的核心价值主张。


二、核心原理:猜-验范式

2.1 直觉:为什么「猜」比「算」快?

想象两个人在写代码:

  • 主模型:是一位资深架构师,代码写得稳,但每次写一行都要思考半天,读几十页参考文档。
  • 草稿模型:是一位刚入职的实习生,手速飞快,但代码质量一般。

如果没有 Speculative Decoding,架构师自己一行一行写,虽然每行都对,但效率极低。

有了 Speculative Decoding,流程变成:

  1. 实习生先唰唰唰写出 5 行代码(草稿)
  2. 架构师扫一眼,把每行标记为「✓」或「✗」
  3. 从第一个错的开始,架构师自己写一行
  4. 重复

关键点在于:架构师「检查」5 行代码的时间,和「自己写」1 行的时间几乎一样。因为 GPU 计算是高度并行的——一次 forward pass 处理 5 个 token 的验证,和一次 forward pass 生成 1 个 token,时间成本差不多。

这就是 Speculative Decoding 的底层逻辑:利用 GPU 的并行计算能力,用一次推理的代价验证多个候选 token。

2.2 数学本质:拒绝采样的并行化

Speculative Decoding 的核心算法是 拒绝采样(Rejection Sampling)的并行化版本

设目标模型为 $p(x)$(大模型,精度高但慢),草稿模型为 $q(x)$(小模型,快但精度低)。

每一步,草稿模型 $q$ 自回归地生成 $\gamma$ 个候选 token:
$$
\hat{x}_1, \hat{x}2, ..., \hat{x}\gamma \sim q(\cdot | \text{context})
$$

然后目标模型 $p$ 用一次 forward pass 计算这 $\gamma$ 个 token 的概率分布:
$$
p(\cdot | \text{context}), p(\cdot | \text{context}, \hat{x}_1), ..., p(\cdot | \text{context}, \hat{x}1, ..., \hat{x}{\gamma-1})
$$

对于每个候选 token $\hat{x}_t$,以概率 $\min(1, \frac{p(\hat{x}_t | \cdot)}{q(\hat{x}_t | \cdot)})$ 接受它。如果被拒绝,则从调整后的分布 $p'(\cdot) = \text{norm}(\max(0, p(\cdot) - q(\cdot)))$ 中采样一个 token 填补,并停止本轮验证。

这个过程的数学保证是:最终输出分布与纯目标模型 $p(x)$ 完全一致。

所以它不是「近似加速」,而是「无损加速」——这点非常关键。如果你的 CTO 怀疑加速会影响质量,Speculative Decoding 的答案很明确:数学上保证一致。

2.3 为什么要并行验证

传统自回归解码为什么慢?因为它每一串行步骤都要重新加载模型权重:

Step 1: Load weights (70ms) → Compute (20ms) → Emit token
Step 2: Load weights (70ms) → Compute (20ms) → Emit token
...
Step N: Load weights (70ms) → Compute (20ms) → Emit token

瓶颈在于每次都要花 70ms 加载参数,真正的计算只占一小部分。如果用个比喻:每次只端一盘菜上餐桌,但为了端这一盘菜,你得从仓库走到厨房再走回来——90% 的时间花在了走路上,真正做菜的时间只有 10%。

Speculative Decoding 的并行验证则不同:

Draft:    q(t1) → q(t2) → q(t3) → q(t4) → q(t5)   (小模型快,总耗时 ≈ 20ms)
Verify:   p(t1,t2,t3,t4,t5)                          (一次推理,耗时 ≈ 90ms)

总耗时 ≈ 110ms,而传统方案需要 5 × 90ms = 450ms。关键在于:草稿模型参数只有目标模型的 1/10 甚至更少,每次生成 token 只需要加载约 15GB 的权重(8B 模型),而不是 140GB。

2.4 接受率:决定加速效果的核心指标

加速效果最终由接受率决定。接受率是指草稿生成的 token 被主模型验证通过的比率。

假设草稿阶段耗时 $T_d$(生成 $\gamma$ 个 token),验证阶段耗时 $T_v$(和生成 1 个 token 差不多),接受率为 $\alpha$,那么平均每个验证周期获得的有效 token 数约为:

$$
E[\text{tokens}] = 1 + \sum_{k=1}^{\gamma} \alpha^k = \frac{1 - \alpha^{\gamma+1}}{1 - \alpha}
$$

加速比就是:

$$
\text{Speedup} = \frac{E[\text{tokens}] \cdot T_{\text{baseline}}}{T_d + T_v}
$$

其中 $T_{\text{baseline}}$ 是传统方式生成 1 个 token 的时间。

来算个具体的:草稿模型速度是目标模型的 10 倍($T_d \approx T_v / 10$),接受率 $\alpha = 0.7$,$\gamma = 6$:

$$
E[\text{tokens}] = \frac{1 - 0.7^7}{1 - 0.7} = \frac{1 - 0.082}{0.3} \approx 3.06
$$

$$
\text{Speedup} = \frac{3.06 \cdot T_v}{0.1T_v + T_v} = \frac{3.06}{1.1} \approx 2.78x
$$

这就是 2-3 倍加速的来源。如果接受率降到 0.5,加速比变成约 1.8x;如果提高到 0.85,加速比可达 3.5x 以上。


三、草稿模型选择:工程上的核心决策

Speculative Decoding 好不好用,核心看草稿模型选得对不对。这是我花时间最多的地方,也是生产环境最容易踩坑的环节。

3.1 草稿模型的核心约束

草稿模型 $q$ 与目标模型 $p$ 之间有三个关键指标:

  • 接受率(Acceptance Rate):草稿生成的 token 被主模型接受的概率,决定了每个验证周期内获得的有效 token 数。接受率主要取决于草稿模型和目标模型分布的对齐程度。
  • 草稿速度比(Draft Speed Ratio):$r = \frac{\text{草稿模型 token/s}}{\text{目标模型 token/s}}$,决定了草稿阶段的损耗占比。
  • 显存开销:额外加载一份草稿模型需要消耗额外的显存,这可能是最现实的约束。

这三个指标构成一个不可能三角——没有草稿模型能同时做到接受率高、速度快、显存占用小。需要在三者之间做出权衡。

3.2 常见的草稿模型组合策略

策略草稿模型目标模型接受率加速比额外显存
同族降级Phi-3-mini (3.8B)Llama-3-8B65-75%2.0-2.5x~7.5GB
同族降级Llama-3-8B (INT4)Llama-3-70B55-65%1.8-2.2x~5GB
自草稿目标模型自身 (KV 投机)目标模型50-60%1.5-2.0x~0GB
MedusaN/A (多个预测头)目标模型60-70%2.0-2.5x~0.5GB
检索增强n-gram 缓存目标模型30-45%1.2-1.5x~0GB
跨族妥协Qwen2-1.5BLlama-3-70B30-45%1.3-1.6x~3GB

最推荐的策略是同族降级——在相同的模型家族内选择一个小的 checkpoint 作为草稿。比如 Llama-3-70B 配 Llama-3-8B,或 Qwen2-72B 配 Qwen2-7B。好处是共享 tokenizer,共享预训练数据分布,接受率天然高。实测中,同族草稿模型的接受率通常在 60-70% 之间,跨家族则断崖式跌到 30-45%。

3.3 自草稿(Self-Speculative Decoding)

一个巧妙的变体:使用目标模型自身的浅层或量化版本作为「草稿」

具体做法是:将目标模型的前 N 层(通常取全部,但用更低精度)用作草稿生成。这样两个模型共享大部分权重的显存映射,额外开销极小。

from transformers import AutoModelForCausalLM, BitsAndBytesConfig
import torch

# 目标模型:FP16
target_model = AutoModelForCausalLM.from_pretrained(
    "meta-llama/Llama-3.1-70B",
    torch_dtype=torch.float16,
    device_map="auto",
)

# 草稿模型:同一模型的 INT4 量化版本
draft_model = AutoModelForCausalLM.from_pretrained(
    "meta-llama/Llama-3.1-70B",
    quantization_config=BitsAndBytesConfig(
        load_in_4bit=True,
        bnb_4bit_compute_dtype=torch.float16,
    ),
    device_map="auto",
)

print(f"Target VRAM: {target_model.get_memory_footprint() / 1e9:.1f} GB")
print(f"Draft VRAM:   {draft_model.get_memory_footprint() / 1e9:.1f} GB")
print(f"Total:        {(target_model.get_memory_footprint() + draft_model.get_memory_footprint()) / 1e9:.1f} GB")

目标模型约 140GB (FP16),草稿模型约 35GB (INT4),合计约 175GB——刚好可以塞进 2×A100-80G。加速比实测约 1.8x,虽然没有独立小模型那么夸张,但好处是:

  1. 不引入额外的模型依赖
  2. 两个模型的分布天然一致,接受率有保证
  3. KV Cache 可以部分共享(通过 past_key_values 传递)

3.4 草稿模型训练:要不要微调对齐

一个生产级的问题:要不要针对特定推理场景微调草稿模型,使它更贴合目标模型的输出分布?

答案是:

  • 对于代码生成场景,建议微调。代码的输出模式高度结构化(函数签名、类型注解、括号配对),草稿模型如果没对齐,接受率会显著低于理论值。
  • 对于通用对话场景,一般不用。通用对话的多样性使得对齐收益有限。

下面是一个简单的草稿模型微调脚本(LoRA):

from peft import LoraConfig, get_peft_model
from datasets import load_dataset

# 用目标模型生成的输出作为训练数据
dataset = load_dataset("json", data_files="target_model_outputs.jsonl")

# LoRA 微调草稿模型
lora_config = LoraConfig(
    r=16,
    lora_alpha=32,
    target_modules=["q_proj", "k_proj", "v_proj", "o_proj"],
    lora_dropout=0.05,
    bias="none",
    task_type="CAUSAL_LM",
)

draft_model = AutoModelForCausalLM.from_pretrained("meta-llama/Llama-3.2-3B")
draft_model = get_peft_model(draft_model, lora_config)

# 训练目标:最小化草稿分布与目标分布之间的 KL 散度
trainer = Trainer(
    model=draft_model,
    train_dataset=dataset,
    args=TrainingArguments(
        per_device_train_batch_size=4,
        learning_rate=2e-4,
        num_train_epochs=1,
        logging_steps=50,
    ),
    data_collator=default_data_collator,
)
trainer.train()

微调后实测接受率从 38% 提升到 58%,加速比从 1.5x 提升到 2.1x——效果相当显著。


四、代码实战:从零实现 Speculative Decoding

4.1 核心算法实现

下面是一个生产可用级别的 Speculative Decoding 实现(基于 PyTorch 和 HuggingFace Transformers),我把它称为「不会在边界条件下崩溃」的工程级别实现:

import torch
import torch.nn.functional as F
from transformers import PreTrainedModel
from typing import Tuple, Optional, List

@torch.no_grad()
def speculative_decode(
    target_model: PreTrainedModel,
    draft_model: PreTrainedModel,
    input_ids: torch.Tensor,
    max_new_tokens: int = 256,
    gamma: int = 6,
    temperature: float = 1.0,
    top_k: Optional[int] = None,
    top_p: Optional[float] = None,
    eos_token_id: int = 2,
    pad_token_id: int = 0,
) -> torch.Tensor:
    """
    推测解码主循环

    Args:
        target_model: 目标模型(大模型)
        draft_model: 草稿模型(小模型)
        input_ids: 输入 token 序列 [batch, seq_len]
        max_new_tokens: 最大生成 token 数
        gamma: 每次推测的 token 数
        temperature: 采样温度
        top_k: Top-K 过滤
        top_p: Top-P (nucleus) 过滤
        eos_token_id: 结束符 ID
        pad_token_id: 填充符 ID

    Returns:
        generated: 生成的 token 序列
    """
    device = input_ids.device
    batch_size = input_ids.shape[0]
    generated = input_ids.clone()
    eos_reached = torch.zeros(batch_size, dtype=torch.bool, device=device)

    # 初始化 KV Cache
    target_past = None
    draft_past = None

    tokens_generated = 0
    total_draft_tokens = 0
    total_accepted_tokens = 0

    while tokens_generated < max_new_tokens:
        if eos_reached.all():
            break

        # 剩余允许生成的 token 数
        remaining = max_new_tokens - tokens_generated
        current_gamma = min(gamma, remaining)

        # ============ 阶段 1:草稿阶段 ============
        draft_sequence = generated[-1:] if draft_past is not None else generated
        draft_tokens: List[torch.Tensor] = []
        draft_logits_list: List[torch.Tensor] = []
        current_input = draft_sequence

        for step in range(current_gamma):
            draft_out = draft_model(
                current_input,
                past_key_values=draft_past,
                use_cache=True,
            )
            draft_logits = draft_out.logits[:, -1, :]
            draft_past = draft_out.past_key_values

            draft_probs = F.softmax(draft_logits / temperature, dim=-1)

            if top_k is not None:
                draft_probs = _top_k_filtering(draft_probs, top_k)
            if top_p is not None:
                draft_probs = _top_p_filtering(draft_probs, top_p)

            next_token = torch.multinomial(draft_probs, num_samples=1)
            draft_tokens.append(next_token)
            draft_logits_list.append(draft_logits)
            current_input = next_token

        # ============ 阶段 2:验证阶段 ============
        if draft_past is not None:
            # 重置草稿 KV Cache 到验证前状态
            # (实际工程中用更精细的做法,这里简化)
            pass

        # 主模型验证所有草稿 token
        verify_input = torch.cat([generated] + draft_tokens, dim=1)
        target_out = target_model(
            verify_input,
            past_key_values=target_past,
            use_cache=True,
        )
        target_logits = target_out.logits
        target_past = target_out.past_key_values

        # 提取验证位置 logits
        context_len = generated.shape[1]
        verify_logits = target_logits[:, context_len:, :]

        # ============ 阶段 3:拒绝采样 ============
        accepted_tokens_list: List[torch.Tensor] = []
        all_accepted = True
        num_accepted = 0

        for step in range(current_gamma):
            target_logits_t = verify_logits[:, step, :]
            draft_logits_t = draft_logits_list[step]
            candidate_token = draft_tokens[step]

            if all_accepted and (not eos_reached.all()):
                target_probs = F.softmax(target_logits_t / temperature, dim=-1)
                draft_probs = F.softmax(draft_logits_t / temperature, dim=-1)

                p_target = target_probs.gather(1, candidate_token).squeeze(-1)
                q_draft = draft_probs.gather(1, candidate_token).squeeze(-1)

                # 接受概率
                accept_prob = torch.minimum(
                    torch.ones_like(p_target),
                    p_target / (q_draft + 1e-8)
                )

                random_vals = torch.rand_like(accept_prob)
                accept_mask = random_vals < accept_prob
                accept_mask = accept_mask & (~eos_reached)

                # 对每个样本决定:接受草稿 token 还是重新采样
                rejected = ~accept_mask
                if rejected.any():
                    all_accepted = False
                    # 被拒绝的位置从修正分布采样
                    adjust_probs = torch.clamp(target_probs - draft_probs, min=0)
                    adjust_probs_sum = adjust_probs.sum(dim=-1, keepdim=True)
                    adjust_probs = torch.where(
                        adjust_probs_sum > 0,
                        adjust_probs / adjust_probs_sum,
                        target_probs
                    )
                    resample_token = torch.multinomial(adjust_probs, num_samples=1)
                    # 用目标 token 替代被拒绝的
                    final_token = torch.where(
                        accept_mask.unsqueeze(-1),
                        candidate_token,
                        resample_token
                    )
                    num_accepted += accept_mask.sum().item()
                else:
                    final_token = candidate_token
                    num_accepted += candidate_token.size(0)
            else:
                # 已有一个被拒绝,后续从目标分布采样
                target_probs = F.softmax(target_logits_t / temperature, dim=-1)
                if top_k is not None:
                    target_probs = _top_k_filtering(target_probs, top_k)
                if top_p is not None:
                    target_probs = _top_p_filtering(target_probs, top_p)
                final_token = torch.multinomial(target_probs, num_samples=1)

            # EOS 检查
            eos_mask = (final_token == eos_token_id).squeeze(-1)
            eos_reached = eos_reached | eos_mask

            accepted_tokens_list.append(final_token)

        # 拼接接受序列
        accepted_sequence = torch.cat(accepted_tokens_list, dim=1)
        generated = torch.cat([generated, accepted_sequence], dim=1)
        tokens_generated += accepted_sequence.shape[1]
        total_accepted_tokens += num_accepted
        total_draft_tokens += current_gamma * batch_size - (
            batch_size - sum(eos_reached.tolist())
        ) * (current_gamma - accepted_sequence.shape[1])

        # Bonus token:如果所有草稿都被接受
        if all_accepted:
            bonus_logits = target_logits[:, -1, :]
            bonus_probs = F.softmax(bonus_logits / temperature, dim=-1)
            if top_k is not None:
                bonus_probs = _top_k_filtering(bonus_probs, top_k)
            if top_p is not None:
                bonus_probs = _top_p_filtering(bonus_probs, top_p)
            bonus_token = torch.multinomial(bonus_probs, num_samples=1)
            generated = torch.cat([generated, bonus_token], dim=1)
            tokens_generated += 1

            eos_mask = (bonus_token == eos_token_id).squeeze(-1)
            eos_reached = eos_reached | eos_mask

    acceptance_rate = total_accepted_tokens / max(total_draft_tokens, 1)
    print(f"[Speculative Decode] Acceptance rate: {acceptance_rate:.2%}, "
          f"Tokens: {tokens_generated}")
    return generated[:, input_ids.shape[1]:]


def _top_k_filtering(probs: torch.Tensor, k: int) -> torch.Tensor:
    """Top-K 过滤"""
    values, _ = torch.topk(probs, k, dim=-1)
    min_values = values[:, -1].unsqueeze(-1)
    probs[probs < min_values] = 0.0
    return probs / probs.sum(dim=-1, keepdim=True)


def _top_p_filtering(probs: torch.Tensor, p: float) -> torch.Tensor:
    """Top-P (nucleus) 过滤"""
    sorted_probs, sorted_indices = torch.sort(probs, descending=True, dim=-1)
    cumulative_probs = sorted_probs.cumsum(dim=-1)
    sorted_indices_to_remove = cumulative_probs > p
    sorted_indices_to_remove[:, 1:] = sorted_indices_to_remove[:, :-1].clone()
    sorted_indices_to_remove[:, 0] = False
    indices_to_remove = sorted_indices_to_remove.scatter(
        1, sorted_indices, sorted_indices_to_remove
    )
    probs[indices_to_remove] = 0.0
    return probs / probs.sum(dim=-1, keepdim=True)

4.2 使用 vLLM 的生产方案

自己实现 Speculative Decoding 虽然能透彻理解原理,但生产环境建议直接用 vLLM——它从 v0.5.3 开始就把 speculative_model 作为一等公民支持,而且经过大量生产环境的考验。

vLLM 配置示例:

from vllm import LLM, SamplingParams

# 配置推测解码
llm = LLM(
    model="meta-llama/Llama-3.1-70B",
    speculative_model="meta-llama/Llama-3.1-8B",
    num_speculative_tokens=6,
    speculative_draft_tensor_parallel_size=1,
    use_v2_block_manager=True,
    enable_chunked_prefill=True,
    max_model_len=8192,
    gpu_memory_utilization=0.90,
)

sampling_params = SamplingParams(
    temperature=0.7,
    top_p=0.9,
    max_tokens=1024,
)

outputs = llm.generate(
    "Explain the concept of speculative decoding with a Python example",
    sampling_params,
)
for output in outputs:
    print(output.outputs[0].text)

实测数据: 在 4×A100 (80GB) 上,Llama-3.1-70B + Llama-3.1-8B 组合:

配置Token/s首 Token 延迟 (TTFT)显存占用接受率
基线(无投机)9.2320ms138GB-
γ=418.5180ms148GB62%
γ=624.1140ms151GB58%
γ=826.3130ms155GB53%
γ=1227.1128ms162GB45%

γ=6 是 sweet spot——再拉高 γ,边际收益递减,显存开销反而明显增加。这是因为草稿序列越长,后面的 token 脱离上下文的约束越远,接受率自然下降。

4.3 与 HuggingFace Transformers 集成

如果不想迁移到 vLLM,HuggingFace Transformers 从 4.42 版本也开始支持 assisted_decoding

from transformers import AutoModelForCausalLM, AutoTokenizer

assistant_model = AutoModelForCausalLM.from_pretrained(
    "meta-llama/Llama-3.2-3B",
    torch_dtype=torch.float16,
    device_map="auto",
)

model = AutoModelForCausalLM.from_pretrained(
    "meta-llama/Llama-3.1-8B",
    torch_dtype=torch.float16,
    device_map="auto",
)
tokenizer = AutoTokenizer.from_pretrained("meta-llama/Llama-3.1-8B")

inputs = tokenizer("Write a Python function to sort a list", return_tensors="pt")

# assisted decoding: 一行开启
outputs = model.generate(
    **inputs,
    assistant_model=assistant_model,
    max_new_tokens=256,
    do_sample=True,
    temperature=0.7,
)

print(tokenizer.decode(outputs[0]))

一个参数的改动就能拿到 2x 左右的加速,这是目前入门最快的路径。


五、进阶变体:不止是草稿-验证

5.1 Medusa:多预测头并行

Google 的 Medusa(发表于 2024 年初)放弃了独立的草稿模型,转而在目标模型最后一层添加多个额外的预测头。每个预测头负责预测往后第 k 个位置的 token:

输入:    "The quick brown"
Head 0:  "fox"    (正常的 LM head)
Head 1:  "jumps"  (预测 +1 位置)
Head 2:  "over"   (预测 +2 位置)
Head 3:  "the"    (预测 +3 位置)

Medusa 的精妙在于:

  1. 单模型架构,不需要独立加载草稿模型的权重
  2. 预测头之间通过 attention mask 建立依赖(后面的树结构)
  3. 生成 token 树而不是单链,用树注意力并行验证

不过 Medusa 的代价在于需要微调——Medusa head 的训练需要约 1 天在 8×A100 上完成(对 70B 模型)。如果你已经在用 LoRA 微调目标模型,可以低成本合并 Medusa head 的训练。

5.2 EAGLE:特征级投机

EAGLE(发表于 2024 年中)的思路更激进:它不训练预测头,而是把草稿阶段嵌入到目标模型的 feature space 中。

EAGLE 的核心是训练一个特征级草稿网络,它接收目标模型某一层的 hidden states 作为输入,直接预测下一层的 token 分布。这样一来,草稿过程几乎不增加额外计算,而且特征空间天然对齐——因为用的是目标模型自己的中间表示。

EAGLE 在一系列基准测试中取得了 2.5-3.5x 的加速比,是目前已知加速效果最好的方法之一。

5.3 DeepSeek DSpark:置信度调度 + 半自回归

2026 年 6 月,DeepSeek 联合北大发布 DSpark(论文标题:DSpark: Confidence-Scheduled Speculative Decoding with Semi-Autoregressive Generation),这是近期最值得关注的 Speculative Decoding 变体。梁文锋本人署名论文作者,分量十足。

DSpark 的两个核心创新:

1. 置信度调度(Confidence Scheduling)

传统方法固定 γ 为常数。DSpark 的做法是:

  • 草稿模型每生成一个 token,同时输出该 token 的置信度(Max Probability)
  • 如果置信度持续高(比如 > 0.9),自动延长 γ
  • 如果置信度断崖下跌(比如 < 0.7),立即停止草稿阶段,进入验证

这就避免了「低质量草稿末尾的 token 白白浪费验证资源」的问题。

def dspark_gamma_schedule(
    draft_logits: torch.Tensor,
    base_gamma: int = 6,
    min_gamma: int = 2,
    max_gamma: int = 12,
    high_conf_threshold: float = 0.90,
    low_conf_threshold: float = 0.70,
) -> int:
    """
    DSpark 风格的动态 γ 调度
    
    根据草稿模型 token 级别的置信度分布决定
    继续推测还是提前截断。
    """
    probs = F.softmax(draft_logits, dim=-1)
    max_probs = probs.max(dim=-1).values  # [seq_len]
    
    # 检查最近的 k 个 token
    recent = max_probs[-3:] if len(max_probs) >= 3 else max_probs
    avg_recent_confidence = recent.mean().item()
    min_recent_confidence = recent.min().item()
    
    if min_recent_confidence > high_conf_threshold:
        # 全部高置信度,继续
        return min(max_gamma, base_gamma + 4)
    elif avg_recent_confidence > 0.80:
        # 平均置信度可接受
        return base_gamma
    elif min_recent_confidence < low_conf_threshold:
        # 出现低置信度,立即截断
        return max(min_gamma, base_gamma - 3)
    else:
        return base_gamma

2. 半自回归生成(Semi-Autoregressive Generation, SAG)

传统草稿模型逐个 token 自回归生成——每步都要前向一次小模型。DSpark 用一个非自回归的编码器-解码器结构,一次性预测多个位置。

SAG 把 γ 个 token 的草稿生成耗时从 O(γ) 降到 O(1),代价是 token 间的依赖关系不那么严格,接受率略有下降。但综合下来,DSpark 在 DeepSeek-V4-Pro 上实现了 60-85% 的单用户推理加速——如果用国产卡(如昇腾 910B)部署,这个加速比更加显著,因为国产卡的带宽瓶颈更严重。

5.4 Lookup Decoding:检索增强

对于重复度较高的场景(代码生成、JSON 输出、SQL 查询、模板填写),Lookup Decoding 用了一个更极简的思路:不从草稿模型生成候选,而从 KV Cache 的历史中直接检索已有的 token 序列

class LookupDraftEngine:
    """从 KV Cache 历史中检索候选序列"""
    
    def __init__(self, ngram_cache: Dict[str, torch.Tensor], max_ngram: int = 5):
        self.ngram_cache = ngram_cache  # "n-gram -> token_ids" 映射
        self.max_ngram = max_ngram
    
    def lookup(self, context: torch.Tensor) -> Optional[List[torch.Tensor]]:
        """根据当前上下文检索最长的匹配序列"""
        context_len = context.shape[-1]
        
        for n in range(self.max_ngram, 0, -1):
            if n >= context_len:
                continue
            key = tuple(context[0, -n:].tolist())
            if key in self.ngram_cache:
                return self.ngram_cache[key]
        return None

在 SQL 生成场景中,Lookup Decoding 的接受率超过 80%——因为大量 SQL 片段(SELECT、FROM、WHERE、JOIN)重复出现。代码补全场景也有类似效果。


六、性能工程:把每一纳秒吃干榨净

6.1 KV Cache 共享

Speculative Decoding 有一个隐藏的显存杀手:草稿模型和目标模型各自维护一份独立的 KV Cache。对于 70B + 8B 组合,两份 KV Cache 在 4K 上下文时约需要 3-4GB,看似不多,但如果调整到 32K 上下文,就会膨胀到 24-32GB。

解决方案是 KV Cache 共享:草稿模型验证产生的 KV Cache 可以直接被目标模型复用,避免重复计算和存储。

6.2 CUDA Graph 优化

草稿模型的自动回归循环包含条件分支(接受/拒绝),这会导致 CUDA graph 在验证阶段被反复重建。vLLM 的解法是预先编译所有可能分支的 CUDA graph,用 bitmask 运行时选择。

简单来说,CUDA graph 把一系列 GPU 操作「录制」成一个静态图,后续执行时可以省掉内核调度开销。但对于有分支的循环,需要为每个分支单独录制。Speculative Decoding 的分支数随 γ 指数增长,所以 vLLM 使用了一种称为「graph flattening」的技术,只录制最频繁的路径(比如全部接受 + 前 k 个接受),用后备逻辑处理罕见分支。

6.3 Batch 场景的取舍

Speculative Decoding 有一个公认的短板:在大的 batch size 下,加速效果会衰减

原因很简单:当 batch 增大到一定程度(比如 batch_size ≥ 16),目标模型本身的 GPU 利用率已经很饱和了——内存带宽不再是瓶颈,计算单元接近满载。这时候塞一个草稿模型只会增加竞争。

Batch Size基线 (token/s)SD (token/s)加速比
19.224.12.6x
428.552.31.8x
848.167.51.4x
1672.380.21.1x
3295.698.81.03x

所以,Speculative Decoding 的最佳应用场景是低并发、追求单用户延迟的在线服务——比如 AI 编程助手(每个用户独享一个推理实例)和交互式对话。高并发批处理场景(如离线生成数据)更适合用常规的 batch 优化。

6.4 流式输出的兼容性

Speculative Decoding 与流式输出(Streaming)有天然的配合。因为每次验证周期可以一次性输出 3-6 个 token,流式输出的「字词间隔」反而比传统方案更平滑——传统方案是逐字输出,Speculative Decoding 可以「一波一波」地批量吐字。

不过需要注意:流式输出的 Yielding 频率与 γ 值直接相关。γ 太大时,用户会感觉到「等了一段时间,突然吐出几个字」的脉冲感。建议流式场景下 γ 控制在 4-6 之间。

6.5 国产硬件的特殊考量

2026 年的国产 AI 芯片(昇腾 910B、寒武纪 MLU590、海光 DCU)有个共同特点:内存带宽远不如 NVIDIA,但计算能力并不差太多

以昇腾 910B 为例,其 HBM 带宽约 1.2TB/s(A100 的 60%),但 FP16 算力可达 400 TFLOPS。这就意味着国产卡的「内存壁」更严重——算力有余而带宽不足。

Speculative Decoding 在这种场景下的优势被放大:因为带宽瓶颈越突出,减少参数加载次数的收益就越显著。在昇腾 910B 上实测,70B 模型的 Speculative Decoding 加速比可达 3.0-3.5x,高于 A100 的 2.5x。


七、质量保障:怎么确认没有「偷工减料」

7.1 数学保证 vs 工程现实

Speculative Decoding 理论上无损,但工程实现中有几个可能引入偏差的地方:

  1. 浮点精度差异:拒绝采样的概率计算用 FP16 时可能产生微小偏差
  2. KV Cache 不同步:草稿模型和目标模型的 KV Cache 在保存机制上可能不一致
  3. 温度/采样策略差异:如果草稿和目标模型使用了不同的采样参数,分布就会偏移

7.2 统计性验证

下面这个验证工具可以帮你在上线前确认质量是否对齐:

from scipy import stats
from collections import Counter
import numpy as np

def validate_speculative_decoding(
    model_only_outputs: List[str],
    sd_outputs: List[str],
    ngram_n: int = 4,
) -> dict:
    """
    统计验证 Speculative Decoding 没有改变输出分布
    """
    def get_ngram_dist(texts, n):
        counter = Counter()
        for text in texts:
            for i in range(len(text) - n + 1):
                counter[text[i:i+n]] += 1
        total = sum(counter.values())
        return {k: v/total for k, v in counter.most_common(2000)}

    baseline_dist = get_ngram_dist(model_only_outputs, ngram_n)
    sd_dist = get_ngram_dist(sd_outputs, ngram_n)
    
    common = set(baseline_dist.keys()) & set(sd_dist.keys())
    results = {
        "sample_size": len(model_only_outputs),
        "ngram_n": ngram_n,
        "common_ngrams": len(common),
    }
    
    if len(common) > 100:  # 足够的样本做卡方检验
        baseline_counts = np.array([baseline_dist[k] * len(model_only_outputs) for k in common])
        sd_counts = np.array([sd_dist[k] * len(sd_outputs) for k in common])
        chi2, p_value = stats.chisquare(sd_counts, f_exp=baseline_counts)
        results["chi2_statistic"] = float(chi2)
        results["p_value"] = float(p_value)
        results["distribution_match"] = p_value > 0.05
    
    # 额外的任务特定验证
    results["avg_len_diff_ratio"] = abs(
        np.mean([len(t) for t in model_only_outputs]) -
        np.mean([len(t) for t in sd_outputs])
    ) / np.mean([len(t) for t in model_only_outputs])
    
    return results

八、生产部署 Checklist

如果你准备在生产环境上 Speculative Decoding,这份 checklist 可以帮你少踩坑。

8.1 硬件选型矩阵

目标模型推荐硬件草稿模型策略预期加速
7-8B单卡 A100-40G同模型 INT4/Self-SD1.5-2.0x
13-20B单卡 A100-80GPhi-3-mini / Qwen2-1.5B2.0-2.5x
70B2×A100-80GLlama-3-8B / 同模型 INT41.8-2.6x
180B+4+×A100-80GLlama-3-8B / TP 切分草稿1.5-2.2x
70B (国产)8×昇腾 910BLlama-3-8B INT42.5-3.5x

8.2 参数调优指南

  1. γ 初始值设为 6,然后看接受率:
    • 接受率 > 65%:加大 γ 到 8-10
    • 接受率 < 40%:减小 γ 到 4,或者换草稿模型
  2. 草稿模型不做 TP 切分:小模型切分后的通信延迟会吃掉所有收益
  3. 开启 chunked prefill:vLLM 的 enable_chunked_prefill=True
  4. 监控接受率:如果持续 < 40%,草稿模型和目标模型差异过大
  5. 流式场景 γ 不要 > 6:否则用户感知到脉冲式输出

8.3 常见踩坑记录

坑 1:共享 tokenizer
不同 tokenizer 的 alignment 在草稿验证中可以编写 50 行以上胶水代码。直接用同族模型可以完美避开。

坑 2:EOS 伪触发
草稿模型生成的 EOS token 在验证阶段很可能被拒绝。如果草稿模型提前输出了 EOS,但 KV Cache 已经写入了 EOS 伪触发标记,需要特殊处理缓存状态。

坑 3:前缀缓存冲突
如果同时用了前缀缓存(Prefix Caching),草稿模型和目标模型的 cache 块可能冲突。vLLM 的 use_v2_block_manager 尝试解决了这个问题,但在自定义实现中要小心。

坑 4:长上下文性能衰减
当上下文长度超过 8K 时,草稿模型的接受率会显著下降——因为小模型的 long-range 依赖捕捉能力弱。建议长上下文场景下把 γ 值动态降低。


九、总结与展望

Speculative Decoding 不依赖任何魔法——它只是把「GPU 算力过剩而带宽不足」这个物理约束,通过重构计算流程巧妙地绕了过去。

回顾关键要点:

  1. 数学保证无损——拒绝采样的并行化,输出分布与纯主模型严格一致
  2. 草稿模型是关键——同族降级是最佳实践,接受率决定加速效果
  3. 适合低并发、高延迟敏感场景——AI 编程、交互式对话、流式推理场景收益最大
  4. 工程实现有深度——KV Cache 共享、CUDA graph 优化、动态 γ 调度,每个细节都可能成为瓶颈
  5. DSpark 等新变体在突破边界——置信度调度和半自回归生成让加速比更进一步,尤其在国产硬件上

展望 2026 年下半年到 2027 年,几个值得关注的演进方向:

  • 多模型共享 KV Cache 池:同一个推理集群内引入多个专业化草稿模型(代码草稿、数学草稿、对话草稿),根据输入动态路由到最合适的草稿模型
  • 硬件原生加速:NVIDIA 下一代 GPU 架构传闻会在 Tensor Core 中集成推测解码的 token 树验证路径
  • 推理调度器原生集成:Kubernetes + 推理网关层感知 Speculative Decoding 的资源配置,自动为开启了 SD 的 deployment 分配额外显存

如果你正在做 LLM 推理服务的性能优化,Speculative Decoding 是目前性价比最高的技术之一——不改变模型、不牺牲质量、不增加架构复杂度,只靠工程技巧就能拿到 2-3 倍的加速。对于大部分在线推理场景来说,这可能是 2026 年最值得投入的一项工程优化。


参考:Leviathan et al., "Fast Inference from Transformers via Speculative Decoding", ICML 2023;Chen et al., "Accelerating Large Language Model Decoding with Speculative Decoding", 2023;vLLM 官方文档 v0.5.3-v0.6+;Stern et al., "Blockwise Parallel Decoding for Deep Autoregressive Models" (Medusa 前身);DeepSeek × PKU, "DSpark: Confidence-Scheduled Speculative Decoding with Semi-Autoregressive Generation", 2026。

推荐文章

Vue3中如何处理异步操作?
2024-11-19 04:06:07 +0800 CST
Vue3 中提供了哪些新的指令
2024-11-19 01:48:20 +0800 CST
PHP 的生成器,用过的都说好!
2024-11-18 04:43:02 +0800 CST
程序员茄子在线接单