上周出去玩了,拖了一周终于简单看了一遍 DeepSeek 的 DSpark,顺手整理成一篇笔记,之前没有了解过这方面,所以还是看了两三天。
它不是提升模型“智力”的训练方法,而是一套偏推理系统侧的加速设计:目标是在尽量不破坏 target model 输出分布的前提下,让同一个模型在推理时更快地产生 token。
相关仓库:DeepSpec
这篇主要想讲清楚两件事:一是标准 speculative decoding 的 verify 机制到底怎么工作;二是 DSpark 在 draft 和 verify 两个阶段分别改了什么。
speculative decoding 背景
大语言模型的标准生成方式是自回归解码:每次只生成一个新 token。这样做最稳,但也带来一个很现实的问题:
- 模型越大,每一步 forward 越贵
- 回答越长,要走的步数越多
- 在线服务里用户一多,延迟和吞吐压力就会很明显
所以推理系统里一直有一个核心问题:
能不能在不破坏目标模型输出分布的前提下,让一次“确认”尽量多产出几个 token?
它的标准思路可以概括成两步:
- draft 先让一个更便宜、更快的 draft model 或 draft 模块预写后面几个 token。
- verify 再让真正的 target model 一次性验证这几个 token,接受能通过的最长前缀。
如果草稿质量足够高,那么 target model 一次验证就不只“确认 1 个 token”,而可能确认 2 个、3 个甚至更多 token。这样平均到每个 token 的延迟就会下降。
speculative decoding 之所以重要,不只是因为它快,而是因为经典路线追求的是 lossless / exact 加速,也就是:
- 最终采样分布仍然和 target model 一致
- 不是简单拿一个小模型直接替代大模型
- 而是让小模型负责“提案”,大模型保留“裁决权”
所以它和“直接用弱一点的模型换速度”不是一个问题。
看起来简单,但实际会遇到很多问题,例如:
draft 太弱,接受率太低,白忙一场;
draft 太强,自己又变得很贵;
一次提太长,后缀 token 很容易被拒;
高并发时,验证过多低质量 token 反而拖垮吞吐;
不同任务域里接受率差异很大,代码、数学、闲聊的表现不一样。所以后来的很多工作,本质上都在围绕三个变量做优化:
- draft 怎么设计得更好
- 一次 draft 多长更合适
- 哪些 token 值得送去 verify
DSpark 正是在这几个变量里,重点优化了前两个半:
- 它让 draft 既保留并行速度,又补一点顺序依赖
- 它让 verify 不再固定吃完整段,而是按置信度截断前缀
下面先把标准 speculative decoding 的 verify 机制讲清楚,再看 DSpark 到底改了什么。
普通自回归怎么做
假设当前已经有前缀:
$$ y $$后面真正要生成的 token 是:
$$ x_1, x_2, x_3 $$普通自回归的 target model 会这样做:
- 跑一次,得到 $p_t(\cdot \mid y)$,采样出 $x_1$
- 再跑一次,得到 $p_t(\cdot \mid y, x_1)$,采样出 $x_2$
- 再跑一次,得到 $p_t(\cdot \mid y, x_1, x_2)$,采样出 $x_3$
所以:
- 3 个新 token
- 要 3 次 target 解码 step
speculative decoding 的验证思路
现在换成 speculative decoding。
先让 draft model 提议一小段,流程也是和上面一样,只是模型更小,成本更低:
$$ \hat{x}_1, \hat{x}_2, \hat{x}_3 $$这里帽子表示“草稿提议”。
然后 target model 不再一步一步自己生成,而是做一件事:
把前缀 $y$ 和 draft 提议的这一小段一起看掉。
也就是一次性处理:
$$ y, \hat{x}_1, \hat{x}_2, \hat{x}_3 $$为什么一次前向能得到多个位置的信息
这是 Transformer 的基本能力。因为在 causal mask 下,一次前向就可以同时输出每个位置的 logits。
如果输入序列是:
$$ [y, \hat{x}_1, \hat{x}_2, \hat{x}_3] $$那么输出时会同时得到这几个位置对应的条件分布:
- 第 1 个位置对应
- 第 2 个位置对应
- 第 3 个位置对应
- 再往后一个位置对应
所以一次前向,其实已经把“如果沿着这条草稿路径继续往下走,每一步 target 会怎么看”都算出来了。
这点特别关键:
target 不是只算第一个草稿 token,而是把整条草稿前缀上的每个位置分布都并行算出来了。
这也是后面“一次 verify 可能接受多个 token”的基础。
验证的正常流程
假设 draft 提议了:
$$ \hat{x}_1, \hat{x}_2, \hat{x}_3 $$draft 还提供了每个位置自己对提议 token 的概率:
$$ p_d^1(\hat{x}_1),\quad p_d^2(\hat{x}_2),\quad p_d^3(\hat{x}_3) $$target 一次前向后,也拿到了对应位置对这些 token 的概率:
$$ p_t^1(\hat{x}_1),\quad p_t^2(\hat{x}_2),\quad p_t^3(\hat{x}_3) $$然后从左到右做接受测试。
第一步,检查第 1 个 token,接受概率是:
$$ \alpha_1 = \min\left(1, \frac{p_t^1(\hat{x}_1)}{p_d^1(\hat{x}_1)}\right) $$如果第 1 个就拒绝了:
- 后面 $\hat{x}_2, \hat{x}_3$ 全部作废
- 从第 1 位 residual 分布补一个 token
- 这一轮结束
如果第 1 个接受,再检查第 2 个,接受概率是:
$$ \alpha_2 = \min\left(1, \frac{p_t^2(\hat{x}_2)}{p_d^2(\hat{x}_2)}\right) $$如果第 2 个拒绝:
- $\hat{x}_3$ 作废
- 从第 2 位 residual 分布补一个 token
- 这一轮结束
如果前两个都接受,再检查第 3 个,接受概率是:
$$ \alpha_3 = \min\left(1, \frac{p_t^3(\hat{x}_3)}{p_d^3(\hat{x}_3)}\right) $$如果第 3 个也接受了,那这一轮就一口气接受了 3 个 draft token。
不过在很多具体实现里,这一轮通常还会顺手再从 target 的最后一个位置补采一个 next_token。所以实际提交回输出序列的,往往是“被接受的 draft 前缀 + 1 个 target token”。
为什么一次可以接受多个
现在就能看出来了:
一次 target 前向虽然只跑了一次,但它已经同时给出了:
- 第 1 位的 target 判断
- 第 2 位的 target 判断
- 第 3 位的 target 判断
所以如果这些位置都通过 acceptance 测试,那么:
- 第 1 个 token 合法
- 第 2 个 token 也合法
- 第 3 个 token 也合法
于是这一轮就能直接提交整个前缀:
$$ \hat{x}_1, \hat{x}_2, \hat{x}_3 $$这就是“一次计算接受多个 token”的本质:
不是一次前向直接“生成了多个 token”,而是一次前向验证了多个 draft token 是否都能作为 target 路径上的合法前缀继续保留。
标准验证公式是什么
设第 $k$ 个 draft token 是 $\hat{x}_k$,draft 和 target 在这个位置给出的概率分别是:
$$ p_d^k(\hat{x}_k), \quad p_t^k(\hat{x}_k) $$标准 speculative decoding 的接受概率是:
$$ \alpha_k = \min\left(1,\frac{p_t^k(\hat{x}_k)}{p_d^k(\hat{x}_k)}\right) $$直觉上:
- 如果 target 比 draft 更喜欢这个 token,就直接接受
- 如果 draft 比 target 更激进,就按比例接受
如果拒绝了,那么就由 target 对应的 residual 分布里采样:
$$ r_k(v)\propto \max\left(p_t^k(v)-p_d^k(v), 0\right) $$这里的 $v$ 表示“词表中的任意 token”,不是某一个已经提议出来的 token。也就是说,residual 是一个全词表分布,不是只在 $\hat{x}_k$ 上做修补。
上面的 $\propto$ 表示“正比于”,也就是这还是一个未归一化的分布。真正采样前还要除以总和,变成:
$$ r_k(v)=\frac{\max\left(p_t^k(v)-p_d^k(v), 0\right)}{\sum_u \max\left(p_t^k(u)-p_d^k(u), 0\right)} $$也就是说:对那些 target 比 draft 更喜欢的 token,把它们还没覆盖到的概率质量补回来,最后再归一化成合法分布。
这里还要特别注意一点:
接受规则是按前缀逐位置执行的,但 target 对这些位置的条件分布可以在一次前向中并行算出来。
为什么理论上输出与 target model 直接输出分布一致
这里只分析单个位置就行,因为每个位置的逻辑是一样的。
什么叫分布一致
如果 target model 在某个位置真正定义的分布是:
$$ p_t(x) $$那么标准 lossless speculative decoding 的目标就是让最终输出仍然满足:
$$ P(\text{final}=x)=p_t(x) $$它是怎么实现的
最终分布由两部分组成:
- draft 提案并被接受 的那部分
- reject 后由 residual 补齐 的那部分
对某个 token $x$ 来说,直接接受贡献的概率质量是:
$$ p_d(x)\cdot \min\left(1,\frac{p_t(x)}{p_d(x)}\right)=\min(p_d(x),p_t(x)) $$而 residual 部分要补的缺口则是:
$$ p_t(x)-\min(p_d(x),p_t(x))=\max(p_t(x)-p_d(x),0) $$也就是说:
- draft 先覆盖一部分 target 概率质量
- 覆盖不到的那部分,再由 residual 分布补齐
两部分合起来,最后正好恢复成 target 分布。
一个数值例子
假设某一位置上,target 只有两个候选 token:$A$ 和 $B$,它们的概率是:
$$ p_t(A)=0.8,\quad p_t(B)=0.2 $$draft 分布是:
$$ p_d(A)=0.5,\quad p_d(B)=0.5 $$第一步:draft 先提案
draft 会先从自己的分布里抽一个 token:
- 抽到 $A$ 的概率是 $0.5$
- 抽到 $B$ 的概率是 $0.5$
第二步:acceptance rule
对 $A$:
$$ \alpha(A)=\min(1,0.8/0.5)=1 $$对 $B$:
$$ \alpha(B)=\min(1,0.2/0.5)=0.4 $$所以:
- 如果 draft 抽到 $A$,一定接受
- 如果 draft 抽到 $B$,只以 $0.4$ 的概率接受,以 $0.6$ 的概率拒绝
第三步:先只算“直接接受”这条分支
$A$ 通过接受分支进入最终输出的概率是:
$$ 0.5\times 1=0.5 $$$B$ 通过接受分支进入最终输出的概率是:
$$ 0.5\times 0.4=0.2 $$所以到这里为止,接受分支已经贡献了:
- $A$:$0.5$
- $B$:$0.2$
但 target 真正想要的是:
- $A$:$0.8$
- $B$:$0.2$
因此当前还缺:
- $A$:$0.3$
- $B$:$0$
第四步:reject 分支总共有多少概率质量
reject 分支只会在“draft 抽到 $B$ 且 $B$ 被拒绝”时触发。
它的总概率是:
$$ 0.5\times (1-0.4)=0.5\times 0.6=0.3 $$也就是说,residual 分支总共会携带 $0.3$ 的概率质量。
第五步:residual distribution 怎么分配这 $0.3$
residual 分布定义为:
$$ r(x)\propto \max(p_t(x)-p_d(x),0) $$所以:
对 $A$:
$$ \max(0.8-0.5,0)=0.3 $$对 $B$:
$$ \max(0.2-0.5,0)=0 $$因此 residual 的未归一化权重是:
- $A$:$0.3$
- $B$:$0$
归一化后得到:
$$ r(A)=1,\quad r(B)=0 $$也就是说,只要进入 residual 分支,就一定会补出 $A$。
第六步:把 residual 分支加回总分布
因为 residual 分支的总概率质量是 $0.3$,且它一定输出 $A$,所以它对最终结果的贡献是:
- $A$:$0.3$
- $B$:$0$
现在把两条分支加起来:
- $A$:$0.5+0.3=0.8$
- $B$:$0.2+0=0.2$
所以最终恢复成:
$$ P(\text{final}=A)=0.8,\quad P(\text{final}=B)=0.2 $$正好等于 target 分布。
并行 draft 是什么
在继续看 DSpark 之前,最好先单独理解一下“并行 draft”这个概念。因为 DSpark 的第一部分创新,正是建立在并行 draft 的优点和缺点之上。
并行 draft 与 MTP(Multi-Token Prediction)相关,但不要直接划等号:MTP 更偏训练/多 token 预测范式;并行 draft 是推理侧“一次给出多个位置草稿”的提案方式。DeepSeek 线上基线里的
MTP-1是另一条对照线。
它和普通自回归有什么本质区别
普通自回归语言模型做的是链式生成:
$$ p(x_{1:B}\mid y)=\prod_{k=1}^{B} p(x_k\mid y,x_{- $y$ 是当前已经接受的前缀
- $x_1,\ldots,x_B$ 是后面准备生成的 token
这个公式的意思是:
- 第 1 位要先生成
- 第 2 位必须建立在第 1 位已经确定的基础上
- 第 3 位又要建立在前两位已经确定的基础上
所以标准自回归天然是串行的。
并行 draft 则故意放松了这个要求。它不再严格建模:
$$ p(x_k\mid y,x_{它们都共享同一个前缀 $y$,但后面位置不显式依赖前面最终采样出来的 token。
这也是并行 draft 为什么快的原因:它可以一次前向,直接给出 block 内多个位置的 hidden states 和 logits。
DSpark 做了什么
DSpark 是 DeepSeek 的一套“先让小草稿模块多写几个 token,再让大模型一次性验证”的推理加速系统;它不是提升模型智力的训练方法,而是提升模型服务速度和吞吐的推理系统设计。
也可以粗略理解成:DSpark 是挂在 DeepSeek 模型上的一套 speculative decoding 模块。
1. 对于 draft 阶段的问题:semi-autoregressive
有一类并行 draft 方法会一次性预测多个 token,速度很快,但因为这些 token 之间缺少充分的前后依赖,越往后越容易写歪。论文里把这个现象叫作 suffix decay,可以理解为“草稿后半段接受率塌得很快”。
DSpark 的做法是:
- 前面仍然保留一个重并行的 draft backbone,保证起草速度快
- 后面再接一个轻量的顺序输出头,给草稿 token 注入局部依赖关系
所以它不是纯并行,也不是完全自回归,而是论文说的 semi-autoregressive。
具体而言:并行 backbone 并没有直接最终决定每个 token 是什么,它只是先产出每个位置的“底稿表示 / base logits”;真正最终采样前,顺序头还可以逐位置改写这些 logits。
可以把它理解成:
并行 backbone 先给每个位置各出一份“初稿分布”
顺序头按
位置 1 -> 位置 2 -> 位置 3 ...依次读取“前一个已经采样出的 token”对当前位置 logits 加一个 bias / 修正项
然后才从修正后的 logits 里采样当前 token
2. 对于 verify 阶段的问题:confidence-scheduled verification
如果每次都把整段草稿全部丢给大模型验证,那么在高并发场景下,很多“本来就大概率会被拒”的尾部 token 也会占掉大模型宝贵的 batch 容量。结果不是更快,而是整体吞吐下降。
DSpark 在这里做了第二个关键设计:confidence-scheduled verification。
DSpark 在 draft 阶段不只产出 draft token,还会额外产出每个位置的一个 confidence。
这个 confidence 可以理解成:
在前面都被接受的前提下,这个位置继续被 target 接受的概率大概有多高。
然后再根据前面的接受情况、confidence,以及当前引擎的吞吐特征和负载水平,动态决定:
- 这次到底验证到第几个 token 为止
- 哪些低置信度尾部 token 先别验证,避免浪费大模型算力
一句话概括就是:
先用便宜的 confidence 估计,筛掉不值得 target 出手的尾部 token。
效果(论文/官方宣称):
按 DSpark 论文公开数字:在 DeepSeek-V4 线上服务里,相比生产基线 MTP-1,在相同吞吐水平下,单用户生成速度大约提升:
V4-Flash:60% 到 85%V4-Pro:57% 到 78%
核心代码
我自己读代码时,大致按下面这条链路往下看:
- 先看
generate_decoding_sample:它是总控循环,负责把draft -> verify -> commit -> update串起来 - 再看
build_dspark_proposal:它负责这一轮 draft 怎么起草、怎么裁掉低置信度尾部 - 然后看
sample_draft_tokens和markov_head.sample_block_tokens:它们负责把“并行初稿”变成“半自回归草稿” - 最后看
verify_draft_tokens:它负责 target 验证、最长前缀接受、reject 后 residual 补采样
也就是说,这几段代码分别回答的是四个不同问题:
- 整轮循环怎么跑
- 一轮 draft 怎么产生 proposal
- proposal 里的 token 是怎么采样出来的
- target 最后怎么裁决
1. base_evaluator.py (line 308) 的 generate_decoding_sample
我会先看这一层,因为它管的是整轮 speculative decoding 怎么循环。
先把它当成一个总控框架来看就行:
- 先让 target 正常吐出第一个 token,完成 prefill 之后的起点初始化
- 调
init_context(...)初始化 draft 侧上下文 - 反复调用
propose(...)生成一轮草稿 proposal - 调
verify_draft_tokens(...)让 target 做裁决 - 把被接受的 token 正式写回输出序列
- 调
update(...)更新 draft 侧状态,进入下一轮
所以这一层我不会先抠某一行怎么算,而是先抓一个整体印象:
DSpark 并没有改掉 speculative decoding 的大框架,它只是把自己的 draft 和 confidence 逻辑,作为
propose / update / post_verify这些钩子接进了通用循环。
展开代码:generate_decoding_sample
@torch.inference_mode() #PyTorch 的推理模式
def generate_decoding_sample(
*,
target_model,
input_ids: torch.Tensor,
max_new_tokens: int,
max_proposal_tokens: int,
temperature: float,
stop_token_ids: list[int] | None,
init_context: Callable[..., Any],
propose: Callable[..., DraftProposal],
update: Callable[[Any, VerificationResult], None],
post_verify: Callable[[DraftProposal, VerificationResult], None] | None = None,
) -> SimpleNamespace:
# 通用的Speculative-decoding loop,任何Speculative-decoding算法都可以接入
"""Speculative-decoding loop.
`init_context(initial_output, output_ids, position_ids, num_input_tokens)` builds the algorithm-specific state once after prefill. //不同 draft 算法需要保存的东西不一样,这里就初始化对应的状态
`propose(context, output_ids, position_ids, start, stop_token_ids)` returns the next DraftProposal.
`update(context, verification)` advances the state when the loop continues.
`post_verify` is an optional diagnostic hook called after every verification (used for confidence calibration).
"""
"""
实现了统一流程:
初始化
提 proposal
verify
更新
循环
"""
assert max_proposal_tokens >= 1
assert input_ids.size(0) == 1, "only bsz=1 is supported"
device = input_ids.device
num_input_tokens = input_ids.shape[1]
max_length = num_input_tokens + int(max_new_tokens)
output_ids = torch.empty(
(1, max_length + max_proposal_tokens + 1),
dtype=torch.long,
device=device,
) # 显存优化操作:预先分配一块足够大的空白 Tensor 用来存放整个推理过程生成的所有 Token(包括最后一步投机可能越界的多余 Token),避免每次生成新字符时都去动态 torch.cat 导致显存碎片和性能开销。
position_ids = torch.arange(output_ids.shape[1], device=device).unsqueeze(0)
past_key_values_target = DynamicCache()
# 做 prefill,先让 target 正常吐出一个 token
output = target_model(
input_ids=input_ids,
position_ids=position_ids[:, :num_input_tokens],
past_key_values=past_key_values_target,
use_cache=True,
output_hidden_states=True,
logits_to_keep=1, #指定只保留最后几个位置的 Logits,这里是1
)
# 把prompt填入output_ids,然后从 target 的输出分布采样出第一个新 token
output_ids[:, :num_input_tokens] = input_ids
output_ids[:, num_input_tokens : num_input_tokens + 1] = sample_from_probs(
logits_to_probs(output.logits, float(temperature))
)
# start 标记当前最新已被验证并采纳的 Token 的索引位置
# (此时正好在 Prompt 结束、新生成 Token 的地方)。
# 后面三个列表分别用于记录每一轮最终大模型的接纳长度、
# 小模型的提议长度、小模型被接纳的纯草稿长度(用于后期统计效率指标)。
# 其中 acceptance_lengths 和 accepted_draft_lengths 的区别是:
# 无论小模型提议的草稿是对是错,大模型在验证时都会在小模型被接纳的草稿向后多预测一个 Token
# (跑了一次的结果得用上啊),除非到结束符了。
start = input_ids.shape[1]
acceptance_lengths: list[int] = []
proposal_lengths: list[int] = []
accepted_draft_lengths: list[int] = []
initial_token = output_ids[:, num_input_tokens : num_input_tokens + 1]
if has_stop_token(initial_token, stop_token_ids):
output_ids = output_ids[:, : num_input_tokens + 1]
output_ids = trim_output_ids(output_ids, num_input_tokens, stop_token_ids)
return SimpleNamespace(
output_ids=output_ids,
num_input_tokens=num_input_tokens,
num_output_tokens=output_ids.shape[1] - num_input_tokens,
acceptance_lengths=acceptance_lengths,
proposal_lengths=proposal_lengths,
accepted_draft_lengths=accepted_draft_lengths,
verify_count=0,
)
# 初始化循环状态
context = init_context(
initial_output=output,
output_ids=output_ids,
position_ids=position_ids,
num_input_tokens=num_input_tokens,
)
# Speculative Decoding 核心循环
while start < max_length:
# draft 进行提议
proposal = propose(
context=context,
output_ids=output_ids,
position_ids=position_ids,
start=start,
stop_token_ids=stop_token_ids,
)
# target 进行验证
"""
这里的verify_draft_tokens的返回值为
return VerificationResult(
target_output=target_output,
target_probs=target_probs,
accept_prefix_mask=accept_prefix_mask,
accepted_draft_tokens=accepted_draft_tokens,
next_token=next_token,
effective_proposal_length=effective_proposal_length,
terminated_by_stop_token=terminated_by_stop_token,
committed_tokens=committed_tokens,
)
"""
verification = verify_draft_tokens(
target_model=target_model,
proposal=proposal,
position_ids=position_ids,
start=start,
past_key_values_target=past_key_values_target,
temperature=temperature,
max_proposal_tokens=max_proposal_tokens,
current_token_ids=output_ids[:, start : start + 1],
stop_token_ids=stop_token_ids,
)
# 如果有传入辅助性的诊断函数,在这步触发。
if post_verify is not None:
post_verify(proposal, verification)
# 更新内存:把已经被大模型认可通过的小模型草稿片段,正式写回到 output_ids 序列中。
proposal_lengths.append(int(verification.effective_proposal_length))
accepted_draft_tokens = int(verification.accepted_draft_tokens)
accepted_draft_lengths.append(accepted_draft_tokens)
output_ids[:, start : start + accepted_draft_tokens + 1] = (
proposal.verify_input_ids[:, : accepted_draft_tokens + 1]
)
# 判断是否碰到结束符,碰到就终止循环
if verification.terminated_by_stop_token:
acceptance_lengths.append(accepted_draft_tokens)
start += accepted_draft_tokens
past_key_values_target.crop(start)
break
# 没碰到就在被接收的draft后面,再补上target修正/产出的那个 next_token(这也是为什么每轮哪怕speculation草稿全错,大模型也必定能保底前进一步)
output_ids[:, start + accepted_draft_tokens + 1] = verification.next_token
new_token_ids = output_ids[:, start + 1 : start + accepted_draft_tokens + 2]
acceptance_lengths.append(accepted_draft_tokens + 1)
start += accepted_draft_tokens + 1
past_key_values_target.crop(start) # KV 缓存裁剪:由于大模型在验证时连同小模型那些错误的草稿也一起计算了 KV 并塞进了缓存,此时必须根据最终实际接纳的正确长度start把后面错误的缓存裁掉
update(context, verification)
if has_stop_token(new_token_ids, stop_token_ids):
break
# 利用start将真正有效的Token序列截取出来
output_ids = output_ids[:, : min(start + 1, max_length)]
output_ids = trim_output_ids(output_ids, num_input_tokens, stop_token_ids)
# 返回并封装
return SimpleNamespace(
output_ids=output_ids,
num_input_tokens=num_input_tokens,
num_output_tokens=output_ids.shape[1] - num_input_tokens,
acceptance_lengths=acceptance_lengths,
proposal_lengths=proposal_lengths,
accepted_draft_lengths=accepted_draft_lengths,
verify_count=len(proposal_lengths),
)2. draft_ops.py (line 96) 的 build_dspark_proposal
顺着总控循环往下看,build_dspark_proposal 就比较好理解了:它回答的是“DSpark 这一轮怎么起草 block”。
按我的理解,这段逻辑大致就是四步:
- 从
block_hidden里取出本轮 proposal 的 hidden states - 先用
compute_logits(...)得到每个位置的base_draft_logits - 再用
sample_draft_tokens(...)采样出草稿 token;如果开了markov_head,这里就会做半自回归修正 - 如果开了
confidence_head,再预测每个位置的 confidence,并把低置信度尾部截掉
所以这段代码的核心不是“直接产出整段草稿”,而是:
先起一版并行初稿,再根据 markov 信息修正,再根据 confidence 决定这次到底把多长前缀送去 verify。
假设大模型当前已经确认的基准 token 是 [A],而小模型这一轮采样出的草稿是 [B, C, D]。
那么这里的 prev_token_ids 可以理解成把“起点 token”和“除最后一个之外的草稿 token”拼起来,也就是 [A, B, C]。这样每个位置在预测 confidence 时,都能看到自己左边那个已经确定下来的 token。
展开代码:build_dspark_proposal
# 结合小模型在每个格子的隐状态,以及小模型刚刚自己采样出来的草稿 Token,去算出一个“自信心分数矩阵”。用这个分数来评估小模型对自己刚写出来的草稿到底有多大的把握。
def _predict_confidence_logits(
model: DSparkModel,
*,
proposal_hidden_states: torch.Tensor,
draft_input_ids: torch.Tensor,
sampled_tokens: torch.Tensor,
block_size: int,
) -> torch.Tensor | None:
# 先前一个token拼接上抛弃draft最后一个token的序列
prev_token_ids = torch.cat(
[draft_input_ids[:, :1], sampled_tokens[:, :-1]],
dim=1,
)
# 上面的prev_token_ids 和 proposal_hidden_states一一对应的时候是错开一位的,比如proposal_hidden_states的第二个对应的是预测的第一个的token,这样就实现了前后关系对应,具体上面有例子
confidence_pred = model.predict_confidence_step(
proposal_hidden_states,
prev_token_ids=prev_token_ids,
)
if confidence_pred is None:
return None
return confidence_pred.float().reshape(
confidence_pred.shape[0],
block_size,
-1,
)[:, :, 0]
def _confident_prefix_length(
confidence_logits: torch.Tensor,
*,
block_size: int,
threshold: float,
) -> int:
if threshold <= 0.0:
return int(block_size)
below_threshold = confidence_logits.sigmoid() < threshold
if not bool(below_threshold[0].any().item()):
return int(block_size)
return int(torch.nonzero(below_threshold[0], as_tuple=False)[0].item())
def build_dspark_proposal(
model: DSparkModel,
*,
draft_input_ids: torch.Tensor,
block_hidden: torch.Tensor,
block_size: int,
temperature: float,
confidence_threshold: float,
) -> DSparkDraftProposal:
assert draft_input_ids.size(0) == 1, "build_dspark_proposal requires batch_size=1"
# draft 提议的每个位置的 hidden_states
proposal_hidden_states = block_hidden[:, :block_size, :]
base_draft_logits = model.compute_logits(proposal_hidden_states)
sampled_tokens, draft_logits = model.sample_draft_tokens(
base_draft_logits,
first_prev_token_ids=draft_input_ids[:, 0],
temperature=temperature,
hidden_states=proposal_hidden_states,
)
# DSpark 的关键技术:动态置信度动态裁剪
proposal_draft_tokens = int(block_size)
confidence_logits = None
if model.confidence_head is not None:
# 如果有confidence_head,则调用_predict_confidence_logits
confidence_logits = _predict_confidence_logits(
model,
proposal_hidden_states=proposal_hidden_states,
draft_input_ids=draft_input_ids,
sampled_tokens=sampled_tokens,
block_size=block_size,
)
if confidence_logits is None:
return _empty_dspark_proposal(draft_input_ids)
# 将confidence_logits经过一个sigmoid,然后超过阈值的才被验证
proposal_draft_tokens = _confident_prefix_length(
confidence_logits,
block_size=block_size,
threshold=float(confidence_threshold),
)
if proposal_draft_tokens == 0:
return _empty_dspark_proposal(draft_input_ids)
verify_input_ids = torch.cat(
[draft_input_ids[:, :1], sampled_tokens[:, :proposal_draft_tokens]],
dim=1,
)
draft_probs = logits_to_probs(
draft_logits[:, :proposal_draft_tokens, :],
temperature,
)
return DSparkDraftProposal(
draft_token_count=proposal_draft_tokens,
verify_input_ids=verify_input_ids,
draft_probs=draft_probs,
confidence_logits=(
confidence_logits[:, :proposal_draft_tokens]
if confidence_logits is not None
else None
),
)其中 predict_confidence_step 的实现在这里:modeling.py (line 293)
这一小段我会和上面的 _predict_confidence_logits(...) 放在一起看。前者做的是“给定当前位置 hidden state,以及可选的前一个 token 信息,输出一个 confidence logit”;后者做的是“把整段 block 每个位置的 confidence 都算出来,再整理成按位置对齐的矩阵”。
展开代码:predict_confidence_step
# 用一个很轻的线性 AcceptRatePredictor,根据每个 draft 位置的 hidden state,以及可选的前一个 token embedding,输出该位置的 confidence logit;推理时再经过 sigmoid,用来判断这一位是否值得继续送进 verify。
def predict_confidence_step(
self,
hidden_states,
prev_token_ids=None,
):
if self.confidence_head is None:
return None
if self.confidence_head_with_markov:
prev_embeddings = self.markov_head.get_prev_embeddings(prev_token_ids)
features = torch.cat([hidden_states, prev_embeddings], dim=-1)
return self.confidence_head(features).float()
return self.confidence_head(hidden_states).float()其中 sample_draft_tokens 在 modeling.py (line 310)
sample_draft_tokens 这一步挺关键,因为它正好把“普通并行 draft”和“DSpark 的 semi-autoregressive draft”分开了:
- 如果
self.markov_head is None,就直接从base_logits并行采样 - 如果
self.markov_head is not None,就进入markov_head.sample_block_tokens(...),逐位置读取前一个已采样 token,对当前 logits 做修正后再采样
所以 sample_draft_tokens(...) 本身像一个分发器:决定这一轮草稿到底走“纯并行”还是“半自回归修正”。
展开代码:sample_draft_tokens
def sample_draft_tokens(...):
batch_size, proposal_len = base_logits.shape[:2]
if proposal_len == 0:
...
if self.markov_head is None:
return sample_tokens(base_logits, temperature), base_logits
return self.markov_head.sample_block_tokens(
base_logits,
first_prev_token_ids=first_prev_token_ids,
hidden_states=hidden_states,
temperature=temperature,
)具体的 markov_head.sample_block_tokens 在 markov_head.py (line 55)
我看这个函数时,主要盯循环里的三件事:
prev_token_ids一开始取的是本轮起点前一个已确认 token- 每一步都会用
apply_step_logits(...)把当前位置base_logits修正成step_logits - 当前步采样出来的
next_token_ids,会立刻变成下一步的prev_token_ids
这就是 semi-autoregressive 的核心。也就是说:
backbone 负责一次性给出整段位置的“底稿”,markov head 再在 block 内部做一个很轻的左到右串行修正。
展开代码:markov_head.sample_block_tokens
sampled_tokens = []
corrected_logits = []
prev_token_ids = first_prev_token_ids.long()
for step_idx in range(proposal_len):
step_hidden = hidden_states[:, step_idx, ...]
# 用前一个位置对这一位做修正 最终 logits = 初稿 logits + 顺序修正 bias
step_logits = self.apply_step_logits(
base_logits[:, step_idx, :],
token_ids=prev_token_ids,
hidden_states=step_hidden,
)
corrected_logits.append(step_logits.unsqueeze(1))
next_token_ids = sample_tokens(step_logits.unsqueeze(1), temperature=temperature).squeeze(1)
sampled_tokens.append(next_token_ids)
prev_token_ids = next_token_ids
return torch.stack(sampled_tokens, dim=1), torch.cat(corrected_logits, dim=1)3. base_evaluator.py (line 186) 的 verify_draft_tokens
再往后看 verify_draft_tokens。这段代码回答的是“target 怎么验证、接受、拒绝,以及 residual 怎么补采样”。
按我的理解,这里的重点可以抓成五步:
- 先检查 proposal 的长度、起始 token 是否合法
- 把“当前已确认 token + draft 提议整段”一次送进 target,得到整段位置的 target logits
- 抽取出 target 和 draft 对这些提议 token 的对应概率,计算逐位置接受概率
- 用
cumprod把逐位置接受变成“最长前缀接受” - 如果中途拒绝,就从 residual 分布补一个
next_token;如果全接受,就直接从 target 的最后一个位置继续采样
所以这段代码最重要的不是某个公式单独成立,而是它把前面讲过的理论完整落到了实现上:
- “一次前向算多个位置” 对应
target_model(...) - “逐位置接受” 对应
selected_target_probs / selected_draft_probs - “最长前缀接受” 对应
accept_mask.cumprod(dim=1) - “reject 后补齐缺口” 对应
sample_residual(...)
展开代码:verify_draft_tokens
def verify_draft_tokens(
*,
target_model,
proposal: DraftProposal,
position_ids: torch.Tensor,
start: int,
past_key_values_target: DynamicCache,
temperature: float,
max_proposal_tokens: int,
current_token_ids: torch.Tensor | None = None,
stop_token_ids: list[int] | None = None,
) -> VerificationResult:
"""Verify draft tokens with the target model and rejection sampling."""
if proposal.draft_token_count > max_proposal_tokens:
raise ValueError(
"DraftProposal.draft_token_count must not exceed "
f"max_proposal_tokens={max_proposal_tokens}, "
f"got {proposal.draft_token_count}."
)
if current_token_ids is not None and not torch.equal(
proposal.verify_input_ids[:, :1],
current_token_ids,
):
raise ValueError(
"DraftProposal.verify_input_ids must start with the current "
"accepted token."
)
# 计算本次需要验证的总长度 verify_length(等于草稿长度 + 1 个基准 Token)。并从全局位置编码中切出对应这段区间的位置 ID verify_position_ids
draft_token_count = int(proposal.draft_token_count)
verify_length = draft_token_count + 1
verify_position_ids = position_ids[:, start : start + verify_length]
# 核心并行计算:把大模型上一轮最后一个正确 Token 和小模型新提议的K个草稿拼接成的整个序列,一次性全部喂给大模型。大模型在内部会并行计算并追加这一整段的KV缓存。
target_output = target_model(
input_ids=proposal.verify_input_ids,
position_ids=verify_position_ids,
past_key_values=past_key_values_target,
use_cache=True,
output_hidden_states=True,
)
if target_output.logits.ndim != 3:
raise ValueError(
"target model must return rank-3 logits [B, S, V], "
f"got ndim={target_output.logits.ndim}."
)
target_probs = logits_to_probs(target_output.logits, float(temperature))
if (
draft_token_count > 0
and proposal.draft_probs is not None
and proposal.draft_probs.size(-1) != target_probs.size(-1)
):
raise ValueError(
"DraftProposal.draft_probs vocab size must match target logits, "
f"got {proposal.draft_probs.size(-1)} and {target_probs.size(-1)}."
)
accept_prefix_mask = None
if draft_token_count > 0:
assert proposal.draft_probs is not None
proposed_tokens = proposal.verify_input_ids[:, 1:]
# 获取大模型对草稿选出的token id的预测概率
selected_target_probs = gather_token_probs(
target_probs[:, :-1, :],# target会在draft后多预测一个,这里要去掉
proposed_tokens,
)
# 获取小模型对自身草稿选出的token id的预测概率
selected_draft_probs = gather_token_probs(
proposal.draft_probs,
proposed_tokens,
).clamp_min(1e-8)
# 计算接受概率(也就是上面提到的公式)
accept_prob = torch.clamp(
selected_target_probs / selected_draft_probs,
max=1.0,
)
accept_mask = (torch.rand_like(accept_prob) < accept_prob).to(torch.int64) # 生成accept_prob大小的0~1随机浮点数,并和接受概率accept_prob比较
accept_prefix_mask = accept_mask.cumprod(dim=1) # cumprod(dim=1) 是累乘(Cumulative Product),会沿着序列方向,让前面的结果和后面一直相乘,也就是只要有一个0后面就全是0。
accepted_draft_tokens = int(accept_prefix_mask.sum(dim=1)[0].item())
else:
accepted_draft_tokens = 0
# 检查大模型接受的草稿Token内部,是否包含了结束符
effective_proposal_length = draft_token_count
terminated_by_stop_token = False
if stop_token_ids and accepted_draft_tokens > 0:
accepted_slice = proposal.verify_input_ids[0, 1 : accepted_draft_tokens + 1]
stop_tensor = torch.tensor(
stop_token_ids,
device=accepted_slice.device,
dtype=accepted_slice.dtype,
)
eos_hits = torch.isin(accepted_slice, stop_tensor).nonzero(as_tuple=True)[0]
if eos_hits.numel() > 0:
eos_pos = int(eos_hits[0].item())
accepted_draft_tokens = eos_pos + 1
effective_proposal_length = eos_pos + 1
terminated_by_stop_token = True
if 0 < draft_token_count and accepted_draft_tokens < draft_token_count:
assert proposal.draft_probs is not None
# 如果草稿有错被拒绝了,利用Residual Sampling采样出next_token,Residual Sampling保证了分布一致
next_token = sample_residual(
target_probs[:, accepted_draft_tokens, :],
proposal.draft_probs[:, accepted_draft_tokens, :],
)
else:
# 如果都接受了,则直接采样下一个
next_token = sample_from_probs(target_probs[:, -1:, :]).squeeze(1)
committed_tokens = torch.cat(
[
proposal.verify_input_ids[:, 1 : accepted_draft_tokens + 1],
next_token.unsqueeze(1),
],
dim=1,
)
return VerificationResult(
target_output=target_output,
target_probs=target_probs,
accept_prefix_mask=accept_prefix_mask,
accepted_draft_tokens=accepted_draft_tokens,
next_token=next_token,
effective_proposal_length=effective_proposal_length,
terminated_by_stop_token=terminated_by_stop_token,
committed_tokens=committed_tokens,
)