深入推测解码:从概念到实现
说明:本文翻译自 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 vocab_size 的 logits 向量。我们使用最后一个 token
这意味着,除非我们每字节运行数百次运算——而 LLM 中并不是这样——我们本质上都是内存带宽受限的。
当我们批处理时,每组权重(
从 KV Cache 到推测解码
每个 token 都要经过多个层,每层又有几个标准模块:归一化、MLP,以及最值得注意的 Transformer 注意力模块。在注意力之前的每个模块中,一个 token 不关心其他 token。而在注意力中,token
LLM 推理引擎带来的一个关键优化是:不必为每个新 token 重新计算那些 K 和 V 张量,而是存储它们,只为新 token 运行计算。这就是著名的 KV cache。当新 token 经过模型时,它们的键和值向量会被加入缓存。
采样出的 token
每运行一次(即产生一个新 token 的一次前向传播),我们都需要加载全部模型权重,下一次又要重新加载,如此往复。由于内存带宽是限制因素,我们大部分时间都在等待权重加载。
这就是推测解码的用武之地。如果我们能提前猜出几个可能的 token 并喂给模型,就可以在一次前向传播中验证它们。如果猜对了,我们几乎免费获得这些 token。
常规方式:
推测解码:
当然,这只有在被猜出的 token(我们称之为草稿 token,draft tokens)通常正确、且猜测过程比目标模型的完整前向传播便宜得多时才有效。否则,我们还不如直接使用那个大型原始模型——我们称之为目标(target)。
有许多技术可以生成这些草稿 token:n-gram、EAGLE、MTP 等。
但核心思想是一样的。一次前向传播通常用于采样一个 token。如果我们能运行一个更便宜的草稿过程并预测几个额外的 token,就可以减少昂贵的目标模型步骤数。如果我们做对了,还可以精确地保持目标模型的分布,就像根本没有草稿模型一样。
伪代码
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,就验证它们:
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保证原始分布
重要的一点是:推测解码不会为了速度而牺牲正确性。最终输出遵循与目标模型相同的分布。
如果我们遵循以下算法,这一点在数学上是可以证明的:
# 这一步应该比常规解码快得多
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 步是关键:它正是我们处理草稿与目标之间差异的地方。通过从
第 3b 步也很重要。如果草稿认为 token 2 的概率是 0.4,但目标认为是 0.2,那么
如果
Speculator(草拟器)
有很多方法可以获取这些草稿 token。我们优化的标准是:尽可能接近目标模型,并且更快。
更小的草稿模型
最直观的可能是使用一个更小的模型。想象你有一个 470B 模型,依赖同家族的 7B 模型。
对于复杂 token,小模型表现不佳;但对于重复和简单的内容,它应该不错。例如:
Q:我们如何解决狭义相对论? A:要解决狭义相对论……
小模型应该能轻松猜到我们会重复问题的一部分,并提供合适的草稿 token。答案的其余部分更难,但如果我们把容易的部分做对了,仍然能占得先机。
我们可以想出不同的技术来生成草稿 token。它只是一个函数,接受现有序列并尝试预测 K 个草稿 token。在 vLLM 中,这被实现为可插拔的 speculator。
下面的代码片段是从 vLLM 中精简出来的,只保留了骨架和关键数据流。
共享的 proposer 基类如下所示:
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 上运行。它的简单也意味着在实践中经常出错,但对于有重复模式的文本非常有用,而代码就是完美的例子。
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
低层、中层和高层的隐藏状态都会被使用。对每个 token,这些隐藏状态向量被拼接,并用一个学习得到的全连接模块组合,该模块被训练用于跨不同阶段挑选和合并相关信息。
# 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 步),这对应于使用来自草稿模型的
当一个草稿 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,尽管准确性可能会下降。
给定上下文:"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 个常规层之后。
让我们看看代码。这里没有什么不寻常的:取输入嵌入和导致它们的隐藏状态,归一化使它们处于相近的量级,然后拼接并投影,使它们可以送入一个常规解码层。
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。[链接待更新]