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 做对比:

蓝色曲线在 Chat 任务上非常典型——第 1 个位置还不错,第 4、5 个位置就明显下跌。DSpark 的绿色曲线则保持平稳甚至略有上升。原因正是加入了序列依赖。
3. Semi-Autoregressive 架构
DSpark 的 drafter 分两阶段,如下图:

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:

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%,因为低质量草稿被剪掉了:

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:

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

并发请求数增多时,DSpark 的 verification budget(下图折线)自动从 ~5 收缩到 ~3.5,始终保持最优系统吞吐。MTP-1 固定在 2,无论负载高低都不动。
7. Offline Benchmark 主结果
如下表,在 Qwen3-4B/8B/14B 和 Gemma4-12B 上,DSpark 的 acceptance length 全面超过 Eagle3 和 DFlash:

以 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 推理效率和参数高效微调,和本文的工程背景有直接关联。
