Ouzhang's picture
Add files using upload-large-folder tool
13c5606 verified
|
Raw
History Blame Contribute Delete
21.3 kB

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 的样本质量。

GPT-generated D-MMD overview

论文信息

项目 信息
标题 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
PDF 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 图

Paper 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

GPT-2 metric

论文认为 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

Posterior sampling settings - masked

Posterior sampling settings - uniform

解释:这组图研究 evaluation 时 posterior sampling 的 temperature 和 top-p 对 CIFAR-10 FID 的影响。它不是训练方法本身,而是说明采样超参数会显著影响最终 FID;不同 diffusion 类型的最佳 temperature/top-p 不完全相同,因此主实验需要调采样设置。

Appendix Figure: MMD teacher temperature

Teacher temperature - masked

Teacher temperature - uniform

解释:这组图固定 D-MMD 框架,改变蒸馏时 teacher 的 temperature。结论是 teacher 的 mode-seeking 强度会改变 student 的最终 FID;温度过高或过低都可能变差,说明 D-MMD 蒸馏的目标分布不是越尖越好,而是需要在质量和多样性之间取合适点。

Appendix Figure: MMD teacher top-p

Teacher top-p - masked

Teacher top-p - uniform

解释:这组图对应 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 崩溃。