Beyond Single Tokens: Distilling Discrete Diffusion Models via Discrete MMD
论文笔记:离散扩散模型通常需要很多去噪步,并且每步的 token 预测是 factorized 的,少步采样时容易累积独立性假设带来的错误。本文提出 Discrete Moment Matching Distillation, D-MMD,把连续扩散中的 Moment Matching Distillation 推广到离散 token / pixel 空间,用 teacher、student generator、auxiliary model 的交替优化来蒸馏少步离散扩散生成器。结果是在 CIFAR-10 和 OpenWebText 上,D-MMD 不仅显著减少 NFE,还经常超过原 teacher 的样本质量。
论文信息
| 项目 | 信息 |
|---|---|
| 标题 | Beyond Single Tokens: Distilling Discrete Diffusion Models via Discrete MMD |
| 作者 | Emiel Hoogeboom, David Ruhe, Jonathan Heek, Thomas Mensink, Tim Salimans |
| 单位 | Google DeepMind Amsterdam |
| arXiv | https://arxiv.org/abs/2603.20155 |
| https://arxiv.org/pdf/2603.20155 | |
| DOI | https://doi.org/10.48550/arXiv.2603.20155 |
| 版本 | arXiv v1, 2026-03-20 提交;PDF 元数据日期为 2026-03-23 |
| 领域 | cs.LG, cs.CV, stat.ML |
| 代码 | 论文 PDF、arXiv 页面和源码包中未给出官方代码链接;截至本笔记编写时未确认公开官方 repo |
| License | CC BY 4.0 |
这篇论文解决什么问题
离散扩散模型适合并行地生成一整段 token 或一整块离散像素,但推理通常要跑很多步。每一步模型输出的是各维度独立的 categorical 分布,短步数采样时,token 之间的相关性很难通过单次 factorized 预测直接表达,导致错误在采样迭代中累积。
连续扩散模型已有很多少步蒸馏方法,尤其是 Moment Matching Distillation, MMD;但离散空间里 hard categorical sample 对 student 参数不可导,直接照搬会困难。本文的核心做法是把 MMD 写成更一般的 min-max 形式,再替换成适合离散概率向量的交叉熵差值目标:
- student generator 输出软概率
x_hat_eta(z_t),不直接对 hard token 求梯度。 - frozen teacher
x_hat_theta(z_s)提供目标分布。 - auxiliary model
x_hat_phi(z_s)学 student 诱导出的分布,并作为 adversarial/moment-matching 参照。 - student 最小化 teacher loss、最大化 auxiliary loss;auxiliary 则学习 student,同时被 teacher regularize。
最终效果很直接:CIFAR-10 上 uniform teacher 1024 步 FID 7.5,而 Uniform D-MMD 32 步 FID 3.7;masked teacher 1024 步 FID 6.4,而 Masked D-MMD 64 步 FID 3.5。OWT 文本上,Masked D-MMD 16/32 步的 GPT-2 Gradient Error 也优于 256/512 步 teacher。
原论文 Overview 图
论文 Fig. 1 传达的是主结论:D-MMD 生成器在文本任务上用更少 function evaluations 就能匹配甚至超过 teacher。右侧的样本文本不是核心指标,而是展示 16-step Masked D-MMD 能生成一段较长的 1024-token 样本,说明少步生成不是只在短文本上成立。
方法拆解
1. 从连续 MMD 到通用 D-MMD
连续 MMD 的目标是让 student 采样分布下的条件一阶矩匹配真实/teacher 分布:
E_g[x | z_t] = E_q[x | z_t]
本文把原本的 alternating MMD 改写成一个更一般的 min-max 形式:
min_eta max_phi E_g [
L_s(x, teacher(z_s), z_s)
- L_s(x, auxiliary(z_s), z_s)
- L_s(teacher(z_s), auxiliary(z_s), z_s)
]
直观理解:
- generator 希望它产生的样本在 teacher 看来 loss 小。
- generator 同时希望 auxiliary 暂时跟不上它,从而形成 moment matching 的驱动力。
- auxiliary 要拟合 generator 诱导出的条件分布,同时靠 teacher regularization 避免跑偏。
在连续平方误差场景中,这个写法能还原原 MMD 的梯度;因此它是对 MMD 的推广,而不是完全另起炉灶。
2. 离散空间里匹配概率而不是 hard token
离散 token 的 categorical sample 不可导,所以本文用 student 的 soft probability vector 参与 generator loss:
L_GEN = CE(x_hat_eta | x_hat_theta) - CE(x_hat_eta | x_hat_phi)
这等价于用 teacher 和 auxiliary 的 log-probability 差值来更新 student 的概率输出。auxiliary loss 为:
L_AUX = CE(x | x_hat_phi) + CE(x_hat_theta | x_hat_phi)
其中 x 可以是 hard sample;在 masked diffusion 中也可以使用 soft target,因为 mask 后的 z_s 不泄漏具体 token。uniform diffusion 中需要用 hard sample 来避免 z_s 与 soft probability 不一致造成 bias。
3. factorized generator 如何学相关输出
表面上,每个位置独立采样的 categorical generator 似乎无法表达 token 相关性。论文的解释是:generator 实际包含两层随机性。
第一层是 x_hat_eta(z_t) 这个 soft sample 本身,它可以随输入噪声和当前状态产生相关变化;第二层才是逐维 categorical hard sample。为了匹配 teacher/auxiliary 的条件矩,模型会降低输出分布的 entropy,让 soft probability 在样本之间形成相关结构。这个现象在噪声条件实验中有实证支持。
4. temperature 与 top-p distillation
语言模型实际采样常常不用原始 logits,而是用 temperature 或 top-p 让分布更偏向高概率区域。D-MMD 也把这种 mode-seeking 行为蒸馏进 student:
- temperature distillation:teacher logits 除以温度
tau。 - top-p distillation:不能把 top-p 外的 logits 设成
-1e20,否则梯度会爆;论文改成对被 mask 的 logits 减去一个常数Delta=2,让低概率项变小但不产生极端梯度。
评价指标:GPT-2 Gradient Error
论文认为 generative perplexity 不适合评估离散扩散文本生成。原因是 reference LM 给高概率,不代表样本典型;低温采样或重复文本可能让 perplexity 变好,但样本质量变差。
因此作者提出 GPT-2 Gradient Error, GGE:把生成样本和真实数据分别喂给训练好的 GPT-2,看它们引起的参数梯度均值有多接近。若生成分布等于数据分布,reference model 在两者上的期望梯度应接近一致,GGE 越低越好。Fig. 2 说明:随着采样变得更 mode-seeking,perplexity 可能持续改善,但 gradient moment 会先改善后恶化,更能暴露过度偏向模式或分布失真的问题。
实验与图表解释
Table 1: CIFAR-10 主结果
| Model | 4 | 8 | 16 | 32 | 64 | 128 | 256 | 512 | 1024 |
|---|---|---|---|---|---|---|---|---|---|
| Uniform Teacher | - | - | 36.3 | 17.1 | 10.7 | 8.6 | 7.9 | 7.6 | 7.5 |
| Uniform D-MMD | 7.1 | 5.0 | 4.1 | 3.7 | 3.8 | - | - | - | - |
| Masked Teacher | - | - | 122.9 | 47.1 | 20.0 | 11.1 | 7.8 | 6.7 | 6.4 |
| Masked D-MMD | 22.3 | 12.7 | 5.3 | 3.8 | 3.5 | - | - | - | - |
解释:这是论文最强的图表之一。teacher 需要 1024 步才能到 FID 7.5/6.4,而 D-MMD 在 32 或 64 步就到 3.7/3.5。它说明 D-MMD 不只是加速近似 teacher,而是改变了采样分布的 Pareto front,少步 student 可以超过长步 teacher。作者也解释了这可能来自 adversarial/reverse-KL 风格的 mode-seeking,能修正最大似然 teacher 的 mode-covering 倾向。
Table 2: OWT 文本主结果
| Model | 8 | 16 | 32 | 64 | 128 | 256 | 512 |
|---|---|---|---|---|---|---|---|
| Uniform Teacher, p=0.50 | - | - | 0.375 | 0.326 | 0.330 | 0.324 | 0.313 |
| Uniform D-MMD, p=0.70/0.70 | 0.337 | 0.310 | 0.307 | 0.316 | - | - | - |
| Masked Teacher, p=0.85 | - | - | 0.402 | 0.307 | 0.297 | 0.275 | 0.275 |
| Masked D-MMD, p=0.85 | 0.456 | 0.236 | 0.225 | 0.231 | - | - | - |
| AR Baseline | 0.061 | - | - | - | - | - | - |
解释:指标是 GGE,越低越好。Masked D-MMD 在 16 步达到 0.236,已经优于 256/512 步 teacher 的 0.275;32 步进一步到 0.225。Uniform D-MMD 也能在 16/32 步略优于 512 步 teacher。AR baseline 仍明显更强,说明这篇论文主要证明离散扩散蒸馏有效,而不是声称小规模扩散文本模型已经超过 AR LM。
Table 3: Block Autoregressive Diffusion
| Model | 16 | 256 |
|---|---|---|
| 256-Block Uniform Teacher, p=0.9 | - | 0.225 |
| 256-Block Uniform D-MMD, p=0.7 | 0.225 | - |
解释:整段 1024 token 无条件生成比较理想化;更现实的方式是 AR encoder 负责上下文,diffusion 并行生成一个 block。这里 block size 为 256,16-step D-MMD 达到和 256-step teacher 相同的 GGE 0.225,说明 D-MMD 也适合 block-level 并行生成场景。
Table 4: CIFAR-10 与相关工作对比
| Method | NFE | FID |
|---|---|---|
| Di4C Teacher | 40 | 8.0 |
| Di4C hybrid | 20 | 9.5 |
| Di4C | 10 | 20.6 |
| Uniform Teacher | 512 / 64 | 7.6 / 10.7 |
| Uniform D-MMD | 8 / 16 / 32 | 5.0 / 4.1 / 3.7 |
| Masked Teacher | 512 / 64 | 6.7 / 20.0 |
| Masked D-MMD | 16 / 32 / 64 | 5.3 / 3.8 / 3.5 |
解释:Di4C 的 teacher 只需 40 步且本身对 CIFAR-10 设置有利,但 D-MMD 8 步 Uniform 已有 FID 5.0,优于 Di4C teacher 的 8.0;64 步 Masked D-MMD 达到 3.5。这个表强调的是与已有离散扩散蒸馏工作的横向差距。
Table 5: OWT 与相关工作对比
| Method | NFE | GGE ↓ | GPT-2 PPL ↓ | Sample entropy |
|---|---|---|---|---|
| Duo + DCD | 4 | - | 108.2 | 4.82 |
| Duo + Di4C | 4 | - | 150.7 | 4.81 |
| MDLM + SDTT | 4 | - | 339.7 | 5.38 |
| MDLM + Di4C | 4 | - | 239.3 | 5.40 |
| FMLM | 4 | - | 76.4 | 5.05 |
| Masked Teacher | 256 | 0.275 | 22.5 | 5.13 |
| SDTT reimpl. | 64 / 32 | 0.293 / 0.340 | 26.9 / 30.4 | 5.17 / 5.18 |
| Masked D-MMD | 4 / 16 / 32 | 0.820 / 0.236 / 0.225 | 20.3 / 17.2 / 19.4 | 4.60 / 5.00 / 5.05 |
| Data | - | 0.000 | 15.4 | 5.44 |
解释:这个表同时报告 GGE、generative perplexity 和 entropy。D-MMD 16/32 步在 GGE 上优于 SDTT 和 teacher;但 4 步 D-MMD 虽然 GPT-2 PPL 低,GGE 和 entropy 都显示分布质量不够好。这正好支持作者关于 perplexity 不可靠的论点:PPL 好不等于真实分布匹配好。
Table 6: Masked D-MMD 是否需要输入噪声
| Masked D-MMD | 指标 | 4 | 8 | 16 | 32 | 64 |
|---|---|---|---|---|---|---|
| without noise | FID | 151 | 37.0 | 14.7 | 7.7 | 6.0 |
| without noise | generator output entropy | 1.26 | 1.37 | 1.57 | 1.86 | 1.91 |
| with noise | FID | 22.3 | 12.7 | 5.3 | 3.8 | 3.5 |
| with noise | generator output entropy | 1.01 | 1.29 | 1.53 | 1.76 | 1.83 |
解释:masked distillation 对额外噪声源很敏感。加入噪声后,低步数 FID 大幅改善,例如 4 步从 151 降到 22.3,16 步从 14.7 降到 5.3。作者的解释是:为了表达相关输出,generator 需要能在 soft probabilities 层面产生随机相关结构;输入噪声给 masked generator 提供了这一自由度。
Appendix Figure: evaluation-time posterior sampling settings
解释:这组图研究 evaluation 时 posterior sampling 的 temperature 和 top-p 对 CIFAR-10 FID 的影响。它不是训练方法本身,而是说明采样超参数会显著影响最终 FID;不同 diffusion 类型的最佳 temperature/top-p 不完全相同,因此主实验需要调采样设置。
Appendix Figure: MMD teacher temperature
解释:这组图固定 D-MMD 框架,改变蒸馏时 teacher 的 temperature。结论是 teacher 的 mode-seeking 强度会改变 student 的最终 FID;温度过高或过低都可能变差,说明 D-MMD 蒸馏的目标分布不是越尖越好,而是需要在质量和多样性之间取合适点。
Appendix Figure: MMD teacher top-p
解释:这组图对应 top-p distillation。它补充了方法部分的设计选择:top-p 可以把 teacher 目标推向更高质量模式,但如果处理不当会梯度爆炸。论文采用降低非 top-p logits 的温和实现,并用这些曲线说明 top-p 值同样需要调参。
Sample Figure: 16-step Masked D-MMD 生成样本
论文还给出了一段随机、非 cherry-pick 的 1024-token 样本。这个图的作用不是定量比较,而是定性展示:少步 Masked D-MMD 可以生成较长文本,并具有局部连贯性;但从样本文字看,长程语义和事实一致性仍不是本文重点,也没有达到强 AR LM 的文本质量。
整体评价
这篇论文的价值在于把连续扩散里有效的 MMD 蒸馏思想比较干净地迁移到了离散扩散,并且没有只停留在公式层面:CIFAR-10、OWT、block-AR、相关工作对比和噪声消融都指向同一个结论,即 D-MMD 能明显改善少步离散扩散的质量。
我认为最有启发的点有三个。第一,student 超过 teacher 并非偶然,而是 adversarial/moment-matching 目标带来的分布偏移,类似从 mode-covering 往更适合采样的方向移动。第二,GGE 指标很好地指出了 generative perplexity 的缺陷,对文本扩散评估很重要。第三,factorized categorical 输出并不等于模型完全无法表达相关性,关键在于 soft probability 与输入噪声形成的随机结构。
局限也比较明显:实验规模仍偏研究验证,文本结果和 AR baseline 还有明显差距;D-MMD 的 adversarial 训练会依赖 teacher temperature、top-p、noise conditioning 等超参数;官方代码未公开也会影响复现。总体来看,这是一篇很适合继续跟进的离散扩散蒸馏工作,尤其适合关注 block-level 并行文本生成和少步离散生成的人阅读。
用一个样本文本串讲 D-MMD
论文里的随机样本大概是这样一段体育/采访风格文本:
"He's in a really good spot. It's the right situation, you don't have to put him in,
you shouldn't be able to get him, and I think that is what I loved to do when he was
young..."
"He's shown this year, with his growth in the system, he's done a really, really,
great job offensively..."
"I think he is definitely on the right path, I think he is playing on a high level..."
这段文本不是 prompt continuation,而是无条件生成的 1024-token sample。可以把它理解成:模型先面对一整段长度为 1024 的空白/噪声 token 序列,然后用 16 次并行 denoising,把整段文本逐步变得像 OpenWebText 里的新闻采访。
1. 16-step Masked D-MMD 生成时发生了什么
以 masked diffusion 为例,初始状态可以想成:
[MASK] [MASK] [MASK] ... [MASK]
在第 1 步,student generator 对每个位置输出一个 token 分布,而不是直接确定一个 token。比如某个位置可能给出:
P("He")=0.18, P("The")=0.12, P("I")=0.08, ...
另一个位置可能给出:
P("'s")=0.21, P("was")=0.10, P("is")=0.09, ...
然后从这些分布中采样,得到一些 hard token,再通过 posterior sampler 回到较早的 noisy state z_s。重复 16 次后,局部片段逐渐稳定成:
"He's shown this year, with his growth in the system..."
关键不是每一步只改一个 token,而是每一步对整段 1024 token 并行给概率。这样理论上速度快,但难点是:每个位置的 categorical 分布表面上是 factorized 的,容易把 token 间关系拆散。D-MMD 的目标就是让 student 虽然最后按位置采样,但它输出的 soft probabilities 本身要带有相关结构。
2. D-MMD 的两个 loss 在这个例子里是什么样子
假设当前 student 生成了一个半成品片段:
"He's shown this year, with his growth in the system..."
对某个位置,比如 shown 后面的位置,teacher 看到上下文和 noisy state z_s 后可能认为:
teacher: P("this")=0.35, P("the")=0.12, P("last")=0.08, ...
auxiliary model 是一个正在追踪 student 生成分布的模型,它可能认为:
auxiliary: P("the")=0.30, P("this")=0.18, P("last")=0.10, ...
student 自己在更早的 z_t 上输出:
student: P("this")=0.25, P("the")=0.20, P("last")=0.06, ...
此时 generator loss 是:
L_GEN = CE(student | teacher) - CE(student | auxiliary)
直观说:
CE(student | teacher)让 student 往 teacher 认为自然的 token 分布靠近,比如提高"this"。- CE(student | auxiliary)让 student 避开 auxiliary 已经能预测到的坏/旧分布,形成 adversarial 的 moment matching 推力。- 最终梯度大致看的是
log teacher - log auxiliary:teacher 比 auxiliary 更偏好的 token 会被 student 增强,反过来会被压低。
auxiliary loss 是:
L_AUX = CE(sampled token | auxiliary) + CE(teacher | auxiliary)
它的含义是:
- 第一项让 auxiliary 学会 student 实际采出来的 token,例如这次采到了
"this",auxiliary 就要提高"this"。 - 第二项让 auxiliary 不要离 teacher 太远,避免 auxiliary 只追着 student 的噪声跑。
交替训练后,如果 student 生成的分布真的像 teacher 的采样分布,auxiliary 会追上 teacher,二者输出接近,此时 log teacher - log auxiliary 接近 0,D-MMD 达到固定点。
3. 为什么这能生成相关文本
看这段样本,很多 token 之间有明显相关性:
"He's ... his growth ... he's done ... offensively"
"right path ... playing on a high level ... made great strides"
如果模型只独立预测每个位置,很容易出现类似:
"He's shown this year, with the growth on they system, it's done offensively path..."
D-MMD 希望 student 的 soft probabilities 在采样前就带有整体倾向。比如当前面高概率选择 "He's" 和 "shown" 时,后面位置的 soft distribution 也要更偏向 "this year"、"his growth"、"in the system" 这类同一语域下的组合。论文说 factorized generator 能学相关性,靠的就是 soft sample 层面的随机结构和低输出 entropy,而不是最后 categorical step 本身。
4. 评测这段文本时具体看什么
论文主要不用普通 likelihood,因为 D-MMD 这种生成器没有 tractable sampling likelihood。它报告三个文本相关量:
Generative perplexity:把这段生成文本喂给 GPT-2 Large,看 GPT-2 觉得它有多像自然文本。低 PPL 通常说明局部语法和常见短语不错。比如这段里 "He's shown this year"、"right path"、"great strides" 都是 GPT-2 熟悉的新闻/体育采访表达,所以 PPL 可能较低。
Sample entropy:看 token 分布是否太单调。若模型一直重复 "he's on the right path",PPL 可能不差,但 entropy 会下降。论文希望 entropy 不要太低。
**GPT-2 Gradient Error (GGE)**:这是更核心的指标。它不只问 GPT-2 给这段文本多高概率,而是问:如果 GPT-2 用这批生成文本训练一步,它的参数梯度方向和用真实 OWT 文本训练一步的梯度方向差多少。若生成文本整体分布像真实数据,生成样本的梯度均值应接近真实数据的梯度均值,GGE 低。
因此,这段样本如果只是局部短语自然,但整体长期重复、话题单一、结构不像真实 OWT,PPL 可能还行,GGE 会变差。
5. 一个能骗过 PPL 的反例
论文批评 generative perplexity,是因为它容易奖励高概率但不典型的文本。比如:
"He is on the right path. He is on the right path. He is on the right path.
He is on the right path. He is on the right path..."
这类文本每个局部 token 都很常见,GPT-2 可能给出不高的 perplexity。但它显然不是高质量 1024-token 新闻样本:重复、信息量低、多样性差。entropy 会偏低,GGE 也会发现它不像真实数据分布。
另一种反例是语法碎片式乱码:
"he growth system right offensively path young court get year done body ready..."
这种可能 entropy 不低,但 GPT-2 PPL 会很高,GGE 也会很差。论文里说的两个失败方向就是:一种是低熵重复骗 PPL,另一种是高熵但不成文的随机 token。
6. 回到这段样本,它说明了什么
这段 16-step Masked D-MMD 样本说明:少步 D-MMD 已经能生成局部连贯、语域一致的长文本,像体育采访或新闻引语。它不是完美文本:人物名、事实、长程语义可能不可靠,也有一些重复表达。但它不像 4-step 失败模型那样全是随机 token,也不像 mode collapse 那样只重复一句话。
所以这段样本和表格结果合起来表达的是:D-MMD 的价值不在于一步到达强 AR LM 水平,而在于让离散扩散模型在少得多的 denoising steps 下仍能保持较好的分布匹配,尤其避免 few-step discrete diffusion 常见的 factorization 崩溃。








