DSpark:DeepSeek 线上跑的投机解码,半自回归 + 置信度调度怎么做到的

DSpark:DeepSeek 线上跑的投机解码,半自回归 + 置信度调度怎么做到的

原文:DSpark: Confidence-Scheduled Speculative Decoding with Semi-Autoregressive Generation


1. 前言

你有没有想过,投机解码的加速上限到底在哪?

一个简单的分析:每轮 decoding 的 per-token 延迟是

\[L = \frac{T_\text{draft} + T_\text{verify}}{\tau}\]

$\tau$ 是平均 acceptance length,分子是每轮花的时间,分母是每轮接受的 token 数。想让 $L$ 小,要么压分子(更快的 draft/verify),要么拉 $\tau$。

现有方案各有问题:

  • AR drafter(如 EAGLE-3):$\tau$ 高,但 $T_\text{draft} \propto \gamma$(串行,草拟 token 数越多越慢),被迫用小 $\gamma$
  • Parallel drafter(如 DFlash):$T_\text{draft}$ 极低(一次 forward),但 $\tau$ 会随 block 内位置衰减——位置越靠后,没有前序 token 的依赖,预测越离谱,acceptance rate 快速下跌

这个”位置越靠后越不准”的现象有个名字:suffix decay

DSpark 是 DeepSeek 在自己的 V4 生产系统上实际部署的方案。它的核心思路是:parallel drafter 的低延迟要保留,但要给 block 内部引入序列依赖来修复 suffix decay;同时,不是所有草稿 token 都值得送给大模型验证——引入置信度调度,按 token 质量和系统负载动态裁减验证长度。


2. Suffix Decay:并行 Drafter 的问题在哪

先把问题讲清楚。

假设你要草拟一个 block [x₁, x₂, x₃, x₄],并行 drafter 一次 forward 同时预测这 4 个位置:

并行预测:
x₁ = argmax P(x₁ | context)         ← 只依赖 context,还好
x₂ = argmax P(x₂ | context)         ← 不知道 x₁ 是什么!
x₃ = argmax P(x₃ | context)         ← 不知道 x₁, x₂ 是什么!
x₄ = argmax P(x₄ | context)         ← 完全盲猜 x₁-x₃

举个极端的例子,假设语言模型在这个位置有两种等可能的续写:

  • "of course"
  • "no problem"

并行 drafter 可能给出 "of problem""no course"——每个位置独立取最高概率,但组合起来逻辑不通,目标模型一验证就拒掉。

如下图,对 DFlash(蓝色,纯并行)和 DSpark(绿色)在不同 draft 位置的条件 acceptance rate 做对比:

位置条件 acceptance rate 对比

蓝色曲线在 Chat 任务上非常典型——第 1 个位置还不错,第 4、5 个位置就明显下跌。DSpark 的绿色曲线则保持平稳甚至略有上升。原因正是加入了序列依赖。


3. Semi-Autoregressive 架构

DSpark 的 drafter 分两阶段,如下图:

DSpark 整体架构和 decoding 循环

3.1 Parallel Stage:骨干并行生成

第一阶段直接用 DFlash 作为 parallel backbone,一次 forward 并行生成整个 block 的 hidden states 和 base logits:

输入:anchor token x₀ + γ 个 <mask>
一次 forward → h₁, ..., hᵧ(hidden states)
              U₁, ..., Uᵧ(base logits,shape: [γ, V])

这一步 $T_\text{draft}$ 极低,不随 $\gamma$ 增长。

3.2 Sequential Stage:轻量序列头修正

第二阶段是一个很轻量的模块,从左到右依次给每个位置的 logits 加上一个”转移偏置” $B_k$,让第 $k$ 个 token 的预测依赖前面已经采样出来的 token:

\[p_k(v \mid x_0, x_{<k}) = \frac{\exp(U_k(v) + B_k(x_0, x_{<k}, v))}{\sum_{u \in \mathcal{V}} \exp(U_k(u) + B_k(x_0, x_{<k}, u))}\]

$U_k$ 是 parallel stage 给出的 base logit(已经有对全局 context 的感知),$B_k$ 是对 intra-block 依赖的修正。

两种实现:

Markov Head(默认):最简单,只依赖前一个 token,低秩分解 $B = W_1 W_2$:

class MarkovHead(nn.Module):
    def __init__(self, vocab_size, rank=256):
        super().__init__()
        self.W1 = nn.Embedding(vocab_size, rank)   # 上一个 token -> rank 维向量
        self.W2 = nn.Linear(rank, vocab_size, bias=False)  # rank -> V

    def forward(self, prev_token_id: torch.Tensor) -> torch.Tensor:
        # prev_token_id: [batch]
        embed = self.W1(prev_token_id)   # [batch, rank]
        return self.W2(embed)            # [batch, V]  ← 转移偏置 B(x_{k-1}, ·)

RNN Head:维护一个跨位置的 recurrent state $s_k$,能看到整个 prefix 历史:

class RNNHead(nn.Module):
    def __init__(self, vocab_size, backbone_dim, rank=256):
        super().__init__()
        self.W1 = nn.Embedding(vocab_size, rank)
        # 单个 linear 拆成 gate / candidate / output 三部分
        self.Wgco = nn.Linear(2 * rank + backbone_dim, 3 * rank, bias=False)
        self.W2 = nn.Linear(rank, vocab_size, bias=False)

    def step(self, s_prev, prev_token_id, h_k):
        # s_prev: [batch, rank],h_k: [batch, backbone_dim]
        z = torch.cat([s_prev, self.W1(prev_token_id), h_k], dim=-1)  # [batch, 2r+d]
        gates = self.Wgco(z).chunk(3, dim=-1)   # gate, cand, out 各 [batch, rank]
        g = torch.sigmoid(gates[0])
        s_k = g * s_prev + (1 - g) * torch.tanh(gates[1])             # GRU-like update
        B_k = self.W2(torch.tanh(gates[2]))                            # [batch, V]
        return s_k, B_k

整个 sequential stage 的推理循环:

def sequential_stage(base_logits, hidden_states, seq_head, anchor_id):
    # base_logits:   [batch, γ, V]  ← 来自 parallel stage
    # hidden_states: [batch, γ, d]  ← 来自 parallel stage
    B, gamma, V = base_logits.shape
    draft_tokens = []
    prev_id = anchor_id           # 第一个"前序 token"是 anchor
    s = torch.zeros(B, seq_head.rank)  # RNN state 初始为 0

    for k in range(gamma):
        s, bias = seq_head.step(s, prev_id, hidden_states[:, k, :])
        logits_k = base_logits[:, k, :] + bias     # 并行 logit + 转移修正
        x_k = torch.multinomial(logits_k.softmax(-1), 1).squeeze(-1)
        draft_tokens.append(x_k)
        prev_id = x_k              # 下一步知道当前采样结果

    return torch.stack(draft_tokens, dim=1)  # [batch, γ]

这个 sequential loop 虽然是串行的,但 seq_head 极其轻量(只有 embedding lookup + 低秩矩阵乘法),每步的计算量远小于 parallel backbone,所以 $T_\text{sequential} \ll T_\text{parallel}$,总延迟仍然由并行阶段主导。

如下图,drafter 层数对 acceptance length 的影响——2 层 DSpark 就已经超过了 5 层 DFlash:

drafter 深度对 acceptance length 的影响


4. Confidence-Scheduled Verification

第二个核心设计解决”验证浪费”的问题。

问题的本质:投机解码送给目标模型验证的 token 越多,一旦某个位置被拒,后面的都白算了。如果 draft 质量很差,每次送 8 个 token 去验证,平均只接受 2 个,剩下 6 个 token 的验证纯属浪费 GPU 显存和计算。

在高并发系统里更严重:每个被白验证的 token 都占用了 target model 的 batch capacity,这些 capacity 本可以服务其他用户请求。

4.1 Confidence Head:预测每个 token 的接受概率

DSpark 在 sequential stage 的基础上,为每个草稿 token 输出一个置信度 $c_k \in (0, 1)$,代表”在前面所有 token 都被接受的条件下,第 $k$ 个 token 也被接受的概率”:

\[c_k = \sigma\left(w^\top [h_k;\ W_1[x_{k-1}]]\right)\]

监督信号来自 draft 和 target 分布之间的 TV 距离:

\[c_k^* = 1 - \frac{1}{2} \|p_d_k - p_t_k\|_1\]

这个公式来自投机解码的理论分析——per-step acceptance rate 等于 1 减去 TV 距离。用这个作为监督,模型就学会了预测自己草稿的质量。

class ConfidenceHead(nn.Module):
    def __init__(self, backbone_dim, markov_rank):
        super().__init__()
        self.w = nn.Linear(backbone_dim + markov_rank, 1, bias=False)

    def forward(self, h_k, prev_embed):
        # h_k: [batch, d], prev_embed: [batch, rank]
        return torch.sigmoid(self.w(torch.cat([h_k, prev_embed], dim=-1))).squeeze(-1)

训练时用 BCE loss:

\[\mathcal{L}_\text{conf} = -\sum_{k=1}^\gamma w_k \left[ c_k^* \log c_k + (1-c_k^*) \log(1-c_k) \right]\]

4.2 Sequential Temperature Scaling(校准)

神经网络输出的概率通常偏高(overconfident),直接用 $c_k$ 来估算 acceptance rate 误差大。

DSpark 引入了 Sequential Temperature Scaling(STS):prefix survival probability 是各位置 confidence 的连乘 $\prod_{i \leq k} c_i$,STS 从左到右逐位置校准,每次用一维 grid search 找让该位置 ECE 最小的温度参数,且已校准的前序位置不变。

这是个 order-preserving 变换(不改变 token 质量的相对排序),只修正绝对概率量级,让 scheduler 的估算更准。

4.3 Hardware-Aware Prefix Scheduler

有了校准后的 $c_k$,position $j$ 的 prefix survival probability 是:

\[a_{r,j} = \prod_{i \leq j} c_{r,i}\]

即”前 $j$ 个 token 全部被接受的概率”,天然单调递减。

Scheduler 的目标:在整个 batch 的 $R$ 个请求里,选择每个请求的验证长度 $\ell_r$,使得系统 期望总吞吐最大化

\[\Theta = \tau \cdot \text{SPS}(B)\]

其中 $B = \sum_r (1 + \ell_r)$ 是送给目标模型的总 token 数,$\text{SPS}(B)$ 是在 batch size $B$ 下的引擎 steps/s(提前 profile 好的)。

由于 $a_{r,j}$ 单调递减,这个优化问题可以用贪心算法高效求解:把所有候选 token 按 $a_{r,j}$ 从高到低排序,依次加入 batch,每加一个 token 计算一次当前期望吞吐,一旦吞吐不再增加就停止。

整个 Algorithm 1 的逻辑是:

def hardware_aware_scheduler(confidence_scores, sps_table, R, gamma):
    # confidence_scores: [R, gamma]
    # sps_table: dict {batch_size -> steps_per_sec}

    # 计算每个请求每个位置的 survival probability
    a = torch.cumprod(confidence_scores, dim=1)  # [R, gamma]

    # 所有候选 (request, position) 按 a 从大到小排序
    candidates = [(a[r, j].item(), r, j)
                  for r in range(R) for j in range(gamma) if a[r, j] > 0]
    candidates.sort(reverse=True)

    lengths = [0] * R          # 初始每个请求验证长度为 0
    B = R                      # 初始 batch size(每个请求至少发 anchor token)
    tau_expected = R           # 初始期望接受数(每个请求至少接受 anchor)
    best_theta = tau_expected * sps_table[B]
    best_lengths = lengths.copy()

    for a_val, r, j in candidates:
        if j != lengths[r]:    # 必须按顺序扩展(prefix 约束)
            continue
        lengths[r] += 1
        B += 1
        tau_expected += a_val
        theta = tau_expected * sps_table.get(B, 0)
        if theta > best_theta:
            best_theta = theta
            best_lengths = lengths.copy()
        else:
            break              # 单调性保证可以早停

    return best_lengths

这个设计的妙处在于:系统负载高时 $\text{SPS}(B)$ 下降快,scheduler 自动收紧每个请求的验证长度;负载低时自动放开,让每个请求验证更多 token。MTP-1(原来的 baseline)验证长度固定为 2,完全没有这种弹性。

如下图,confidence threshold 的 sweep 实验——随着阈值提高,acceptance rate 从 ~77% 上升到 92.5%,因为低质量草稿被剪掉了:

Confidence threshold 和 acceptance rate 的 tradeoff


5. 训练目标

DSpark 的训练损失由三项组成,都带位置衰减权重 $w_k = e^{-(k-1)/\gamma}$(越靠前权重越大):

\[\mathcal{L} = \alpha_\text{ce} \mathcal{L}_\text{ce} + \alpha_\text{tv} \mathcal{L}_\text{tv} + \alpha_\text{conf} \mathcal{L}_\text{conf}\]
  • $\mathcal{L}_\text{ce}$(交叉熵,$\alpha=0.1$):让 drafter 预测正确 token
  • $\mathcal{L}_\text{tv}$(TV 距离,$\alpha=0.9$):最小化 draft 和 target 分布之间的 total variation distance,直接最大化期望 acceptance rate
  • $\mathcal{L}_\text{conf}$(置信度,$\alpha=1.0$):训练 confidence head 准确预测自己草稿会不会被接受

$\mathcal{L}\text{tv}$ 权重远大于 $\mathcal{L}\text{ce}$ 是有意设计的——因为 acceptance rate 和 TV 距离直接挂钩,最小化 TV 比最大化 token 准确率对最终 speedup 更有帮助。


6. 线上部署结果

这篇论文最有说服力的部分是 DeepSeek-V4 生产系统的真实流量数据。

如下图,throughput(系统吞吐)vs. TPS(每个用户的生成速度)的 Pareto frontier:

Throughput vs TPS:DeepSeek-V4 线上表现

和 MTP-1(生产 baseline)相比,DSpark 把这条 Pareto 曲线整体右移——相同用户 TPS 下吞吐提升 +51~52%,相同系统吞吐下 TPS 提升 +60~57%,在极低负载的高 TPS 区间提升高达 +85%

再看 load-adaptive 的调度效果:

负载自适应吞吐和 verification budget

并发请求数增多时,DSpark 的 verification budget(下图折线)自动从 ~5 收缩到 ~3.5,始终保持最优系统吞吐。MTP-1 固定在 2,无论负载高低都不动。


7. Offline Benchmark 主结果

如下表,在 Qwen3-4B/8B/14B 和 Gemma4-12B 上,DSpark 的 acceptance length 全面超过 Eagle3 和 DFlash:

主结果:acceptance length 对比

以 Qwen3-8B + GSM8K 为例:Eagle3 是 5.30,DFlash 是 5.33,DSpark 是 6.17。Chat 类任务(MT-Bench、Alpaca)提升更显著,恰恰说明序列依赖建模对”语境连贯性”要求高的任务更重要。


8. 个人 Take

DSpark 有两个地方让我觉得设计得很干净。

一个是 semi-autoregressive 架构的分工:parallel backbone 承担全部的”大计算量”(对全局 context 的感知),sequential head 只做”小修正”(block 内依赖),整体延迟由前者主导,后者几乎是免费的。这种”重骨干 + 轻修正”的组合是个值得借鉴的范式。

另一个是 confidence-scheduled verification 的视角——它把投机解码从”单请求 drafting 问题”扩展成了”多请求 batch scheduling 问题”。这个视角转变很重要:在高并发系统里,draft token 的边际价值不是固定的,取决于当前系统负载。把这两者联合优化,才能真正释放 speculative decoding 在生产场景下的潜力。

这类工作的说服力有相当一部分来自真实线上数据。DeepSeek 能把自家 V4 的生产流量拿来验证,这个条件不是每个研究组都有的。但方法本身的思路——用轻量序列头修 suffix decay,用置信度驱动动态调度——是普适的,不依赖特定的系统环境。


如果这篇文章涉及的 LLM 推理效率优化你想系统深入,可以看看我之前出版的《动手学 AutoML:从 NAS 到大语言模型优化实战》,书里有专章讲 LLM 推理效率和参数高效微调,和本文的工程背景有直接关联。

动手学AutoML书籍封面

Flag Counter