数据不够算力够时,掩码扩散语言模型胜过自回归

Diffusion Beats Autoregressive in Data-Constrained Settings

Mihir Prabhudesai, Mengning Wu, Amir Zadeh, Katerina Fragkiadaki, Deepak Pathak

cs.LG, cs.AI, cs.CV, cs.RO

2025-07-22

CMU在重复有限数据上对照训练约200个模型:扩散数据复用半衰期约494个epoch,自回归约31个;100M独特token上验证损失3.55对3.71。

这篇在解决什么

高质量文本有上限,算力还在涨。机器人、医疗这些领域一开始就没有海量独特数据。自回归语言模型一直按从左到右的下一词训练,单遍数据上很划算。掩码扩散语言模型把生成做成随机遮挡再补全,双向都能看,先验工作却说它要大约 16 倍算力才能追上同样的验证 NLL。

那 16 倍是在单 epoch、每个 token 只见一次的设定里量的。算力放大时模型和独特数据一起加,分不清扩散吃亏的是算力效率还是样本效率。CMU 把独特数据量钉死,反复读同一批数据,问的是:数据才是瓶颈时,谁更值。

方法

两家共用 GPT-2 风格加 RoPE 的 Transformer,C4 英文、GPT-2 BPE、序列长 2048。独特 token 预算 25M / 50M / 100M,最多训 800 个 epoch,模型从 7M 到 2.5B,一共约 200 个模型。超参跟 Muennighoff 等人给自回归调的那套,对 AR 略有利。

自回归用因果掩码做下一词。扩散每步抽一个掩码比例 r,按 r 独立把 token 换成 [MASK],在双向注意下还原被遮位置,损失按 1/r 加权,对应似然的 ELBO。掩码图案每次重采样,模型等于在大量不同的条件预测顺序上训练。

缩放沿用 data-constrained Chinchilla:重复数据的效用按指数衰减,拟合复用半衰期 RD。超过这个 epoch 数,再读同一批数据的收益会掉得很厉害。

结果

单 epoch、Chinchilla 最优点附近,扩散明显更差:100M unique token 上验证损失 10.65 对 7.07。把同一批数据反复读下去,AR 大约 50 个 epoch 就开始过拟合,扩散在实验预算内没有过拟合,500 个 epoch 时损失 3.55,低于 AR 最好的 3.71。相对各自单 epoch 起点,扩散降了 67%,AR 降了 48%。

拟合出的 RD,扩散约 494,AR 约 31。重复数据几乎等价于新数据的区间,AR 大约 4 个 epoch,扩散大约 100 个。扩散开始超过 AR 的临界算力随独特 token 数呈幂律,指数约 2.174。

下游也跟着走。100M unique token 上,按验证损失选出的最好扩散模型:

任务AR扩散
SciQ58.0568.67
LAMBADA10.9115.19
ARC-Easy35.6337.84
HellaSwag27.3730.24
PiQA60.9460.72

500M unique token、按临界算力训的 2.3B 扩散模型(130 epoch,仍未收敛)在 SciQ 上是 79.13 对 AR 的 67.82,LAMBADA 是 22.30 对 15.07。PiQA 上 AR 仍略高。

机制对照:给 AR 加 attention dropout 或把输入 token 的注意力置零,过拟合照旧。改成在 N 种固定排列上做下一词,N=16 时 100 epoch 的验证损失逼近扩散。随机掩码带来的多种条件顺序,是样本效率的主要来源。

为什么重要

先前「扩散语言模型要 16 倍算力」把样本效率和算力效率绑在一起了。数据会先于算力耗尽,或者根本没有互联网级语料时,扩散把同一批数据重复利用的能力强一个数量级。作者给的口诀很硬:缺算力用自回归,缺数据用扩散。混合模型如果能在顺序多样性和每步监督密度之间插值,是这篇自己标出的下一步。

局限与存疑

验证损失的绝对值不能在两家之间直接比,缩放律里的熵常数 E0 不同,作者在机制节里把这项拿掉再比。下游表在 100M 设定下模型仍然很小,不少任务刚过随机基线,PiQA 上扩散还略输。超参原本为 AR 调,扩散可能还能更好,方向却不能反过来帮 AR 开脱。500M 那次扩散因算力中止,没有看到过拟合边界。实验是语言建模,对机器人或医疗序列的外推没有直接证据。推理侧扩散多步去噪的成本,这篇几乎没谈。

术语

原文与代码

社区讨论

相关论文

全部论文解读