ICML'26 | DFlash:用 Diffusion 模型干掉 AR Drafter,投机解码提速 6×

ICML’26 | DFlash:用 Diffusion 模型干掉 AR Drafter,投机解码提速 6×

原文:DFlash: Block Diffusion for Flash Speculative Decoding


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:

EAGLE-3 vs DFlash drafting 延迟对比

这个延迟优势意味着 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):

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:

DFlash 训练时的 attention mask 设计

具体做法:

  1. 从 response 序列里随机采样 anchor token 位置
  2. 以 anchor 为起点,把接下来 block_size - 1 个 token 随机 mask 掉一部分(绿色)
  3. 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 反而更高:

主结果表:DFlash vs EAGLE-3 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 推理效率和参数高效微调,和本文的工程背景有直接关联。

动手学AutoML书籍封面

Flag Counter