上周出去玩了,拖了一周终于简单看了一遍 DeepSeek 的 DSpark,顺手整理成一篇笔记,之前没有了解过这方面,所以还是看了两三天。

它不是提升模型“智力”的训练方法,而是一套偏推理系统侧的加速设计:目标是在尽量不破坏 target model 输出分布的前提下,让同一个模型在推理时更快地产生 token。

相关仓库:DeepSpec

提示

这篇主要想讲清楚两件事:一是标准 speculative decoding 的 verify 机制到底怎么工作;二是 DSpark 在 draft 和 verify 两个阶段分别改了什么。

speculative decoding 背景

大语言模型的标准生成方式是自回归解码:每次只生成一个新 token。这样做最稳,但也带来一个很现实的问题:

  • 模型越大,每一步 forward 越贵
  • 回答越长,要走的步数越多
  • 在线服务里用户一多,延迟和吞吐压力就会很明显

所以推理系统里一直有一个核心问题:

能不能在不破坏目标模型输出分布的前提下,让一次“确认”尽量多产出几个 token?

它的标准思路可以概括成两步:

  1. draft 先让一个更便宜、更快的 draft model 或 draft 模块预写后面几个 token。
  2. verify 再让真正的 target model 一次性验证这几个 token,接受能通过的最长前缀。

如果草稿质量足够高,那么 target model 一次验证就不只“确认 1 个 token”,而可能确认 2 个、3 个甚至更多 token。这样平均到每个 token 的延迟就会下降。

speculative decoding 之所以重要,不只是因为它快,而是因为经典路线追求的是 lossless / exact 加速,也就是:

  • 最终采样分布仍然和 target model 一致
  • 不是简单拿一个小模型直接替代大模型
  • 而是让小模型负责“提案”,大模型保留“裁决权”

所以它和“直接用弱一点的模型换速度”不是一个问题。

看起来简单,但实际会遇到很多问题,例如:

text
draft 太弱,接受率太低,白忙一场;
draft 太强,自己又变得很贵;
一次提太长,后缀 token 很容易被拒;
高并发时,验证过多低质量 token 反而拖垮吞吐;
不同任务域里接受率差异很大,代码、数学、闲聊的表现不一样。

所以后来的很多工作,本质上都在围绕三个变量做优化:

  1. draft 怎么设计得更好
  2. 一次 draft 多长更合适
  3. 哪些 token 值得送去 verify

DSpark 正是在这几个变量里,重点优化了前两个半:

  • 它让 draft 既保留并行速度,又补一点顺序依赖
  • 它让 verify 不再固定吃完整段,而是按置信度截断前缀

下面先把标准 speculative decoding 的 verify 机制讲清楚,再看 DSpark 到底改了什么。


普通自回归怎么做

假设当前已经有前缀:

$$ y $$

后面真正要生成的 token 是:

$$ x_1, x_2, x_3 $$

普通自回归的 target model 会这样做:

  1. 跑一次,得到 $p_t(\cdot \mid y)$,采样出 $x_1$
  2. 再跑一次,得到 $p_t(\cdot \mid y, x_1)$,采样出 $x_2$
  3. 再跑一次,得到 $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 个位置对应
$$ p_t(\cdot \mid y) $$
  • 第 2 个位置对应
$$ p_t(\cdot \mid y, \hat{x}_1) $$
  • 第 3 个位置对应
$$ p_t(\cdot \mid y, \hat{x}_1, \hat{x}_2) $$
  • 再往后一个位置对应
$$ p_t(\cdot \mid y, \hat{x}_1, \hat{x}_2, \hat{x}_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) $$

它是怎么实现的

最终分布由两部分组成:

  1. draft 提案并被接受 的那部分
  2. 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_{而是更像同时学一组“未来相对位置”的草稿分布:

$$ q_1(\cdot\mid y),\quad q_2(\cdot\mid y),\quad \ldots,\quad q_B(\cdot\mid y) $$

它们都共享同一个前缀 $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%

核心代码

我自己读代码时,大致按下面这条链路往下看:

  1. 先看 generate_decoding_sample:它是总控循环,负责把 draft -> verify -> commit -> update 串起来
  2. 再看 build_dspark_proposal:它负责这一轮 draft 怎么起草、怎么裁掉低置信度尾部
  3. 然后看 sample_draft_tokensmarkov_head.sample_block_tokens:它们负责把“并行初稿”变成“半自回归草稿”
  4. 最后看 verify_draft_tokens:它负责 target 验证、最长前缀接受、reject 后 residual 补采样

也就是说,这几段代码分别回答的是四个不同问题:

  • 整轮循环怎么跑
  • 一轮 draft 怎么产生 proposal
  • proposal 里的 token 是怎么采样出来的
  • target 最后怎么裁决

1. base_evaluator.py (line 308)generate_decoding_sample

我会先看这一层,因为它管的是整轮 speculative decoding 怎么循环。

先把它当成一个总控框架来看就行:

  1. 先让 target 正常吐出第一个 token,完成 prefill 之后的起点初始化
  2. init_context(...) 初始化 draft 侧上下文
  3. 反复调用 propose(...) 生成一轮草稿 proposal
  4. verify_draft_tokens(...) 让 target 做裁决
  5. 把被接受的 token 正式写回输出序列
  6. update(...) 更新 draft 侧状态,进入下一轮

所以这一层我不会先抠某一行怎么算,而是先抓一个整体印象:

DSpark 并没有改掉 speculative decoding 的大框架,它只是把自己的 draft 和 confidence 逻辑,作为 propose / update / post_verify 这些钩子接进了通用循环。

展开代码:generate_decoding_sample
python
@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”。

按我的理解,这段逻辑大致就是四步:

  1. block_hidden 里取出本轮 proposal 的 hidden states
  2. 先用 compute_logits(...) 得到每个位置的 base_draft_logits
  3. 再用 sample_draft_tokens(...) 采样出草稿 token;如果开了 markov_head,这里就会做半自回归修正
  4. 如果开了 confidence_head,再预测每个位置的 confidence,并把低置信度尾部截掉

所以这段代码的核心不是“直接产出整段草稿”,而是:

先起一版并行初稿,再根据 markov 信息修正,再根据 confidence 决定这次到底把多长前缀送去 verify。

假设大模型当前已经确认的基准 token 是 [A],而小模型这一轮采样出的草稿是 [B, C, D]

那么这里的 prev_token_ids 可以理解成把“起点 token”和“除最后一个之外的草稿 token”拼起来,也就是 [A, B, C]。这样每个位置在预测 confidence 时,都能看到自己左边那个已经确定下来的 token。

展开代码:build_dspark_proposal
python
# 结合小模型在每个格子的隐状态,以及小模型刚刚自己采样出来的草稿 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
python
# 用一个很轻的线性 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_tokensmodeling.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
python
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_tokensmarkov_head.py (line 55)

我看这个函数时,主要盯循环里的三件事:

  1. prev_token_ids 一开始取的是本轮起点前一个已确认 token
  2. 每一步都会用 apply_step_logits(...) 把当前位置 base_logits 修正成 step_logits
  3. 当前步采样出来的 next_token_ids,会立刻变成下一步的 prev_token_ids

这就是 semi-autoregressive 的核心。也就是说:

backbone 负责一次性给出整段位置的“底稿”,markov head 再在 block 内部做一个很轻的左到右串行修正。

展开代码:markov_head.sample_block_tokens
python
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 怎么补采样”。

按我的理解,这里的重点可以抓成五步:

  1. 先检查 proposal 的长度、起始 token 是否合法
  2. 把“当前已确认 token + draft 提议整段”一次送进 target,得到整段位置的 target logits
  3. 抽取出 target 和 draft 对这些提议 token 的对应概率,计算逐位置接受概率
  4. cumprod 把逐位置接受变成“最长前缀接受”
  5. 如果中途拒绝,就从 residual 分布补一个 next_token;如果全接受,就直接从 target 的最后一个位置继续采样

所以这段代码最重要的不是某个公式单独成立,而是它把前面讲过的理论完整落到了实现上:

  • “一次前向算多个位置” 对应 target_model(...)
  • “逐位置接受” 对应 selected_target_probs / selected_draft_probs
  • “最长前缀接受” 对应 accept_mask.cumprod(dim=1)
  • “reject 后补齐缺口” 对应 sample_residual(...)
展开代码:verify_draft_tokens
python
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,
    )