@_yucheng_lu: MTP 使自回归 LLM 变快。同样的技巧能否用于扩散语言模型?与 @modal 进行了一次有趣的合作,探索……
摘要
介绍了多令牌残差预测(MRP)技术,该技术通过预测相邻去噪步骤之间的残差来加速扩散语言模型推理,在 SGLang 中实现了最高 1.56 倍的速度提升,并在激进解码设置中恢复了最高 +16 的准确率点数。
查看缓存全文
缓存时间: 2026/07/03 02:31
MTP 让自回归语言模型变得更快。同样的技巧能用于扩散语言模型吗? 我与 @modal 进行了一次有趣的合作,正是围绕这一点:多令牌残差预测 (MRP)
关键变化:我们不再训练一个小型头部来预测下一个去噪步骤的完整分布,而是预测相邻步骤之间的残差。这是一个简单得多的目标,因此一个只有3层的小模块就能准确地学习它,并将其应用于多个步骤。
我们在两种场景下应用了 MRP: • 静态场景 → (近乎)无损加速,在 SGLang 中最高可达 1.56 倍。 • 动态场景 → 可恢复因激进的低阈值解码而损失的准确率,最高可达 +16 个百分点。
代码、SGLang 实现和模型均在博客中
多令牌残差预测 | Modal 博客
来源:https://modal.com/blog/multi-token-residual-prediction 返回 (https://modal.com/blog) 研究
2026 年 7 月 1 日 · 7 分钟阅读
编者按:这篇客座博文描述了 Modal Research 与上海纽约大学 HeavyBall Research 实验室 (https://www.yucheng-lu.me/lab.html) 之间的一项研究合作成果。
MRP 是一个拥有两种用途的小模块。左图:在静态场景中,它可以无损地加速解码(推测模式,前瞻步数 K = 3),或者以微小的质量代价进一步加速(直接模式,前瞻步数 K = 1),在 SGLang 中实现最高 1.56 倍的吞吐量。右图:在动态场景中,它可以恢复因激进的低阈值解码(阈值 τ = 0.5)而损失的准确率,最高可达 +16 个百分点。结果基于 SDAR-1.7B/4B/8B 模型在 GSM8K、MATH500、HumanEval 和 MBPP 数据集上的平均值。 TL;DR: 多令牌预测 (MTP) (https://arxiv.org/abs/2404.19737) 通过一次前向传播预测多个令牌来加速自回归模型 (Gloeckle et al., 2024)。我们将这个想法引入扩散语言模型,并做了一项使其生效的修改:我们不再训练一个小型头部来预测下一个去噪步骤的完整分布,而是训练它预测相邻步骤之间的残差。残差是一个简单得多的目标,因此一个微小的模块就能准确地预测它,并将其应用于多个步骤。然后,同一个模块可以服务于 DLM 推理中通常需要相互权衡的两种场景:在静态场景中,当必须保持输出质量时,它能带来(近乎)无损的加速(在 SGLang 中平均可达 1.56 倍);在动态场景中,它能恢复因激进的吞吐量设置而损失的大部分质量(平均最高 +16 个百分点)。
论文: https://arxiv.org/abs/2605.18817
代码: https://github.com/heavyball-research/multi-token-residual-prediction
SGLang 实现: https://github.com/heavyball-research/sglang
模型: https://huggingface.co/collections/heavyball/sdar-mrp
从 MTP 开始
如果你关注过快 LLM 推理,你一定知道多令牌预测 (MTP) 的故事。自回归模型每次前向传播生成一个令牌,这非常昂贵。因此,人们会附加一些轻量级头部,例如 Medusa (https://arxiv.org/abs/2401.10774)、EAGLE (https://arxiv.org/abs/2401.15077)、DeepSeek 的 MTP (https://arxiv.org/abs/2412.19437),它们会窥探主干的隐藏状态,并在一次前向传播中猜出接下来的几个令牌。再结合推测验证,我们就能获得真正的加速。
这是一个美妙的想法,并且行之有效。所以我们自然要问:我们能否让它超越自回归模型,在其他模型上也奏效?
扩散语言模型 (DLM) 是一个自然的尝试方向,因为它们并非从左到右解码。DLM 从完全掩码的序列开始,逐步去噪,每次少量地 unmask 高置信度的位置。并行性内置于该过程中,但有一个权衡:如果在单步中 unmask 过多的令牌,质量就会下降,因为每个令牌解码时并未看到与其他令牌同时被提交的情况。大多数 DLM 加速文献都处在这条帕累托曲线上,用质量换取速度。
这正是我们想要打破的权衡,而 MTP 看起来是合适的工具。如果一个轻量级头部能够根据主干的隐藏状态预测额外的令牌,那么它每次前向传播就能解码多个位置,同时保持每个位置对其他位置的感知,而不是简单地 unmask 更多位置并为此付出质量代价。问题在于 MTP 的配方能否直接迁移到扩散设置中。正如我们发现的那样,它不能直接迁移,而理解其原因正是我们提出自己方法的起点。
一次天真的尝试
我们首先训练一个小型头部,使其 непосредственно 根据当前步骤的隐藏状态预测下一步的完整 log 密度。运行几次,即可在每次主干前向传播中 unmask 多个令牌。
问题在我们要求头部执行超过一步的预测时就立刻出现了。蒸馏整个分布意味着每一步都必须从头复现一个庞大、高动态范围的目标,并且每一步的误差会累积。我们发现,这样训练的头部在执行一步预测时表现尚可,但到三步或四步时,在我们尝试的每个模型规模上都会崩溃。在 SDAR-4B 主干上,我们的直接蒸馏头部在一步预测时 GSM8K 准确率为 84.8%,但到四步时骤降至个位数。它根本无法在多次迭代中维持分布。
| 方法 | K=1 | K=2 | K=3 | K=4 |
|---|---|---|---|---|
| 朴素 MTP | 84.8 | 16.9 | 5.9 | 1.9 |
SDAR-4B 上的 GSM8K(0-shot,思维链)准确率。我们在此直接应用 MTP 配方,准确率在第一步 MTP 之后急剧下降。此处 K 表示预测步数。
主要洞察
上述失败同时也是一条线索。如果完整的下一步分布难以在多个步骤中预测,那么问题在于是否存在一个更容易预测的目标。答案是肯定的,这个目标来自观察主干输出在相邻两个去噪步骤之间的变化有多小。
这就是全部洞察,它改变了目标。我们不需要一个小型模块从头复现下一步的完整分布。我们需要它对一个已经相当不错的预测进行微小的修正。
这并非仅仅是幸运的经验事实。它源于去噪的马尔可夫结构:每一步只扰动少数几个位置,因此根据 Lipschitz 论证,未扰动位置上的预测分布只能移动有限的量,并且随着去噪过程的推进和模型变得更加自信,这个界限会收紧。我们要求一个小型模块学习的信号确实是低复杂度的。这正是一个小型模块可以学习它的原因。
多令牌残差预测 (MRP)
多令牌残差预测 (MRP) 是一个小型 Transformer(在我们的主要配置中为 3 层),附加在一个冻结的 DLM 主干上。它读取主干的隐藏状态,预测步骤间的 logit 残差,并将其添加到主干自身的 logits 中。主干、其 LM 头部和令牌嵌入都被冻结;只有 MRP 模块被训练。
训练目标是 MTP 蒸馏损失的残差版本。我们运行冻结的主干两次:一次在揭示一组令牌之前,一次在之后,并对仍被掩码的位置使用 KL 散度来训练 MRP,以最小化两次输出之间的差异。由于归一化常数在 softmax 下被抵消,这等价于匹配真实的条件分布,而模块本身只表示修正量。
| 方法 | K=1 | K=2 | K=3 | K=4 |
|---|---|---|---|---|
| 朴素 MTP | 84.8 | 16.9 | 5.9 | 1.9 |
| MRP | 88.6 | 84.9 | 70.9 | 57.2 |
SDAR-4B 上的 GSM8K(0-shot,思维链)准确率。使用辅助模块预测残差而非整个分布使得学习变得容易得多。此处 K 表示预测步数。残差框架的优势在 K 增大时显现出来。直接建模分布(朴素 MTP)与残差学习(MRP)之间的差距在 K = 1 时仅有几个百分点,但此后急剧扩大:在 K = 2 时,残差学习已经在 GSM8K 上领先 +65 个百分点,而直接变体在 K = 4 时完全崩溃。
在推理中的应用
通过一个在冻结主干上训练好的模块,我们可以廉价地近似计算下一步去噪会产生的结果。如何最佳地利用这个近似值取决于主干被解码的方式。
DLM 通常在两种场景下运行。在静态去噪中,每一步 unmask 固定且少量的位置;这能保持高质量但吞吐量低,是你希望输出正确时采用的场景。在动态去噪中,所有置信度超过阈值的位置会被一次性 unmask;这能提高吞吐量,但在低阈值下,主干每步会提交大量令牌,导致质量下降。
MRP 适用于这两种场景,但在其中扮演不同的角色:在静态场景中,它增加揭示次数以加快速度;而在动态场景中,它撤销过度的揭示以恢复质量。
应用一:静态去噪中的无损加速
在静态场景中,MRP 将推理变成了一个可调节的旋钮。同一个训练好的模块为你提供了一系列操作点,从完全匹配主干的输出到大幅加速且伴有微小的、可测量的质量成本,你可以根据应用的实际需求选择所处位置。
一端是推测解码,用于输出必须与主干产生结果完全匹配的场景。这里 MRP 充当起草者:它廉价地提出下一批令牌,而主干在一次前向传播中验证它们。主干同意草案的位置会被接受;不同意的主干会重新掩码并重新生成。由于扩散模型在一次前向传播中对每个位置进行评分,这种逐位置的验证非常自然。验证传递并未浪费:其隐藏状态和 logits 为下一次迭代提供了种子,因此当接受率高时,它同时充当了下一步的主干传递,其成本被摊销。
| 主干 | GSM8K | MATH500 | HumanEval | MBPP |
|---|---|---|---|---|
| SDAR-4B | 90.0 / 1.36x | 68.0 / 1.26x | 67.7 / 1.35x | 66.5 / 1.27x |
| SDAR-8B | 90.4 / 1.40x | 74.8 / 1.39x | 72.6 / 1.34x | 67.3 / 1.34x |
SGLang 中的推测模式。准确率 (%) 后跟相对于仅主干基线的吞吐量加速比。质量默认与主干一致。在单个 H100 上测量;我们在 SGLang 中提供了实现。
另一端是直接解码,它跳过验证,直接提交 MRP 修正后的 logits。验证需要一次完整的主干前向传播,因此跳过验证可以提高加速上限,并且残差预测本身已经足够准确,在推理任务上,质量损失很小:
| 设置 | GSM8K | MATH500 | HumanEval | MBPP |
|---|---|---|---|---|
| 基线 | 90.9 / 1x | 72.2 / 1x | 73.8 / 1x | 67.7 / 1x |
| 直接 (MRP Step 1) | 90.1 / 1.59x | 71.4 / 1.61x | 67.1 / 1.53x | 63.8 / 1.51x |
| 直接 (MRP Step 2) | 89.2 / 1.89x | 70.8 / 1.91x | 64.0 / 1.78x | 59.9 / 1.75x |
SDAR-8B 上的直接解码。准确率 (%) 后跟相对于主干速度的吞吐量加速比。令牌不经验证即被提交。K 设置每次主干前向传播中运行的 MRP 步数,因此它本身就是一个调节旋钮:K = 1 时,在 1.6 倍加速下,推理任务与主干相差一个百分点以内;K = 2 时,超过 1.8 倍,但代码任务下降更明显。较小的模型遵循相同模式(完整表格见论文 (https://arxiv.org/abs/2605.18817))。
关键是你选择操作点。当正确性不容妥协时,选择无损;当延迟占主导地位且可接受微小质量成本时,选择更快;可根据任务甚至每次请求进行调优。这种控制只有在你拥有自己应用的推理栈时才存在。在封闭的 API 背后,这种权衡由提供商替你决定:提供者固定了解码策略,你只能接受他们提供的成本-质量点。用自己的主干运行 MRP 将调节旋钮重新交回你手中。
应用二:质量恢复
现在考虑相反的场景:一个交互式环境,延迟最为重要,因此 unmask 阈值设置得很低,每步揭示大量令牌。在激进的阈值下,这会造成损害,因为主干在一次步骤中就提交了一批令牌,每个令牌在选择时都无法考虑到其他令牌同时被提交的影响。
在这里,MRP 反向运行。在主干因低阈值而过度揭示之后,一次单独的 MRP 传递根据那些新揭示的令牌预测残差,修正后的 logits 会重新评估刚刚提交的令牌。任何修正后置信度低于阈值的令牌会被重新掩码,并推迟到后续步骤,届时将有更多上下文。由于残差编码了每个预测在其新邻居被考虑后如何变化,MRP 能够识别出那些仅在孤立环境中才自信的揭示。相同的阈值 τ 同时控制揭示和重新掩码,因此无需引入额外的调优。
| 模型 | τ | GSM8K | MATH500 | HumanEval | MBPP |
|---|---|---|---|---|---|
| 1.7B | 0.5 | 41.6 → 59.1 (+17.5) | 26.0 → 37.4 (+11.4) | 17.7 → 28.7 (+11.0) | 26.9 → 41.3 (+14.4) |
| 1.7B | 0.6 | 56.3 → 67.0 (+10.7) | 33.4 → 40.4 (+7.0) | 31.7 → 43.3 (+11.6) | 42.4 → 49.0 (+6.6) |
| 1.7B | 0.7 | 65.4 → 71.8 (+6.4) | 39.4 → 48.6 (+9.2) | 40.9 → 45.1 (+4.2) | 49.8 → 51.0 (+1.2) |
| 1.7B | 0.8 | 70.6 → 75.4 (+4.8) | 47.4 → 52.0 (+4.6) | 45.7 → 48.8 (+3.1) | 51.8 → 51.8 (0.0) |
| 1.7B | 0.9 | 76.2 → 77.3 (+1.1) | 51.2 → 57.0 (+5.8) | 49.4 → 52.4 (+3.0) | 53.7 → 54.1 (+0.4) |
| 4B | 0.5 | 63.4 → 81.1 (+17.7) | 44.2 → 58.4 (+14.2) | 32.3 → 53.1 (+20.8) | 38.5 → 50.6 (+12.1) |
| 4B | 0.6 | 76.4 → 85.5 (+9.1) | 53.6 → 61.4 (+7.8) | 49.4 → 57.9 (+8.5) | 49.8 → 57.2 (+7.4) |
| 4B | 0.7 | 84.6 → 88.5 (+3.9) | 60.4 → 65.6 (+5.2) | 60.4 → 62.2 (+1.8) | 61.1 → 63.4 (+2.3) |
| 4B | 0.8 | 87.9 → 90.1 (+2.2) | 66.8 → 70.6 (+3.8) | 64.6 → 62.8 (−1.8) | 63.8 → 64.2 (+0.4) |
| 4B | 0.9 | 88.5 → 90.1 (+1.6) | 69.0 → 70.6 (+1.6) | 67.1 → 65.9 (−1.2) | 65.4 → 64.6 (−0.8) |
| 8B | 0.5 | 67.9 → 82.3 (+14.4) | 45.2 → 58.0 (+12.8) | 32.3 → 54.9 (+22.6) | 34.6 → 49.4 (+14.8) |
| 8B | 0.6 | 79.6 → 86.8 (+7.2) | 54.8 → 63.8 (+9.0) | 48.8 → 63.4 (+14.6) | 48.3 → 59.9 (+11.6) |
| 8B | 0.7 | 85.9 → 89.0 (+3.1) | 60.8 → 69.0 (+8.2) | 64.6 → 72.6 (+8.0) | 54.9 → 60.3 (+5.4) |
| 8B | 0.8 | 89.3 → 91.0 (+1.7) | 68.0 → 70.2 (+2.2) | 74.4 → 75.0 (+0.6) | 62.3 → 66.5 (+4.2) |
| 8B | 0.9 | 90.8 → 91.4 (+0.6) | 70.0 → 72.0 (+2.0) | 75.0 → 75.0 (0.0) | 66.9 → 68.5 (+1.6) |
MRP 重新掩码提高了低阈值动态解码的准确性。对于每个 SDAR 主干(1.7B / 4B / 8B)和 unmask 阈值 τ,每个单元格显示纯基于阈值的动态解码 → MRP 重新掩码后的结果。相同的 τ 同时控制揭示和重新掩码,因此无需额外调优。增益在激进(低)阈值下最大,此时主干过度提交,重新掩码有最多的撤销空间(在 τ = 0.5 时,8B HumanEval 上最高达 +22.6),并随着 τ 升高而趋于零,因为此时主干已经保守地 unmask。推理任务(GSM8K, MATH500)在每个操作点都有所改善;极少数的轻微退化限于高 τ 下的代码任务。准确率以 % 表示。
重新审视刚揭示的令牌并重新掩码那些不再成立的令牌的想法出现在之前的工作中,例如 DMax (https://arxiv.org/abs/2604.08302)、RCD (https://arxiv.org/abs/2601.22954) 和 WINO (https://arxiv.org/abs/2507.18578)。这里的新颖之处在于 MRP 用于做出该决策的信号。MRP 没有花费一次完整的主干前向传递来重新评估揭示结果,而是读取步骤间的残差,该残差已经编码了每个预测在新令牌被考虑后如何变化,并将修正后置信度回落到阈值以下的位置重新掩码。那些方法需要付出一次完整主干前向传递的代价进行重新检查,而 MRP 仅通过一次轻量级的残差传递就能获得。
我们学到了什么
在这个过程中,有几点特别突出:
- 深度在 2-3 层存在最佳点。 准确率从 1 层到 3 层稳步提升,然后趋于平稳,而吞吐量持续下降,因为额外的层只增加了每步延迟;到 8 层
相似文章
多令牌残差预测
引入多令牌残差预测(MRP),这是一个用于扩散语言模型的轻量级模块,能够在单次主干前向传播中实现依赖感知的多令牌去噪,实现高达1.42倍的无损加速。
基于时空并行解码与置信度外推的高效扩散LLMs
本文介绍了时空并行解码(TSPD)和置信度外推(CE),通过动态判断令牌何时收敛并预测logit趋势,来加速基于扩散的大语言模型的推理,减少不必要的去噪步骤,同时保持输出质量。
多块扩散语言模型
本文提出多块扩散语言模型(MBD-LMs),将单块扩散扩展为并发多块解码,并采用优化训练策略如多块教师强制(Multi-block Teacher Forcing)和优化的块缓冲区解码算法。实验表明,每次前向传递的令牌数增加,基准测试准确率提升。
$R^2$-dLLM:通过时空冗余削减加速扩散大语言模型
R²-dLLM 引入时空冗余削减技术,在保持生成质量的同时将扩散 LLM 的解码步数最多压缩 75%,直击部署瓶颈。
@simplifyinAI: 研究人员刚刚通过零精度损失将 LLM 的速度提升了 8.5 倍。这项技术被称为 DFlash。它取代了缓慢的自回归…
研究人员提出了 DFlash,这是一种用块扩散模型替代自回归草稿模型的方法,在零精度损失的情况下实现了 8.5 倍的 LLM 推理加速。