Skip to content

深入推测解码:从概念到实现

说明:本文翻译自 Moncef Abboud 的博客 Exploring Speculative Decoding: From Concept to Implementation,版权归原作者所有。

在本文中,我们将通过一个具体的、以 vLLM 为焦点的实现来探索推测解码,涵盖草稿模型、EAGLE、MTP 以及其中涉及的权衡。

图:EAGLE-3 在 vLLM 上的推测解码(在 RTX 3090 上实测)

引言

在本文中,我将讨论推测解码(speculative decoding)——一种用于优化 LLM 推理的技术。它属于那种第一次了解时,会让你有一种"哦对!有道理"恍然领悟感的东西。

但首先,我想说明为什么 LLM 推理优化很重要。

LLM 为用户查询生成响应。我们训练大模型一次,往往耗费 8 位或 9 位数的成本,但服务(serving)却要发生数百万次。运行一个模型需要相当强大的硬件,而说 GPU 昂贵且供应紧张几乎是一种保守的说法。一个微小的效率提升,长期来看意味着巨大的节省。

LLM 推理基础回顾

现代 GPU 是令人印象深刻的猛兽。但它们也有自己的怪癖。一块 GPU 每秒能执行数百亿亿次运算,却只能从 GPU 内存向计算单元移动几万亿字节。

LLM 推理是自回归的。如果我们有一些输入 token [t1tn],模型会为序列中的每个 token 给出一个大小为 vocab_size 的 logits 向量。我们使用最后一个 token tn 的 logits 来采样下一个 token tn+1

这意味着,除非我们每字节运行数百次运算——而 LLM 中并不是这样——我们本质上都是内存带宽受限的。

当我们批处理时,每组权重(X@W)会执行更多运算:批次越大,我们越能复用从内存取出的权重 W,GPU 的利用率就越高。

从 KV Cache 到推测解码

每个 token 都要经过多个层,每层又有几个标准模块:归一化、MLP,以及最值得注意的 Transformer 注意力模块。在注意力之前的每个模块中,一个 token 不关心其他 token。而在注意力中,token ti 需要知道 token t0ti1。具体来说,它需要访问这些 token 在该特定层对应的键和值。

LLM 推理引擎带来的一个关键优化是:不必为每个新 token 重新计算那些 K 和 V 张量,而是存储它们,只为新 token 运行计算。这就是著名的 KV cache。当新 token 经过模型时,它们的键和值向量会被加入缓存。

采样出的 token tn+1 被加入输入,用它的 logits 采样 tn+2,依此类推。在第一次模型运行时,即使我们只需要最后一个 logits 用于生成,也会在输出中计算出 n 个 logits 向量。而在同一次前向传播中计算 1 个还是几个候选 token,其差异往往远小于连续运行多个解码步骤。

每运行一次(即产生一个新 token 的一次前向传播),我们都需要加载全部模型权重,下一次又要重新加载,如此往复。由于内存带宽是限制因素,我们大部分时间都在等待权重加载。

这就是推测解码的用武之地。如果我们能提前猜出几个可能的 token 并喂给模型,就可以在一次前向传播中验证它们。如果猜对了,我们几乎免费获得这些 token。

常规方式:

[t1tn]tn+1

推测解码:

[t1tn]tn+1,tn+2,tn+3,

当然,这只有在被猜出的 token(我们称之为草稿 token,draft tokens)通常正确、且猜测过程比目标模型的完整前向传播便宜得多时才有效。否则,我们还不如直接使用那个大型原始模型——我们称之为目标(target)

有许多技术可以生成这些草稿 token:n-gram、EAGLE、MTP 等。

但核心思想是一样的。一次前向传播通常用于采样一个 token。如果我们能运行一个更便宜的草稿过程并预测几个额外的 token,就可以减少昂贵的目标模型步骤数。如果我们做对了,还可以精确地保持目标模型的分布,就像根本没有草稿模型一样。

伪代码

python
def propose(tokens, num_speculative_tokens):
    draft_tokens = []

    for _ in range(num_speculative_tokens):
        logits = generate_draft_token(tokens)  # 相比目标模型,这必须非常快
        token = sample(logits)
        draft_tokens.append(token)
        tokens.append(token)

    return tokens, draft_tokens

一旦我们有了草稿 token,就验证它们:

python
def sample_verify(tokens, num_drafts):
    logits = target_model(tokens)  # 运行目标模型前向传播,这是真实分布

    for i, logit in enumerate(logits[-num_drafts:]):
        target_token = sample(logit)
        draft_token = tokens[-num_drafts + i]

        if target_token == draft_token:
            accept(draft_token)
        else:
            reject(draft_token)
            break

保证原始分布

重要的一点是:推测解码不会为了速度而牺牲正确性。最终输出遵循与目标模型相同的分布。

如果我们遵循以下算法,这一点在数学上是可以证明的:

text
# 这一步应该比常规解码快得多
1. 通过草稿模型的概率分布 q(x) 采样 token
# 在一次前向传播中廉价地验证多个 token
2. 目标模型计算真实分布 p(x)
3. 以概率 min(1, p(x)/q(x)) 接受采样出的 token。分两种情况:
   a. 如果 p(x) >= q(x),以概率 1 接受。目标模型对 token 的喜爱程度至少与草稿一样,所以我们保留它。
   b. 如果 p(x) < q(x),以概率 p(x)/q(x) 接受,该概率小于 1。
4. 如果被拒绝,从 max(0, p(x)-q(x)) 中采样一个修正 token

第 4 步是关键:它正是我们处理草稿与目标之间差异的地方。通过从 pq 中重新采样,我们覆盖了草稿模型遗漏的 token。

第 3b 步也很重要。如果草稿认为 token 2 的概率是 0.4,但目标认为是 0.2,那么 p/q=0.2/0.4=0.5,所以我们一半的时间接受它。这很合理:草稿过于自信,我们通过比率 p/q 加以纠正。

如果 p>q,目标模型比草稿更喜欢这个 token,所以我们直接接受它。

Speculator(草拟器)

有很多方法可以获取这些草稿 token。我们优化的标准是:尽可能接近目标模型,并且更快

更小的草稿模型

最直观的可能是使用一个更小的模型。想象你有一个 470B 模型,依赖同家族的 7B 模型。

对于复杂 token,小模型表现不佳;但对于重复和简单的内容,它应该不错。例如:

Q:我们如何解决狭义相对论? A:要解决狭义相对论……

小模型应该能轻松猜到我们会重复问题的一部分,并提供合适的草稿 token。答案的其余部分更难,但如果我们把容易的部分做对了,仍然能占得先机。

我们可以想出不同的技术来生成草稿 token。它只是一个函数,接受现有序列并尝试预测 K 个草稿 token。在 vLLM 中,这被实现为可插拔的 speculator

下面的代码片段是从 vLLM 中精简出来的,只保留了骨架和关键数据流。

共享的 proposer 基类如下所示:

python
class SpecDecodeBaseProposer:
    def __init__(self, vllm_config, device, pass_hidden_states_to_model, runner=None):

    @torch.inference_mode()
    def propose(
        self,
        target_token_ids,
        target_positions,
        target_hidden_states,
        next_token_ids,
        token_indices_to_sample,
        common_attn_metadata,
        sampling_metadata,
        mm_embed_inputs=None,
        num_rejected_tokens_gpu=None,
        slot_mappings=None,
    ):
        # 接受当前 token 以及可能用到的额外数据(取决于 proposer 类型)
        # 然后生成草稿 token 及其草稿 logits/probs
        ...

n-gram

事物往往会重复。风水轮流转。这是计算中支撑缓存的一般原则:如果我们看到某些数据,很可能会再次用到它。这就是局部性(locality)。遵循这一原则,我们可以查看序列中最后 N 个 token,搜索之前出现过的位置。我们随后把那次较早出现之后的 K 个 token 作为草稿 token。

从上面的例子看,序列的后缀是"解决狭义相对论",它之前出现在几个词之前,而它之后的内容是"……"。我们猜测那些是可能的草稿 token,结果猜对了。

n-gram 是一种非常简单廉价的运行技术,甚至可以在 CPU 上运行。它的简单也意味着在实践中经常出错,但对于有重复模式的文本非常有用,而代码就是完美的例子。

python
class NgramProposer:
    def __init__(self, vllm_config):
        # 草稿长度和匹配窗口来自 speculative config。
        self.min_n = vllm_config.speculative_config.prompt_lookup_min
        self.max_n = vllm_config.speculative_config.prompt_lookup_max
        self.k = vllm_config.speculative_config.num_speculative_tokens
        self.max_model_len = vllm_config.model_config.max_model_len

    def propose(self, sampled_token_ids, num_tokens_no_spec, token_ids_cpu, slot_mappings=None):
        # 只为实际采样了 token 的请求做推测。
        valid_requests = [i for i, sampled_ids in enumerate(sampled_token_ids)
                          if sampled_ids and num_tokens_no_spec[i] < self.max_model_len]
        return self.batch_propose(len(sampled_token_ids), valid_requests, num_tokens_no_spec, token_ids_cpu)

    def batch_propose(...):
        for i in prange(len(valid_ngram_requests)):
            idx = valid_ngram_requests[i]
            num_tokens = num_tokens_no_spec[idx]
            context_token_ids = token_ids_cpu[idx, :num_tokens]
            drafter_output = _find_longest_matched_ngram_and_propose_tokens(
                origin_tokens=context_token_ids,
                min_ngram=min_n,
                max_ngram=max_n,
                max_model_len=max_model_len,
                k=k,
            )

            valid_ngram_num_drafts[idx] = drafter_output.shape[0]
            if len(drafter_output):
                valid_ngram_draft[idx, : drafter_output.shape[0]] = drafter_output


def _find_longest_matched_ngram_and_propose_tokens(origin_tokens, min_ngram, max_ngram, max_model_len, k):
    # 使用 Knuth–Morris–Pratt (KMP) 算法匹配最长模式
    # 这个视频讲得很清楚:https://www.youtube.com/watch?v=JoF0Z7nVSrA
    # 如果你不熟悉它,值得一看

EAGLE

虽然小草稿模型理论上看起来不错,但在实践中仍然有所欠缺,因为这本质上是两个学习内容不同的模型。

一个有趣的技术是 EAGLE,它有多个迭代版本:EAGLE1、EAGLE2、EAGLE3,以及最新的 EAGLE3.1。其核心思想是:目标模型已经在做大部分繁重的工作,并拥有预测下一个 token 所需的全部信息。其中一些信息就存在于中间隐藏状态中。

因此,EAGLE 不只是依赖 token 嵌入,而是把嵌入加上目标模型的隐藏状态作为轻量草稿网络的输入。该草稿网络随后预测下一个 token。

在 vLLM 的 EAGLE-3 设置中,目标模型在选定层产生隐藏状态。这些状态被拼接起来,通过一个全连接层投影,再经过轻量解码器层和 LM 头,产生草稿 logits。草稿模型仍然是自回归的,但它依赖目标模型的隐藏状态。

本质上,我们添加一个新的解码层,它以所有 token(包括最后一个采样的 token)及其隐藏层为输入,然后利用这些信息预测下一个草稿 token。

对于第一个草稿 token,我们还没有新 token 的目标模型隐藏状态,所以草稿模型不得不依赖它已有的隐藏状态。这就是隐藏状态设计很重要的原因之一。

早期 EAGLE 版本与 EAGLE3 的主要区别在于:早期版本聚焦于最后一个隐藏层,而 EAGLE3 使用多个层,通常跨越早期、中期和后期阶段,以捕获模型推理的更广阔视角。

图:EAGLE-3 论文中的架构示意图(低/中/高层隐藏状态融合)

上面的 EAGLE3 论文示意图说明了这个思想。我们有 "How can",刚刚用目标模型预测了 "I"。对每个预测的 token i,导致它的隐藏状态来自 i1。所以 "How" 的隐藏状态导致 "can","can" 的隐藏状态导致 "I"。在草稿模型中,我们把 i1 的隐藏状态与 i 的 token 嵌入结合起来。

低层、中层和高层的隐藏状态都会被使用。对每个 token,这些隐藏状态向量被拼接,并用一个学习得到的全连接模块组合,该模块被训练用于跨不同阶段挑选和合并相关信息。

python
# vllm/v1/worker/gpu/spec_decode/eagle/speculator.py:405-567

@torch.inference_mode()
def propose(
    self,
    input_batch: InputBatch,
    attn_metadata: dict[str, Any],
    slot_mappings: dict[str, torch.Tensor],
    last_hidden_states: torch.Tensor,     # [num_tokens, H] — 目标的最后一层
    aux_hidden_states: list[torch.Tensor] | None,  # EAGLE-3: 3 × [num_tokens, H]
    num_sampled: torch.Tensor,            # [num_reqs] — 上一迭代接受的个数
    num_rejected: torch.Tensor,           # [num_reqs] — 上一迭代拒绝的个数
    last_sampled: torch.Tensor,           # [max_num_reqs] — 每请求最后接受的 token
    next_prefill_tokens: torch.Tensor,    # [max_num_reqs] — 用于分块 prefill
    temperature: torch.Tensor,
    seeds: torch.Tensor,
    ...
) -> torch.Tensor:
    num_tokens = input_batch.num_tokens_after_padding
    num_reqs = input_batch.num_reqs
    max_query_len = input_batch.num_scheduled_tokens.max()

    # STEP 1: FC 融合(仅 EAGLE-3)
    if aux_hidden_states:
        assert self.method == "eagle3"
        hidden_states = self.model.combine_hidden_states(
            torch.cat(aux_hidden_states, dim=-1)
        )
    else:
        # EAGLE-1/2:直接使用最终隐藏状态(无融合)
        hidden_states = last_hidden_states

    # STEP 2: 准备 EAGLE 输入(Triton kernel)
    prepare_eagle_inputs(
        self.input_buffers, input_batch, self.last_token_indices,
        num_sampled, num_rejected, last_sampled, next_prefill_tokens,
        self.max_num_reqs,
    )

    # STEP 3: PREFILL — 生成草稿 token 0
    self.prefill(
        num_reqs, prefill_batch_desc.num_tokens,
        attn_metadata, slot_mappings,
        num_tokens_across_dp=num_tokens_across_dp,
        cudagraph_runtime_mode=prefill_batch_desc.cg_mode,
        mm_inputs=mm_inputs,
    )

    # STEP 4: 准备 DECODE — 切换到自回归模式
    prepare_eagle_decode(
        self.draft_tokens[:num_reqs, 0], input_batch.seq_lens,
        num_rejected, self.input_buffers, self.max_model_len, self.max_num_reqs,
    )

    # STEP 5: DECODE 循环 — 生成草稿 token 1..K-1
    self.generate_draft(
        num_reqs, decode_batch_desc.num_tokens,
        attn_metadata_updated, slot_mappings_updated,
        num_tokens_across_dp=num_tokens_across_dp,
        cudagraph_runtime_mode=decode_batch_desc.cg_mode,
    )

    return self.draft_tokens[:num_reqs]  # [num_reqs, K]

# 生成草稿 token 仍然是自回归的,因此是 for 循环
def generate_draft(self, num_reqs, num_tokens_padded, attn_metadata, slot_mappings, ...):
    pos = self.input_buffers.positions[:num_reqs]
    query_start_loc = self.input_buffers.query_start_loc[:num_reqs + 1]
    idx_mapping = self.idx_mapping[:num_reqs]

    # ── 遍历草稿位置 1, 2, ..., K-1 ──
    for step in range(1, self.num_speculative_steps):
        # EAGLE 前向:每请求 1 个 token(decode 模式)
        # 输入 = 上一步的 hidden_states + embed(prev_draft_token)
        last_hidden_states, hidden_states = self.run_model(
            num_tokens_padded, attn_metadata, slot_mappings, ...
        )
        last_hidden_states = last_hidden_states[:num_reqs]
        hidden_states = hidden_states[:num_reqs]

        # 我们有了 EAGLE 模型的最终输出
        # 计算 logits 然后采样草稿 token
        logits = self.model.compute_logits(last_hidden_states)
        draft_tokens = self._sample_draft(logits, idx_mapping, pos, step=step)
        self.draft_tokens[:num_reqs, step] = draft_tokens

        # ── 为下一步更新状态(除非这是最后一步)──
        if step < self.num_speculative_steps - 1:
            # ...
            update_eagle_inputs(
                draft_tokens, hidden_states,
                self.input_buffers, self.hidden_states, self.max_model_len,
            )
            # ...

有一个值得指出的微妙之处,因为它相当有趣。当生成第一个之后的草稿 token(token 2 到 k)时,我们使用的是草稿模型自身的隐藏状态。在上图中(右侧第 2、3 步),这对应于使用来自草稿模型的 aiai1,而不是来自目标模型、包含真实隐藏状态的 gigi1

当一个草稿 token 被验证并接受时,来自目标模型的相应真实隐藏状态会被传回草稿模型。此时,草稿的"prefill"步骤(propose 方法中的第 3 步)会用这些纠正后的信息重新计算并重新填充草稿 KV cache,因为"prefill"使用了与目标相同的 slot/attention 元数据。

EAGLE 模型是单独训练、独立存在的模型。我们可以为各种开源模型找到 EAGLE3 模型。vLLM 还有一个用于训练草稿模型(如 EAGLE)的项目,叫做 speculators,它与 vLLM 无缝集成。

MTP

EAGLE 是一个增加额外解码路径的草稿模型,让我们能高效预测额外的草稿 token。

我们能否把类似的额外层合并进目标模型本身,让它成为模型的一部分?这就是 MTP(Multi-Token Prediction,多 token 预测)。一些模型,如 DeepSeek 系列模型,在网络的末尾包含一个额外的多 token 预测层。当一个新 token 被采样时,它的嵌入加上最后一层的隐藏状态会被送入 MTP 层,以预测紧随其后的 token。相同的 LM 头和嵌入会被复用。

这与 EAGLE 非常相似,区别在于模型从训练之初就带着 MTP 训练,它甚至是损失函数的一部分。我们可以有多个 MTP 层来生成更多草稿 token,也可以复用一个 MTP 层来预测额外的草稿 token,尽管准确性可能会下降。

text
给定上下文:"The quick brown fox jumps over the"

目标模型前向传播:
  h = model("The quick brown fox jumps over the")
  token_1 = sample(lm_head(h)) = "lazy"
  Store: h ("the" 位置的隐藏状态)

# 我们用 "the" 位置的隐藏状态预测了 "lazy"
# 两者都用于 MTP 层 0
MTP 层 0:
  输入:embed("lazy") ⊕ h
  输出:h_mtp0
  token_2 = argmax(lm_head(h_mtp0)) = "dog"

# 逻辑同上
MTP 层 1(或复用的层 0):
  输入:embed("dog") ⊕ h_mtp0
  输出:h_mtp1
  token_3 = argmax(lm_head(h_mtp1)) = "and"

草稿:["lazy", "dog", "and"]

如果我们去 Hugging Face 查看 DeepSeek-v4 的权重,可以观察到单个 MTP.0 层独自位于其他 61 个常规层之后。

让我们看看代码。这里没有什么不寻常的:取输入嵌入和导致它们的隐藏状态,归一化使它们处于相近的量级,然后拼接并投影,使它们可以送入一个常规解码层。

python
class DeepSeekMultiTokenPredictorLayer(nn.Module):
    def __init__(self, vllm_config, prefix):
        # MTP 复用模型自身的嵌入 + 隐藏状态路径。
        config = vllm_config.speculative_config.draft_model_config.hf_config
        self.enorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
        self.hnorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
        self.eh_proj = nn.Linear(config.hidden_size * 2, config.hidden_size, bias=False)
        self.shared_head = SharedHead(config=config, prefix=prefix, quant_config=...)
        self.mtp_block = DeepseekV2DecoderLayer(vllm_config, prefix, config=config, topk_indices_buffer=...)

    def forward(self, input_ids, positions, previous_hidden_states, inputs_embeds=None, spec_step_index=0):
        assert inputs_embeds is not None
        # 位置 0 被掩码掉,因为 MTP 只需要移位后的上下文。
        inputs_embeds = torch.where(positions.unsqueeze(-1) == 0, 0, inputs_embeds)
        inputs_embeds = self.enorm(inputs_embeds)
        previous_hidden_states = self.hnorm(previous_hidden_states)
        # 将当前嵌入与前一个隐藏状态融合。
        hidden_states = self.eh_proj(torch.cat([inputs_embeds, previous_hidden_states], dim=-1))
        # 一个额外的解码块把融合后的状态变成草稿 logits。
        hidden_states, residual = self.mtp_block(positions=positions, hidden_states=hidden_states, residual=None)
        return residual + hidden_states


class DeepSeekMultiTokenPredictor(nn.Module):
    def __init__(self, vllm_config, prefix=""):
        config = vllm_config.model_config.hf_config
        self.mtp_start_layer_idx = config.num_hidden_layers
        self.num_mtp_layers = config.num_nextn_predict_layers
        self.layers = nn.ModuleDict({...})
        self.embed_tokens = VocabParallelEmbedding(...)
        self.logits_processor = LogitsProcessor(config.vocab_size)

    def forward(self, input_ids, positions, previous_hidden_states, inputs_embeds=None, spec_step_idx=0):
        current_step_idx = spec_step_idx % self.num_mtp_layers  # 如果草稿 token 数大于 MTP 层数则循环使用层
        return self.layers[str(self.mtp_start_layer_idx + current_step_idx)](
            input_ids, positions, previous_hidden_states, inputs_embeds, current_step_idx
        )

    def compute_logits(self, hidden_states, spec_step_idx=0):
        mtp_layer = self.layers[str(self.mtp_start_layer_idx + (spec_step_idx % self.num_mtp_layers))]
        # 注意这里的 "shared_head"
        return self.logits_processor(mtp_layer.shared_head.head, mtp_layer.shared_head(hidden_states))

生成的草稿 token 随后通过运行所有模型层来验证。

延伸阅读

  • 使用 llm-d 进行分布式 LLM 推理。[链接待更新]
  • 探索混合专家:从概念到推理引擎。[链接待更新]
  • 深入高效 LLM 推理:nano-vLLM。[链接待更新]

Maintained by Robin