DSpark 投机采样:串行采样与置信度裁剪

DFlash 一次并行出整块 draft,但块内各位置互不商量,而且整块都要送进 target model。DSpark 只做两件事——draft 侧加一个很轻的串行采样循环让位置之间互相看得见,target 侧加一个裁剪决定这一轮到底验多长。采样规则和接受规则都没变,所以依然无损。

归属:DeepSeek-AI 与北京大学,DSpark: Confidence-Scheduled Speculative Decoding with Semi-Autoregressive Generation,arXiv:2607.05147,作者含 Wenfeng Liang(梁文锋)。代码 DeepSpec,权重 DeepSeek-V4-Pro-DSpark,CC BY 4.0。

前置:只讲相对 DFlash 的增量。块扩散 draft、KV 注入、拒绝采样见 DFlash 投机解码。

一、出发点

DFlash 的链条是:draft model 一次前向算出整块 γ 个位置的分数 → 每位置各自取 top-1 → 整块送去验证。这条链有两个口子,正好在两头:

  • draft 侧:每个位置独立取 top-1。它是在对「不确定的前驱」取期望,所以可能拼出并不存在的组合,越靠块尾越容易错。
  • target 侧:整块都验。尾部位置本来大概率被拒,验它们等于白占 target model 的 batch 容量。

DSpark 的两处改动分别堵这两个口子。

DSpark 的完整一轮

图 1:一轮五步。标绿(串行采样)和标橙(裁剪)是全部新增,其余沿用 DFlash。

二、串行采样:把块内各位置串起来

2.1 它改了什么

draft model 里那层完全不动——仍然是块内双向注意力、一次前向出整块的分数,延迟与 γ 几乎无关。变的只有「怎么从这些分数得到候选」:

  • DFlash:每个位置各自取 top-1,位置之间不通信;
  • DSpark:从左到右逐位采样,采到第 k 位时,用刚刚采出的第 k−1 位去修正第 k 位的分数,再采样。

修正量就是一个加在分数上的偏置:

pk(v)  ∝  exp⁡(Uk(v)+Bk(v))p_k(v)\;\propto\;\exp\big(U_k(v)+B_k(v)\big)

UkU_k 是 draft model 在每位给出的分数,BkB_k 是 draft 侧串行头给的偏置。这是全文唯一的算法公式,读法就一句:draft model 的分数 + 前文修正,再 softmax 采样。

偏置怎么算出来的?查表 + 一次小矩阵乘:BkB_k 只依赖紧邻的前一位 token,所以维护一张「token → 向量」的表,查出来乘个小矩阵,就得到加在词表上的偏置。没有额外的 Transformer 前向——这就是它「轻」的原因。

论文还试了能记住整块前缀的递归版本(RNN 头),结论是收益只多一点、主要出现在很长的块上,但实现复杂不少,所以生产用的是只看前一位那版。

2.2 例子

假设上文已定,接下来该写「很大程度」。draft model 给位置 2 的 top-1 是「可」0.40,「大」只拿到 0.31 —— 因为它不知道位置 1 会采出什么:若位置 1 是「稍」,这里本该是「微」。两个模式被平均掉了,所以纯并行 draft 在这里容易采偏。

串行采样把「不确定的前驱」换成「已经采出的前驱」:

  • 第 1 步没有前文,直接用原始分数,采到 很;
  • 第 2 步用「很」查表得到偏置,抬高「大」、压低「可」,采到 大(0.62 vs 0.19);
  • 第 3 步用「大」的偏置抬高「程」,采到 程(0.88);
  • 第 4 步同理采到 度 —— 但这一位「猜得准不准」是另一回事,见第三节。

一旦「很」被真正采出来,「稍」那条分支就不存在了,概率质量自然集中回「大」。DFlash 在这里可能采到「可」或「慢」,DSpark 不会。

同时这个循环顺手产出每个位置的条件存活概率 ckc_k——在第 k 位之前全部被接受的前提下、第 k 位还能活下来的概率。注意它的输入只有 draft model 隐状态和前一位 token。四个位置依次给出 c1=0.95c_1=0.95、c2=0.92c_2=0.92、c3=0.85c_3=0.85,而 c4c_4 只有 0.300.30。

一个走到底的例子

图 2:上例的完整推演(上半部分是串行采样,下半部分是接着做的裁剪)。

三、裁剪:决定这一轮验多长

3.1 为什么要裁

DFlash 把整块送进 target model。但尾部位置本来大概率被拒,验它们纯属浪费;而且在高并发下,这些浪费的验证位会挤占别的请求的 batch 容量。

论文还有一个观察:结构化任务(数学、代码)的接受率明显高于开放式闲聊,可是验证花费是一样的。所以「写死一个验证长度」天然不合理。

3.2 三个量各自的算法

裁剪要最大化系统总吞吐。所有参与决策的量只有三个,且都能从上一步的输入算出来:

量 算法 说明
累积存活概率 sks_k sk=c1c2⋯cks_k = c_1c_2\cdots c_k,即 sk=sk−1×cks_k = s_{k-1}\times c_k 位置 k 能被验到,前提是前面每位都被接受,所以是连乘
期望被接受 EE E=1+s1+s2+⋯E = 1 + s_1 + s_2 + \cdots 1 是 anchor 自己;被验的位不一定被接受,所以按概率加权
吞吐 Θ\Theta Θ=SPS(B)×E\Theta = \mathrm{SPS}(B)\times E 一次前向平均能产出多少 token,单位 tokens/s

其中 BB 是这次前向的 batch 大小(B=1B = 1 个 anchor + 已准入的 draft 位数),SPS(B)\mathrm{SPS}(B) 是引擎在该 batch 下的 step 速度(在每次前向B个token时,每秒跑多少次前向)——它随 BB 单调不增,引擎初始化时 profile 一次存成表。

于是判据是:逐个纳入,Θ\Theta 不再上升就停。

3.3 接上例子,把每个数字算一遍

沿用上例,四位 draft 的存活概率取 c1,c2,c3,c4=0.95, 0.92, 0.85, 0.30c_1,c_2,c_3,c_4 = 0.95,\,0.92,\,0.85,\,0.30,并取一条示意的 step 速度曲线(真实曲线由引擎实测得到、有数千档,这里取整数方便逐格验算):

SPS(1,2,3,4,5,… )=1000, 520, 400, 340, 280, … steps/s\mathrm{SPS}(1,2,3,4,5,\dots)=1000,\ 520,\ 400,\ 340,\ 280,\ \dots\ \text{steps/s}

它随 BB 单调不增。先把连乘算出来:s1=0.950s_1=0.950,s2=0.950×0.92=0.874s_2=0.950\times0.92=0.874,s3=0.874×0.85=0.743s_3=0.874\times0.85=0.743,s4=0.743×0.30=0.223s_4=0.743\times0.30=0.223。然后逐个纳入:

纳入到 BB 期望被接受 EE SPS(B)\mathrm{SPS}(B) Θ=SPS×E\Theta=\mathrm{SPS}\times E 算法动作
只验 anchor 1 E0=1E_0=1 1000 1000×1=10001000\times1=\mathbf{1000} 起点,best = 1000
位置 1(很) 2 1+0.950=1.9501+0.950=\mathbf{1.950} 520 520×1.950=1014.0520\times1.950=\mathbf{1014.0} 1014>10001014>1000 → 留
位置 2(大) 3 1.950+0.874=2.8241.950+0.874=\mathbf{2.824} 400 400×2.824=1129.6400\times2.824=\mathbf{1129.6} 1130>10141130>1014 → 留
位置 3(程) 4 2.824+0.743=3.5672.824+0.743=\mathbf{3.567} 340 340×3.567=1212.7340\times3.567=\mathbf{1212.7} 1213>11301213>1130 → 留
位置 4(度) 5 3.567+0.223=3.7903.567+0.223=\mathbf{3.790} 280 280×3.790=1061.1280\times3.790=\mathbf{1061.1} 1061<12131061<1213 → break

判决列只做一件事:新的 Θ\Theta 有没有超过目前的最好值。

四、采样流程

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
def dspark_cycle(prefix, target, draft, head, scheduler, gamma=16):
"""一个解码轮次。gamma 为上界,本轮实际验多长由 scheduler 决定。"""

# ① target model 前向(本来就做):产出 anchor,并留下上下文特征
x0, H_ctx = target.forward(prefix)

# ② 并行段:一次前向算出整块 draft 分数(与 DFlash 相同,也是最重的一步)
# 输入 anchor + (gamma-1) 个 mask,得出 gamma 个位置的 draft model 分数
hidden, U = draft.parallel_backbone([x0] + [MASK] * (gamma - 1), H_ctx)

# ③ 串行段:左→右逐位采样,每位只用「刚刚采出的前一位」修正自己
cand, conf = [], []
prev = None
for k in range(gamma):
bias = head.bias(prev) if prev is not None else 0
token = sample(softmax(U[k] + bias)) # p_k ∝ exp(U_k + B_k)
cand.append(token)
conf.append(head.confidence(hidden[k], prev_token=token))
prev = token # 只传一位:循环是串行的,但每步极轻

# ④ 裁剪:按累积存活概率贪心准入,吞吐不再上升就停
keep = scheduler.admit(conf)

# ⑤ target model 只并行验证裁剪后的前缀(接受规则与 DFlash 完全一致)
# 对齐关系:第 i 个 logits 预测的是 cand[i];accepted = 被接受的前缀长度
logits = target.forward_parallel(prefix + [x0] + cand[:keep])
accepted = longest_accepted_prefix(cand[:keep], logits) # 从左碰到第一个拒绝

# ⑥ 提交:接受前缀 + 被拒位置的残差重采(无拒绝时即 bonus token,也是下一轮 anchor)
anchor_next = sample_residual(logits[accepted])
prefix += [x0] + cand[:accepted] + [anchor_next]
target.rollback_kv(keep=len(prefix)) # 丢弃未接受 draft 留下的 KV
return prefix

调度器本体:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
def admit(conf, base_batch=1):
"""
conf : 逐位条件存活概率 c_1..c_gamma
base_batch : 不验任何 draft 时占的 batch 大小(anchor 自己算 1 个位置)
返回 : 本轮该验的 draft 位数

依赖一个全局函数 speed(B):验证 batch 为 B 时的相对 step 速度,
引擎初始化时 profile 一次存成表,对 B 单调不增。
"""
# 位置 k 能被验到,前提是前面每位都被接受 → 累积连乘
survival, s = [], 1.0
for c in conf:
s *= c
survival.append(s)

# 起点:一位 draft 都不验,期望接受 = anchor 自身
batch, expect = base_batch, float(base_batch)
best = speed(batch) * expect

# 贪心:多验第 k 位,多赚 survival[k],代价是 batch +1
# 注意 batch / expect 随循环累加,就是当前已准入的状态
keep = 0
for k in range(len(conf)):
batch += 1
expect += survival[k]
score = speed(batch) * expect
if score <= best: # 吞吐不再上升就停(见 §3.4:这一步是正确性要求)
break
best, keep = score, k + 1

return keep

五、效果

对比 结果
平均接受长度 vs DFlash Qwen3-4B/8B/14B 上 +16.3% / +18.4% / +18.3%
块长 4→16 的额外延迟 仅 0.2%–1.3%(串行循环几乎不要钱)
V4 线上同吞吐下单用户速度 +60%–85%(Flash)、+57%–78%(Pro)

参考资料

DSpark 本体

  • DSpark: Confidence-Scheduled Speculative Decoding with Semi-Autoregressive Generation — arXiv:2607.05147(CC BY 4.0)
    • DFlash 背景:§2.2;半自回归与两种串行头:§3.1;置信头与 STS:§3.2.1;调度器与 Algorithm 1:§3.2.2
    • 生产适配(异步、零开销调度、变长 kernel):§5.2–5.3;线上结果与局限:§5.4;选择偏差反例:Appendix A
  • 代码 deepseek-ai/DeepSpec 权重 deepseek-ai/DeepSeek-V4-Pro-DSpark

前代与对照

站内相关