分布匹配蒸馏

分布匹配蒸馏

一种将多步扩散模型蒸馏为单步 / 少步生成器的方法系列。


1. 背景与动机

扩散模型(Diffusion Models)能生成高质量图像,但采样需要数十到上百次网络前向计算(迭代去噪),速度慢、成本高。已有的蒸馏方法(Progressive Distillation、Consistency Models、Rectified Flow 等)试图让学生网络逐点拟合教师模型"噪声 → 图像"的确定性映射,这一映射本身极其复杂,学生难以完美模仿,导致蒸馏后质量明显下降。

DMD 系列换了一个思路:不追求复现教师的具体采样轨迹,只要求学生的输出分布在整体上与教师一致(distribution matching)。这类似 GAN 的思想,但用扩散模型来充当"分布的评分器",从而兼得扩散模型训练的稳定性与单步生成的高速度。


2. DMD:单步分布匹配蒸馏

2.1 核心思想

DMD 方法框架图

给定一个预训练的多步扩散教师模型 μbase\mu_{\text{base}},训练一个单步生成器 GθG_\theta(结构与教师相同、去掉时间条件,用教师权重初始化),使得 GθG_\theta 的输出分布 pfakep_{\text{fake}} 与真实数据分布 prealp_{\text{real}} 匹配。目标是最小化 KL 散度:

DKL(pfake∥preal)=Ez∼N(0;I), x=Gθ(z)[−(log⁡preal(x)−log⁡pfake(x))]D_{\text{KL}}(p_{\text{fake}} \| p_{\text{real}}) = \mathbb{E}_{z \sim \mathcal{N}(0;I),\, x = G_\theta(z)} \left[ -\left( \log p_{\text{real}}(x) - \log p_{\text{fake}}(x) \right) \right]

概率密度不可计算,但训练只需要梯度。对 θ\theta 求梯度后:

∇θDKL=E[−(sreal(x)−sfake(x))dGdθ]\nabla_\theta D_{\text{KL}} = \mathbb{E}\left[ -\left( s_{\text{real}}(x) - s_{\text{fake}}(x) \right) \frac{dG}{d\theta} \right]

其中 sreal(x)=∇xlog⁡preal(x)s_{\text{real}}(x) = \nabla_x \log p_{\text{real}}(x)、sfake(x)=∇xlog⁡pfake(x)s_{\text{fake}}(x) = \nabla_x \log p_{\text{fake}}(x) 是两个分布的 score 函数。直觉上:

  • sreals_{\text{real}} 把生成样本拉向真实分布的众数(更"真实");
  • −sfake-s_{\text{fake}} 把样本推离当前的生成分布(防止坍缩、更"不假");
  • 两者之差就是"更真实、更不假"的更新方向。

2.2 用扩散模型估计两个 score

直接计算 score 有两个障碍:低概率区域 score 发散;且扩散模型只能给出加噪后分布的 score。Score-SDE 理论给出了解法——对样本注入不同强度的随机高斯噪声 xt∼q(xt∣x)∼N(αtx;σt2I)x_t \sim q(x_t|x) \sim \mathcal{N}(\alpha_t x; \sigma_t^2 I),使两个分布"模糊化"后处处重叠,梯度良定义,而扩散模型的去噪输出恰好近似加噪分布的 score:

sreal(xt,t)=−xt−αtμreal(xt,t)σt2,sfake(xt,t)=−xt−αtμfakeϕ(xt,t)σt2s_{\text{real}}(x_t, t) = -\frac{x_t - \alpha_t \mu_{\text{real}}(x_t, t)}{\sigma_t^2}, \qquad s_{\text{fake}}(x_t, t) = -\frac{x_t - \alpha_t \mu^{\phi}_{\text{fake}}(x_t, t)}{\sigma_t^2}

  • Real score:真实分布固定,直接用冻结的预训练教师 μreal=μbase\mu_{\text{real}} = \mu_{\text{base}} 建模;
  • Fake score:生成分布随训练不断变化,因此用一个动态训练的 fake 扩散模型 μfakeϕ\mu^\phi_{\text{fake}}(同样从教师初始化),在生成器输出的 fake 样本上用标准去噪损失持续训练:

Ldenoiseϕ=∥μfakeϕ(xt,t)−x0∥22\mathcal{L}^{\phi}_{\text{denoise}} = \left\| \mu^\phi_{\text{fake}}(x_t, t) - x_0 \right\|_2^2

最终的分布匹配梯度(对时间步 t∼U(Tmin⁡,Tmax⁡)t \sim \mathcal{U}(T_{\min}, T_{\max}) 取期望,Tmin⁡=0.02TT_{\min}=0.02T, Tmax⁡=0.98TT_{\max}=0.98T):

∇θDKL≃Ez,t,x,xt[wtαt(sfake(xt,t)−sreal(xt,t))dGdθ]\nabla_\theta D_{\text{KL}} \simeq \mathbb{E}_{z,t,x,x_t} \left[ w_t \alpha_t \left( s_{\text{fake}}(x_t, t) - s_{\text{real}}(x_t, t) \right) \frac{dG}{d\theta} \right]

权重 wt=σt2αt⋅CS∥μbase(xt,t)−x∥1w_t = \frac{\sigma_t^2}{\alpha_t} \cdot \frac{CS}{\|\mu_{\text{base}}(x_t,t) - x\|_1} 用于归一化不同噪声水平下的梯度幅度,稳定训练(比 DreamFusion/ProlificDreamer 的 σt/αt\sigma_t/\alpha_t、σt3/αt\sigma_t^3/\alpha_t 加权方案好约 0.9 FID)。

2.3 回归损失(Regression Loss)

纯分布匹配目标在小噪声水平下不可靠,且 score 对概率密度的缩放不敏感,容易模式坍缩(mode collapse)。因此 DMD 额外引入回归损失:

  • 离线构造配对数据集 D={z,y}\mathcal{D} = \{z, y\}:用教师模型 + 确定性 ODE 求解器(Heun / PNDM,18~256 步)预先生成"噪声-图像"对(成本不到总训练的 1%);
  • 强制单步生成器在同一噪声输入下匹配教师输出的大尺度结构:

Lreg=E(z,y)∼D ℓ(Gθ(z),y),ℓ=LPIPS\mathcal{L}_{\text{reg}} = \mathbb{E}_{(z, y) \sim \mathcal{D}} \, \ell(G_\theta(z), y), \quad \ell = \text{LPIPS}

总目标:生成器优化 DKL+λregLregD_{\text{KL}} + \lambda_{\text{reg}} \mathcal{L}_{\text{reg}}(λreg=0.25\lambda_{\text{reg}} = 0.25),fake 评分网络优化 Ldenoiseϕ\mathcal{L}^{\phi}_{\text{denoise}},两者交替更新。两个损失作用在不同数据流上:KL 用随机噪声生成的 unpaired 样本,回归用配对数据集。

2.4 DMD 训练流程伪代码

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
算法 1:DMD 训练流程
──────────────────────────────────────────────────────
输入:预训练教师扩散模型 μ_real(冻结)
离线配对数据集 D = {(z_ref, y_ref)}(教师多步采样生成)
输出:训练好的单步生成器 G

初始化:
G ← copyWeights(μ_real) # 单步生成器,去掉时间条件
μ_fake ← copyWeights(μ_real) # fake 分布 score 估计网络(可训练)

while 训练未收敛 do

# ---------- 采样 ----------
采样随机噪声批 z ~ N(0, I)
采样配对数据批 (z_ref, y_ref) ~ D
x ← G(z) # 单步生成 fake 图像
x_ref ← G(z_ref)

# ---------- 1) 更新生成器 G ----------
# (a) 分布匹配梯度(对应论文 Eq.7)
t ~ U(T_min, T_max) # 随机时间步
ε ~ N(0, I)
x_t ← α_t · x + σ_t · ε # 对 fake 样本加噪(前向扩散)
with no_grad:
μ_real_pred ← μ_real(x_t, t) # real score:冻结教师
μ_fake_pred ← μ_fake(x_t, t) # fake score:动态网络
w_t ← σ_t²/α_t · 1/mean|μ_real_pred − x| # 梯度归一化权重
grad ← w_t · (μ_fake_pred − μ_real_pred) # score 之差
L_KL ← 0.5 · MSE(x, stopgrad(x − grad)) # 等价地注入该梯度

# (b) 回归损失(对应论文 Eq.9)
L_reg ← LPIPS(x_ref, y_ref)

L_G ← L_KL + λ_reg · L_reg # λ_reg = 0.25
G ← optimizer_G.update(∇L_G)

# ---------- 2) 更新 fake score 网络 μ_fake ----------
t' ~ U(0, 1)
ε' ~ N(0, I)
x_t' ← forwardDiffusion(stopgrad(x), ε', t') # 对最新 fake 样本加噪
L_denoise ← weight(t') · MSE(μ_fake(x_t', t'), stopgrad(x)) # Eq.6
μ_fake ← optimizer_fake.update(∇L_denoise)

end while

2.5 结果

任务 指标 备注
ImageNet 64×64 FID 2.62(1 步) 教师 EDM 需 511 步,FID 2.32
CIFAR-10 FID 2.66(1 步) 超越 Consistency Model(6.20)约 2.4×
zero-shot COCO-30k(SD v1.5,CFG=3) FID 11.49(1 步,90ms) 教师 50 步 2590ms,FID 8.78
推理速度 20 FPS(FP16,512×512) 相比教师加速约 30~100×

3. DMD2:改进的分布匹配蒸馏

3.1 DMD 的痛点与 DMD2 的改进总览

DMD2 方法框架图

DMD 的回归损失虽然保证了稳定性,但带来两个问题:

  1. 成本高昂:需要用教师多步采样构造数百万噪声-图像对。以 SDXL 为例,覆盖 LAION 1200 万条 prompt 约需 700 A100 天,是 DMD2 全部训练开销的 4 倍以上;
  2. 性能天花板:回归损失把学生绑死在教师的采样轨迹上,学生质量无法超越教师。

DMD2 通过三项改进解决这些问题,最终让学生超越教师:

# 改进 作用
① 去除回归损失 + 双时间尺度更新规则(TTUR) 消除离线建集成本,解开学生与教师轨迹的绑定;用 5:1 的更新频率保证 μ_fake 准确跟踪生成分布,维持训练稳定
② 引入 GAN 损失(判别真实图像 vs 生成图像) 学生首次接触真实数据,纠正教师 real score 的近似误差,使学生质量可超越教师
③ 多步生成器 + 反向模拟(Backward Simulation) 支持 4 步采样提升质量上限;训练时模拟推理时的中间输入分布,消除训练-推理失配

3.2 改进①:去除回归损失 + TTUR

直接去掉回归损失会导致训练不稳定(生成图像的平均亮度等统计量剧烈震荡、不收敛)。DMD2 将其归因于:fake 扩散批评器 μfake\mu_{\text{fake}} 在非平稳的生成分布上动态训练,score 估计不准,导致生成器梯度有偏。

借鉴 GAN 训练中的双时间尺度更新规则(Two Time-scale Update Rule, TTUR):每更新 1 次生成器,更新 5 次 μfake\mu_{\text{fake}},确保 fake score 估计始终准确跟踪当前生成分布。实验表明这样即使不用回归损失也能达到与 DMD 相当的稳定性和质量,且收敛更快。

3.3 改进②:GAN 损失——用真实数据超越教师

DMD 中学生从未见过真实数据,教师 real score 的近似误差会传播给学生且无法纠正。DMD2 在 pipeline 中加入 GAN 目标:

  • 判别器设计极简:在 fake 扩散模型 μfake\mu_{\text{fake}} 的 UNet bottleneck 上加一个分类分支,复用其编码特征;
  • 对真实图像和生成图像都先注入噪声 F(⋅,t)F(\cdot, t)(平滑分布、稳定训练),再进行真假分类,采用标准 non-saturating GAN 目标:

LGAN=Ex∼preal, t[log⁡D(F(x,t))]+Ez∼pnoise, t[−log⁡D(F(Gθ(z),t))]\mathcal{L}_{\text{GAN}} = \mathbb{E}_{x \sim p_{\text{real}},\, t}\left[\log D(F(x, t))\right] + \mathbb{E}_{z \sim p_{\text{noise}},\, t}\left[-\log D(F(G_\theta(z), t))\right]

判别器在真实数据上训练,不受教师误差限制,因此学生质量得以超越教师。GAN 目标也是分布级的匹配,与 DMD 的哲学一致(无需配对数据、不依赖教师轨迹)。消融显示:单独去掉 GAN 损失,SDXL 蒸馏结果 FID 从 19.32 退化到 26.90,且图像过饱和、过平滑。

3.4 改进③:多步生成器与反向模拟

SDXL 这类大模型单步蒸馏困难(模型容量有限、噪声→高多样性图像的直接映射难学),因此 DMD2 支持多步采样:

  • 推理:固定 NN 个时间步 {t1,…,tN}\{t_1, \dots, t_N\}(4 步模型用 {999,749,499,249}\{999, 749, 499, 249\}),从 z0∼N(0,I)z_0 \sim \mathcal{N}(0, I) 出发,交替执行去噪 x^ti=Gθ(xti,ti)\hat{x}_{t_i} = G_\theta(x_{t_i}, t_i) 与加噪 xti+1=αti+1x^ti+σti+1ϵx_{t_{i+1}} = \alpha_{t_{i+1}}\hat{x}_{t_i} + \sigma_{t_{i+1}}\epsilon(借鉴 Consistency Model);
  • 训练-推理失配问题:以往多步蒸馏方法训练时输入的是"加噪的真实图像",但推理时(除第一步外)输入的是"上一步生成器输出的加噪结果",两者分布不同,损害质量;
  • 反向模拟(Backward Simulation):训练时先用当前学生生成器跑几步、模拟推理过程得到中间样本,再对其去噪并用上述损失监督。由于学生只需跑很少几步,模拟代价可接受;且分布匹配损失不依赖学生的输入,天然不受模拟误差影响(对比同期 Imagine Flash 的回归损失会沿采样路径累积误差)。

3.5 DMD2 训练流程伪代码

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
算法 2:DMD2 训练流程
──────────────────────────────────────────────────────
输入:预训练教师扩散模型 μ_real(冻结)
真实图像数据集 D_real(训练 GAN 判别器,无需与 prompt 配对)
多步时间步序列 {t1, ..., tN}(单步时 N=1)
输出:训练好的少步生成器 G

初始化:
G ← copyWeights(μ_real) # 多步生成器(保留时间条件)
μ_fake ← copyWeights(μ_real) # fake score 网络(可训练)
D ← 在 μ_fake 的 UNet bottleneck 上添加分类分支 # GAN 判别器

while 训练未收敛 do

# ---------- 采样:反向模拟推理过程 ----------
采样 z_0 ~ N(0, I),随机选择起始时间步 t_i
x_in ← 加噪后的 z_0(t_i = t1 时)
或用 G 从更早时间步先跑若干步、模拟推理得到 x_{t_i} # Backward Simulation
x ← G(x_in, t_i) # 生成(去噪)输出

# ---------- 1) 更新生成器 G ----------
# (a) 分布匹配梯度(同 DMD,Eq.2)
t ~ U(T_min, T_max),ε ~ N(0, I)
x_t ← α_t · x + σ_t · ε
with no_grad:
s_real ∝ μ_real(x_t, t) # 冻结教师
s_fake ∝ μ_fake(x_t, t)
grad ← w_t · (s_fake − s_real)
L_KL ← 0.5 · MSE(x, stopgrad(x − grad))

# (b) GAN 生成器损失(对生成图像加噪后判真)
t_g ~ U(0, T)
L_GAN_G ← −log D( F(x, t_g) )

L_G ← L_KL + λ_GAN · L_GAN_G # 不再需要回归损失!
G ← optimizer_G.update(∇L_G)

# ---------- 2) 更新 μ_fake 与判别器 D(每 1 次 G 更新对应 5 次,TTUR) ----------
for k = 1 to 5 do
# (a) fake score 去噪损失
t' ~ U(0, 1),ε' ~ N(0, I)
x_t' ← forwardDiffusion(stopgrad(x), ε', t')
L_denoise ← weight(t') · MSE(μ_fake(x_t', t'), stopgrad(x))

# (b) GAN 判别器损失(真实图像 vs 生成图像)
x_real ~ D_real,t_d ~ U(0, T)
L_GAN_D ← −log D(F(x_real, t_d)) − log(1 − D(F(stopgrad(x), t_d)))

L_critic ← L_denoise + L_GAN_D
(μ_fake, D) ← optimizer_critic.update(∇L_critic)
end for

end while

3.6 结果

任务 指标 对比
ImageNet 64×64(1 步) FID 1.51,加长训练后 1.28 超越教师(EDM ODE 采样 511 步,FID 2.32)
zero-shot COCO(SD v1.5,1 步) FID 8.35 比 DMD(11.49)提升 3.14;超越 50 步教师
SDXL 1024×1024(4 步) FID 19.32,Patch FID 20.86 与 100 步教师(FID 19.36)相当,人评 24% 样本质量优于教师
SDXL(1 步) FID 19.01 大幅优于 SDXL-Turbo(24.57)、LCM-SDXL(81.62)

消融实验(ImageNet)验证了各组件的贡献:

配置 FID
DMD 原版(含回归损失) 2.62
去掉回归损失 3.48(不稳定)
去掉回归损失 + TTUR 2.61(恢复稳定)
去掉回归损失 + TTUR + GAN(完整 DMD2) 1.51
仅用 GAN(无分布匹配) 2.56

4. DMD 与 DMD2 对比总结

维度 DMD DMD2
核心目标 KL 散度分布匹配(双 score 之差) 同左 + GAN 目标
回归损失 必需(LPIPS,配对数据集) 去除(真正的分布匹配)
离线数据构造 需要教师多步采样生成噪声-图像对(昂贵) 不需要
稳定性保障 回归损失 TTUR(μ_fake 与 G 更新频率 5:1)
是否用真实数据 否(学生只见教师输出) 是(GAN 判别器用真实图像)
采样步数 仅 1 步 1 步或多步(反向模拟消除训练-推理失配)
质量上限 受教师采样轨迹限制 可超越教师
代表成绩 ImageNet FID 2.62;COCO FID 11.49 ImageNet FID 1.28;COCO FID 8.35;SDXL 4 步媲美教师
判别器 无 μ_fake bottleneck 上加分类分支(参数高效)

一句话总结:DMD 提出"用两个扩散模型的 score 之差作为分布匹配梯度"将多步扩散蒸馏为单步生成器,但依赖离线配对数据的回归损失维稳;DMD2 通过 TTUR 稳定化、GAN 损失引入真实数据监督、反向模拟支持多步采样,去掉了回归损失及其带来的成本与性能天花板,实现了超越教师模型的少步生成。


5. 局限性与后续方向

  • DMD:与教师更精细的采样(100/1000 步)仍有轻微质量差距;同时微调 fake score 网络与生成器,显存开销大(可用 LoRA 缓解)。
  • DMD2:图像多样性相比教师略有下降;最大规模模型(SDXL)仍需 4 步才能匹配教师质量;训练时 guidance scale 固定,用户灵活性受限;未来可结合人类反馈(RLHF)或奖励模型进一步提升。

分布匹配蒸馏
https://huan-yin.github.io/2026/08/17/分布匹配蒸馏/
作者
李相越
发布于
2026年8月17日
许可协议