ICML'26 | DFlash:用 Diffusion 模型干掉 AR Drafter,投机解码提速 6×
ICML’26 | DFlash:用 Diffusion 模型干掉 AR Drafter,投机解码提速 6×
1. 前言
投机解码(Speculative Decoding)的思路很简单:用一个小模型快速草拟几个候选 token,然后扔给大模型并行验证——验证是 free 的,因为大模型本来就支持并行 prefill。
问题出在 drafting 这一步。现在主流方案(比如 EAGLE-3)用的是轻量 AR 头,一次只生成 1 个 token,想草拟 16 个就得跑 16 次串行 forward。序列越长,drafting 延迟越高,慢慢地就把验证那边省下来的时间给吃掉了。
DFlash(ICML 2026)给了一个很直接的解法:把 AR drafter 换成 Diffusion 模型,整个 block 的 token 一次 forward pass 全部并行生成,drafting latency 不再随生成长度增长。再配合从目标模型注入 context feature,草稿质量也上来了。
最终在 Qwen3 系列上,对比 AR baseline 加速 超过 6×,比 EAGLE-3(树大小 60)还快 2.5×。
2. 为什么换成 Diffusion
先把两种 drafter 的本质区别说清楚。
AR drafter 的生成过程:
t=0: 已知 [x₁, x₂, x₃](prompt),预测 x₄
t=1: 已知 [x₁, x₂, x₃, x₄],预测 x₅
t=2: 已知 [x₁, ..., x₅],预测 x₆
...(草拟 γ 个 token 需要 γ 次 forward)
Diffusion drafter 的生成过程:
初始化:把 γ 个待生成位置全部设为 <mask>
输入 = [x₁, x₂, x₃, <mask>, <mask>, <mask>]
一次 forward:同时预测 x₄, x₅, x₆ 的 logits
γ 个 token,1 次 forward
代价是 block 内部 token 之间没有自回归依赖——预测 x₆ 时不知道 x₄、x₅ 是什么,只能靠全局 context 盲猜。但 DFlash 用 target context feature 注入来弥补这个缺陷(后面展开)。
如下图,5 层 DFlash 草拟 16 个 token 只需 ~5ms,而 1 层 EAGLE-3 草拟 16 个 token 需要 ~26ms:

这个延迟优势意味着 DFlash 可以在不超过 latency budget 的前提下,用更深的 drafter(更高质量),这就是它能超过 AR drafter 的根本原因。
3. 核心设计:从目标模型注入 Context Feature
Block diffusion 单靠自己能做到的 acceptance length 很有限——论文里专门验证过,不加任何 conditioning 的 5 层 diffusion drafter,speedup 只有 2-3×。
原因直接:drafter 必须在不知道目标模型内部状态的情况下凭空预测未来 token,等于从零开始猜。
DFlash 的解法:在目标模型 prefill 的时候,把它中间层的 hidden state 抽出来,注入到 drafter 的每一层 KV cache 里。
3.1 提取 Target Context Feature
目标模型跑 prefill 时,从若干中间层(均匀采样 5 层,从第 2 层到倒数第 3 层)抽取 hidden state,拼接后过一个轻量投影层:
\[H_t = \text{RMSNorm}\left(W_c[H^{(l_1)}; \ldots; H^{(l_5)}]\right)\]$W_c \in \mathbb{R}^{D \times 5D}$,把多层信息压缩到 draft 模型的 hidden dimension,得到 context feature $H_t$。
用 PyTorch 写大概是这样:
class ContextFeatureExtractor(nn.Module):
def __init__(self, target_hidden, draft_hidden, num_layers=5):
super().__init__()
self.proj = nn.Linear(target_hidden * num_layers, draft_hidden, bias=False)
self.norm = nn.RMSNorm(draft_hidden)
def forward(self, hidden_states: list[torch.Tensor]) -> torch.Tensor:
# hidden_states: list of [seq_len, target_hidden], len=num_layers
fused = torch.cat(hidden_states, dim=-1) # [seq_len, target_hidden * num_layers]
return self.norm(self.proj(fused)) # [seq_len, draft_hidden]
这个 context feature 只需要计算一次,之后所有 drafting iteration 都复用。
3.2 KV Injection:注入到每一层
拿到 $H_t$ 后,不是把它和 draft token embedding 拼在一起作为输入(EAGLE-3 的做法),而是直接注入到 drafter 每一层的 KV 里:
\(Q_i = W^Q_i H_d\) \(K_i = [W^K_i H_t;\ W^K_i H_d]_{\text{seq}}\) \(V_i = [W^V_i H_t;\ W^V_i H_d]_{\text{seq}}\)
$H_d$ 是 draft token 的 hidden state,$[\ ;\ ]_{\text{seq}}$ 表示在 sequence 维度拼接。
用 PyTorch 写大概是这样:
class DraftAttentionLayer(nn.Module):
def forward(self, h_draft, h_ctx, attn_mask=None):
# h_draft: [batch, γ, d] ← γ 个 mask token 的 hidden states
# h_ctx: [batch, N, d] ← target context feature(整个序列)
Q = self.wq(h_draft) # [batch, γ, d]
K_ctx = self.wk(h_ctx) # [batch, N, d]
K_draft = self.wk(h_draft) # [batch, γ, d]
K = torch.cat([K_ctx, K_draft], dim=1) # [batch, N+γ, d]
V_ctx = self.wv(h_ctx)
V_draft = self.wv(h_draft)
V = torch.cat([V_ctx, V_draft], dim=1) # [batch, N+γ, d]
# bidirectional attention:block 内部 mask token 互相可见
out = F.scaled_dot_product_attention(Q, K, V, attn_mask=attn_mask)
return self.wo(out) # [batch, γ, d]
关键点:draft token 只产生 Query,context feature 只作为额外的 KV entry,不走 Q、output projection、FFN,只相当于给每层的 attention 提供了更丰富的 key-value 池。
这样无论 drafter 有多少层,每一层都能直接 attend 到目标模型的 context feature,信息不会随深度衰减。
如下图,DFlash 的整体推理设计(可以看到 context feature 如何注入进 KV cache):

3.3 一次 Forward 完成 Block 预测
有了 context feature 之后,整个 drafting 过程是:
def draft_block(anchor_token, ctx_feature, drafter, block_size=16):
# anchor_token: 目标模型上一轮验证后产生的最后一个 token
# ctx_feature: 从目标模型 prefill 时提取的 context feature
# 初始化:anchor + γ 个 mask token
mask_token_id = tokenizer.mask_token_id
draft_input = torch.cat([
anchor_token.unsqueeze(1), # [B, 1]
torch.full((B, block_size), mask_token_id) # [B, γ]
], dim=1) # [B, 1+γ]
# embed → drafter forward(bidirectional attention within block)
h = drafter.embed(draft_input) # [B, 1+γ, d]
for layer in drafter.layers:
h = layer(h, ctx_feature) # KV injection
# 只取 mask 位置的 logits
logits = drafter.lm_head(h[:, 1:, :]) # [B, γ, V]
# 并行采样(或 argmax)
draft_tokens = logits.argmax(dim=-1) # [B, γ]
return draft_tokens
整个 block 的 γ 个 token 一次 forward 搞定,无论 γ 多大延迟都不变。
4. 训练:Block Masking
推理时 drafter 拿到的是全部 mask 的 block,训练时也要对齐这个行为。
如下图,DFlash 的训练时 attention mask:

具体做法:
- 从 response 序列里随机采样 anchor token 位置
- 以 anchor 为起点,把接下来
block_size - 1个 token 随机 mask 掉一部分(绿色) - drafter 的任务:给定 anchor + 剩余可见 token + context feature,预测所有 mask 位置
训练时多个 block 拼在一起,用 block 内部双向 attention + 跨 block 单向 causal mask(Flex Attention 实现),可以在一次 forward + backward 里处理整个长序列,效率很高。
损失函数带一个位置衰减权重 $w_k = e^{-(k-1)/\gamma}$,让靠前的位置有更高的训练权重——因为前几个 token 被接受的概率更高,对整体 acceptance length 的贡献更大:
\[\mathcal{L} = -\sum_{k=1}^{\gamma} w_k \log p_d(x_k^*)\]5. 实验结果
5.1 主结果
如下表,在 Qwen3-4B/8B(thinking 模式关闭)上,DFlash 对比 AR baseline 平均加速 4.9×(greedy),在 Math 任务上 acceptance length τ 最高达到 7.87。
EAGLE-3 用了树大小 60(验证 overhead 很大),DFlash 只用 block size 16,speedup 反而更高:

5.2 为什么 Diffusion 能超过 AR?
这里有个看起来反直觉的地方:diffusion 没有 block 内依赖,质量理论上比 AR 差,为什么 acceptance length 更高?
根本原因是 latency budget 的不同分配:
- EAGLE-3 草拟 16 个 token 需要 16 次 forward,延迟 ~26ms,只能用浅层 drafter(否则更慢)
- DFlash 草拟 16 个 token 只需 1 次 forward,延迟 ~5ms,可以用 5 层更深的 drafter
同样的 latency budget 下,DFlash 能用更大、更深的模型,综合质量反而更高。这是典型的”用延迟换质量”的 Pareto 优化。
6. 个人 Take
DFlash 抓住了一个真正的系统瓶颈:之前大家都在想怎么让 drafter 预测得更准(更大的 tree、更好的 AR 头),但核心矛盾是串行 drafting 的延迟随草拟长度线性增长,从根本上限制了能草拟多少 token。
换成 diffusion 之后,drafting latency 和草拟长度解耦,这个约束就没了。可以放心堆更深的 drafter、更大的 block,在相同延迟下拿到更高的 acceptance length。
当然,diffusion drafter 的训练比 AR 复杂——block masking 设计、context conditioning、position-weighted loss,每个细节都有讲究,开箱复现的成本不低。目前结果也主要在 Qwen3 系列上,跨架构的泛化性还需要更多验证。
但”用并行替换串行”这个方向是对的,DFlash 算是把这条路走通的第一个工作。
如果这篇文章涉及的 LLM 推理效率优化你想系统深入,可以看看我之前出版的《动手学 AutoML:从 NAS 到大语言模型优化实战》,书里有专章讲 LLM 推理效率和参数高效微调,和本文的工程背景有直接关联。
