块扩散把离散LM困惑度压到28.23,并支持近万token生成

Block Diffusion: Interpolating Between Autoregressive and Diffusion Language Models

Marianne Arriola, Aaron Gokaslan, Justin T. Chiu, Zhihan Yang, Zhixuan Qi, Jiaqi Han, Subham Sekhar Sahoo, Volodymyr Kuleshov

ICLR 2025 Oral

cs.LG, cs.AI

2025-03-13

提出BD3-LM:块间自回归、块内掩码扩散,补上KV缓存与变长生成。110M模型在LM1B上困惑度≤28.23,优于MDLM的31.78,最长样本接近一万token。

这篇在解决什么

离散扩散语言模型相对自回归,理论上能在块内并行出 token,也更容易往生成过程里塞约束。落地时却连着卡在三件事上。

现有离散扩散几乎都按训练时定好的长度出完整向量,聊天这种「答完就停」的场景接不住。去噪时上下文是双向的,前面算过的 KV 没法缓存,推理比自回归还笨。似然也一直落后:LM1B 上 MDLM 的测试困惑度还停在 ≤31.78,同样 110M 的自回归 Transformer 已经到 22.83。

Cornell Tech 的 Marianne Arriola、Volodymyr Kuleshov 等人在 ICLR 2025 Oral 给出一个插值:块与块之间自回归,块内部做掩码扩散。模型叫 BD3-LM(Block Discrete Denoising Diffusion Language Models)。块大小 L' 是旋钮,L'=1 时退回自回归,L' 等于全长时退回整段扩散。

方法

长度为 L 的序列切成 B 个长度为 L' 的块。对数似然按块因式分解,每个块的条件分布用离散去噪扩散来写,条件是前面已经生成好的干净块。

骨干是 Transformer,注意力换成块因果掩码:当前块内部双向(去噪需要看见块内其他位置),对历史块单向可见,看不见未来块。推理时历史块的 KV 缓存下来,当前块内部并行采样。这同时补上了变长和 KV 缓存,是相对整段离散扩散最直接的架构改动。

训练有个硬约束。去噪当前块必须喂噪声输入,下一块去噪又需要当前块的干净表示,朴素实现每个 token 至少过两遍网络。做法是把干净序列和噪声序列拼成一段,配专用注意力:噪声 token 只看本块其他噪声 token,以及前面块的干净 token。一次前向算完全部块的损失,比两次前向快 20%–25%,整体训练速度压在普通扩散的 2 倍以内。实验里还先按 L'=L 用标准扩散损失预训练 850K step,再在目标块大小上微调 150K step,把额外开销再削一截。

拉开数字差距的,是梯度方差。L'=1 时期望上扩散目标等于自回归负对数似然,LM1B 上训 16B token 仍差大约 2 点(≤25.56 vs 22.88)。掩码扩散平均只在一半 token 上算交叉熵,有效 batch 小了一半。把前向过程改成一律全掩码后,目标与自回归对齐,困惑度回到 22.88,NELBO 方差从 1.52 掉到 0.11。

L'>1 时用 clipped 噪声日程:掩码率从 U[0,1] 裁成 U[β,ω],躲开「几乎不掩」和「几乎全掩」。两端重建都太容易,梯度噪声大、学习信号弱。训练中每隔约 5K step 网格搜索 β、ω,用 NELBO 方差当代理。最优区间跟块大小绑在一起:L'=4 偏重掩码(大约 U[0.45,0.95] 或 U[0.5,1]),L'=16 更靠近中间。消融里 clipped 日程的测试困惑度和方差都低于 linear、log、cosine、square。

模型 12 层、隐维 768、12 头,约 110M,RoPE,不加时间步条件。LM1B 上下文 128,OWT 上下文 1024,batch 512。

结果

扩散数字都是 NELBO 上界,不能当成真实负对数似然跟自回归逐位咬。

LM1B 训 65B token:

方法测试 PPL↓
AR Transformer22.83
Transformer-XL Base23.5
D3PM (absorb)≤82.34
SEDD≤32.68
MDLM≤31.78
BD3-LM L'=16≤30.60
BD3-LM L'=8≤29.83
BD3-LM L'=4≤28.23

块越小越接近自回归。论文写相对 MDLM 最多约 13%;表上 L'=4 对 31.78 大约低 3.5 点。

OpenWebText 训 524B token:AR 17.54,SEDD ≤24.10,MDLM ≤22.98,BD3-LM L'=16/8/4 分别 ≤22.27 / ≤21.68 / ≤20.73。零样本(同样这套 OWT 模型)Pubmed 上 L'=4 到 42.52,低于 AR 的 48.59;Wikitext、LM1B、AG News 是扩散里最好。PTB 上 96.81,差于 MDLM 的 90.96 和 AR 的 81.07。Lambada、Arxiv 也没有赢过 MDLM。

变长生成抽 500 条,SEDD 被锁在训练上下文,最长 1024。BD3-LM L'=16 中位 798、最长 9982,大约 10 倍。停条件是采到 [EOS],或最近 256 token 平均熵掉到 4 以下,用来拦住崩掉之后的无限续写。AR 中位 4008、最长 131K,跟训练集上限一样。

样本质量用 GPT-2 Large 算 generative perplexity,300 条:

方法L=1024 Gen.PPL / NFEL=2048
AR14.1 / 1K13.2 / 2K
SEDD52.0 / 1K
MDLM46.8 / 1K41.3 / 2K
SSD-LM L'=2537.2 / 40K35.3 / 80K
BD3-LM L'=425.7 / 1K23.6 / 2K

SSD-LM 是连续嵌入上的高斯块扩散,400M 参数。函数评估次数压到与 BD3 同量级时,它的 Gen.PPL 崩到约 281。离散路线在少一个数量级的 NFE 下数字更好。

为什么重要

离散扩散要进实用解码,缺的就是 KV 缓存和变长。BD3-LM 把这两件补上,同时把似然从「明显落后」拉到能跟自回归放在一张表里看。对做扩散 LM 的人,这篇更像一份配方:块因果注意力、干净与噪声拼接训练、按块大小搜 clipped 日程。噪声日程那段对普通 MDLM / D3PM 也适用,作者自己点过这一点。

规模停在 110M,不是 7B 对话模型。块越小质量越好、块内并行度越低;块越大越像整段扩散、KV 缓存收益也越大。这是刻意留下的旋钮。跟自回归的似然差距还在,LM1B 上 28.23 vs 22.83。

局限与存疑

作者承认训练比普通扩散贵,向量化后仍可能接近 2 倍;块之间串行,小块会把扩散的并行和可控优势吐回去;最优块大小跟任务有关,需要更大块才谈得上块内编辑。NELBO 在 L'>1 时是松的,表里扩散数字全是上界。

实验没上到 LLM 尺度,也没有指令、代码、数学这些下游。零样本并非全胜。OWT 为了测变长,训练时不再往样本两端塞 [BOS]/[EOS],和部分公开基线的预处理不完全同构。Gen.PPL 用 GPT-2 Large,已知会偏向表面流畅,不能单独当质量判决。代码和权重已公开,复现门槛不算高,但结论目前只覆盖这个量级。

术语

原文与代码

社区讨论

相关论文

全部论文解读