投机解码深度实战:从草稿-验证采样到 EAGLE 树状草稿,把 LLM 解码延迟砍一半

如果你做过 LLM 线上服务调优,一定撞过这堵墙:首 token 出来之后,后面每个 token 都要等,而 GPU 大部分时间在空转。

这不是玄学。自回归解码每生成一个 token 就要把整个模型权重从 HBM 搬进 SM 一次,算算术强度(arithmetic intensity)低得可怜。以 70B 模型 FP8 推理为例,权重约 70GB,H100 的 HBM 带宽 3.35TB/s,单步解码的理论下限就是 70GB / 3.35TB/s ≈ 21ms,也就是单用户上限约 47 token/s——不管你有多少 FLOPS 都没用,这是内存墙。

增大 batch 能把权重搬运摊薄,吞吐上去了,但每个请求的延迟也上去了,而且高 batch 下 GPU 已经饱和,再优化收益递减。要在低延迟(小 batch)场景下提速,只剩一条路:让一次前向解码出多个 token。

这就是投机解码(Speculative Decoding)。

一、第一性原理:草稿 + 并行验证

思路非常朴素:

  1. 用一个又快又小的草稿模型(draft model)自回归生成 γ 个候选 token。它慢不到哪去,因为小模型一次前向也就几毫秒。
  2. 把 prefix + γ 个候选一并喂给目标模型(target model),一次前向拿到 γ+1 个位置的真实概率分布。注意这一步是并行的,代价约等于单次解码。
  3. 逐位置做接受-拒绝采样:草稿和目标分布越接近,接受率越高。
  4. 若全部接受,还能从目标模型第 γ+1 个位置免费多采样一个 token(bonus token)。

关键在于:这套算法在数学上保证输出分布与目标模型逐 token 采样严格同分布。也就是说它不是近似加速,是无损加速——这一点是它区别于量化、剪枝、稀疏化的根本优势,也是它能被 vLLM、SGLang、TensorRT-LLM 全线支持的原因。

接受率的数学结论很干净。设 α 为草稿与目标分布的平均不匹配率,则单轮期望接受 token 数为:

E[accepted] = (1 - α^(γ+1)) / (1 - α)

当 γ=5、α=0.4 时约 1.6 个;α 降到 0.2 时约 3.2 个。所以整个游戏的核心不是把 γ 调大,而是把 α 做小——即让草稿模型尽可能"像"目标模型。

二、核心算法:30 行讲透接受-拒绝采样

下面是一个可直接跑的最小实现(来自 Leviathan et al. 的经典算法,我做了工程化改写):

import torch
from torch import Tensor

@torch.no_grad()
def speculative_sample(
    prefix: Tensor,              # (1, T) 已生成 token
    draft_model, target_model,
    gamma: int = 5,
    max_new_tokens: int = 256,
    temperature: float = 0.0,
):
    x = prefix
    n = 0
    stats = {"draft": 0, "accepted": 0, "rounds": 0}

    while n < max_new_tokens:
        # ---------- 1. 草稿阶段:小模型自回归 gamma 步 ----------
        draft_ids, draft_probs = [], []
        cur = x
        for _ in range(gamma):
            logits = draft_model(cur).logits[:, -1, :]
            p = logits.softmax(-1) if temperature == 0 else (logits / temperature).softmax(-1)
            nxt = p.argmax(-1, keepdim=True) if temperature == 0 else torch.multinomial(p, 1)
            draft_ids.append(nxt)
            draft_probs.append(p)
            cur = torch.cat([cur, nxt], dim=-1)

        draft_ids = torch.cat(draft_ids, dim=-1)          # (1, gamma)
        draft_probs = torch.stack(draft_probs, dim=1)     # (1, gamma, V)
        stats["draft"] += gamma

        # ---------- 2. 验证阶段:目标模型单次前向 ----------
        tgt = target_model(torch.cat([x, draft_ids], dim=-1)).logits
        tgt = tgt[:, -(gamma + 1):, :]                    # (1, gamma+1, V)
        q = tgt.softmax(-1) if temperature == 0 else (tgt / temperature).softmax(-1)

        # ---------- 3. 接受-拒绝采样 ----------
        accepted = 0
        rejected = False
        for i in range(gamma):
            xi = draft_ids[:, i]
            p_i = draft_probs[:, i, :]
            q_i = q[:, i, :]
            ratio = q_i.gather(-1, xi[:, None]).squeeze(-1) / p_i.gather(-1, xi[:, None]).squeeze(-1).clamp_min(1e-12)
            if torch.rand((), device=x.device) < torch.clamp(ratio, max=1.0):
                accepted += 1
            else:
                # 关键:拒绝后不是丢弃,而是从残差分布 (q - p)+ 重采样
                resid = (q_i - p_i).clamp_min(0.0)
                resid = resid / resid.sum(-1, keepdim=True)
                draft_ids[:, i] = torch.multinomial(resid, 1).squeeze(-1)
                accepted += 1
                rejected = True
                break

        keep = draft_ids[:, :accepted]
        x = torch.cat([x, keep], dim=-1)
        n += accepted
        stats["accepted"] += accepted
        stats["rounds"] += 1

        if not rejected:
            # 全部接受 → 从目标分布的最后一个位置补采一个 bonus token
            bonus = q[:, gamma, :]
            nxt = bonus.argmax(-1, keepdim=True) if temperature == 0 else torch.multinomial(bonus, 1)
            x = torch.cat([x, nxt], dim=-1)
            n += 1
            stats["accepted"] += 1

    stats["accept_rate"] = stats["accepted"] / max(stats["draft"], 1)
    return x, stats

三个容易被写错的地方,值得单独说:

  • 残差重采样不能省。很多手写实现拒绝后直接 fallback 成"用目标模型重新采样",这在数学上就偏了,破坏了无损性质。必须是 (q - p)⁺ / Σ(q - p)⁺。
  • 验证时要算 γ+1 个位置。第 γ+1 个位置是 bonus token 的来源,少了它等于白扔一步。
  • temperature=0(贪心)时行为退化。贪心下接受条件变成"argmax 一致",接受率通常更高,投机解码在低温度场景收益最明显——这解释了为什么它在代码补全、结构化输出、Agent 工具调用(普遍 temperature≈0)上特别有效。

三、草稿从哪来:三种路线的取舍

1)独立小模型(Draft Model) 最经典的方案,如用 Qwen2.5-0.5B 给 32B 打草稿。优点是实现简单、框架支持最好;缺点是必须同词表,且要额外部署一份权重,显存和运维成本实打实增加。小模型和你业务分布不 match 时,接受率会崩到 20% 以下,加速比直接变负。

2)自草稿(Self-Drafting):Medusa / MTP / EAGLE 这是 2025-2026 年的主流。不再用独立模型,而是在目标模型上挂轻量头:

  • Medusa:在最后一层隐藏状态上并联多个 MLP head,分别预测第 +2、+3… 个 token。零额外权重搬运,但接受率中等。
  • MTP(Multi-Token Prediction):DeepSeek 系列原生自带,训练时就让模型学会预测后续 token,推理时天然高接受率,是目前性价比最高的路线之一。
  • EAGLE / EAGLE3:把目标模型倒数第二层的特征和 token 嵌入拼起来,用一个极小的自回归头做草稿。因为它复用了目标模型的语义特征,接受率显著高于 Medusa,实测 0.6~0.8 是常见区间。EAGLE3 进一步去掉了特征层面的不确定性累积。

3)无模型草稿:n-gram / 检索式 从 prompt 自身或历史语料里用 n-gram 匹配做草稿(如 REST、PLL)。零成本、零显存、零训练,在摘要、抽取、代码补全、RAG 重述这类"输出高度重复输入"的场景里,接受率能意外地高。我的建议是:这类场景优先试 n-gram,成本几乎为零。

选型结论:有原生 MTP 的权重就用 MTP,没有就上 EAGLE3,业务输出高度重复输入就先试 n-gram,只有在以上都不成立时才考虑独立草稿模型。

四、从链式到树状:把接受长度再翻一倍

链式草稿有个结构性问题:γ 个候选串成一条链,中间任何一个位置被拒,后面全废。γ=5 时如果第 2 个被拒,只拿到 2 个 token,白算 3 个。

树状草稿(Tree-based Drafting)的解法是:草稿阶段用 top-k 分支生成一棵候选树,验证阶段用 tree attention mask 让一次前向并行验证整棵树的所有路径。

# 构造树注意力掩码:节点只能看到自己的祖先
def build_tree_mask(parents: list[int], n: int) -> Tensor:
    # parents[i] = -1 表示根节点(接在 prefix 之后)
    mask = torch.zeros(n, n, dtype=torch.bool)
    anc = [set() for _ in range(n)]
    for i, p in enumerate(parents):
        if p >= 0:
            mask[i] = mask[p].clone()
            mask[i, p] = True
            anc[i] = anc[p] | {p}
    mask = mask | torch.eye(n, dtype=torch.bool)
    return mask, anc  # mask[i,j]=True 表示 i 可 attend j

# 典型 EAGLE 配置:深度 5、每层保留 top-8、总候选 token 数 64
# 实测期望接受长度 3~4,比同 gamma 的链式高 40%+

代价是验证阶段的计算量随候选总数上升,而收益随最长路径上升。工程上的甜点区是总候选 token 数 40~80、深度 4~6,再大验证开销就吃掉了收益。这个数字在 vLLM 和 SGLang 里都作为默认超参存在,不是巧合。

五、生产落地清单

vLLM 示例(EAGLE3):

VLLM_USE_V1=1 vllm serve meta-llama/Llama-3.1-70B-Instruct \
  --tensor-parallel-size 4 \
  --speculative-config '{
      "method": "eagle3",
      "model":  "yuhuili/EAGLE3-LLaMA3.1-Instruct-70B",
      "num_speculative_tokens": 5
  }'

SGLang 示例:

python -m sglang.launch_server --model meta-llama/Llama-3.1-70B-Instruct \
  --speculative-algo EAGLE3 \
  --speculative-draft-model-path yuhuili/EAGLE3-LLaMA3.1-Instruct-70B \
  --speculative-num-steps 5 --speculative-eagle-topk 8 \
  --speculative-num-draft-tokens 64

DeepSeek 系列直接用原生 MTP 即可,无需外部草稿权重。具体方法名以你所用框架版本为准。

什么时候不该开投机解码:

  • 高 QPS 场景。batch 已经 128+,GPU 计算饱和,加投机只会增加调度开销,吞吐反而下降。投机解码是延迟优化手段,不是吞吐优化手段——这句话值得贴墙上。
  • 超长输出。输出越长接受率越容易累积性衰减,测出来加速比可能只有 1.1x。
  • 显存已经吃紧。草稿权重 + 树验证的 KV cache 都要额外显存,OOM 得不偿失。
  • 追求可复现。投机引入了额外随机源,即使 seed 固定,跨版本的 token 序列也可能漂移。

怎么测才不骗自己: 一定要分开看 TPOT(每输出 token 时间)和 E2E 延迟,并记录真实接受率。框架日志里的 acceptance_rate 低于 0.4 时,请先去解决草稿质量问题(换 EAGLE、调 top-k、检查词表是否一致),而不是继续调 γ。γ 从 5 加到 8 在低接受率下的收益远不如把 α 从 0.5 降到 0.3。

六、我的判断

投机解码在 2026 年已经从"论文里的花活"变成了推理引擎的默认配置项。真正的工程门槛不在算法——算法就那 30 行——而在于三件事:草稿方案与业务分布的匹配度、树形超参的调优、以及接受率的持续可观测。

把接受率做成大盘指标、按业务场景分组监控,比纠结用哪个框架重要得多。当你的 Agent 工具调用、代码补全、结构化输出这类低温度场景占到七成流量时,投机解码是投入产出比最高的那一项优化——不用改模型、不用重新训练、不损失任何输出质量,改几行配置就有 1.5~2.5x 的 TPOT 下降。这种便宜,线上服务里不多见。

点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部