一步扩散模型与 f f f-散度分布匹配

paper是英伟达发布在Arxiv 2025的工作

paper title:One-step Diffusion Models with f f f-Divergence Distribution Matching

Code:链接

Abstract

从扩散模型中采样涉及一个缓慢的迭代过程,这阻碍了其在实际应用中的部署,尤其是在交互式应用中。为了加速生成速度,近年来的方法通过变分评分蒸馏(variational score distillation)将多步扩散模型蒸馏到单步学生生成器中,从而使得学生生成的样本分布匹配教师模型的分布。然而,这些方法使用逆 Kullback-Leibler(KL)散度进行分布匹配,而这种方式已知具有模式塌陷的倾向。
在本文中,我们提出了一种基于 f f f-散度最小化的新框架,称为 f f f-distill,它能够涵盖不同的散度,并在模式覆盖(mode coverage)与训练方差(training variance)之间实现不同的权衡。我们推导了教师分布和学生分布之间 f f f-散度 的梯度,并表明该梯度可以表达为它们的评分(score)差异的乘积,以及由密度比(density ratio)决定的加权函数。该加权函数在教师分布具有较高密度的样本上赋予更高权重,从而在使用较弱的模式塌陷度(less mode-seeking)的 f f f-散度时,能够更好地匹配分布。
我们观察到,传统的变分评分蒸馏方法(variational score distillation)使用逆 KL 散度进行匹配,是我们框架的一个特例。通过实验,我们证明了其他 f f f-散度(如前向 KL(forward-KL)和Jensen-Shannon 散度)在多个图像生成任务中,能够超越基于逆 KL 散度的最佳变分评分蒸馏方法。特别是,在使用Jensen-Shannon 散度时, f f f-distill 在 ImageNet64 的一步图像生成和MS-COCO 的零样本文本生成任务中达到了当前最先进(state-of-the-art)的生成性能。

1. Introduction

扩散模型 [12, 58] 正在改变视觉领域的生成建模,在图像生成 [1, 43, 44]、视频生成 [13, 53]、3D 物体 [32, 75]、运动建模 [74, 76] 等任务上取得了令人印象深刻的成功。然而,在现实世界应用中,扩散模型的一个关键限制是其采样过程缓慢且计算成本高昂,因为它涉及对去噪神经网络的多次迭代调用。

早期加速扩散模型的方法依赖于改进的数值求解器,以求解描述扩散模型采样过程的常微分方程(ODEs)或随机微分方程(SDEs)[15, 17, 29, 54, 67]。然而,由于离散化误差的存在,这些方法最多只能将采样步骤减少到几十步。

最近,基于蒸馏的方法致力于将采样步骤减少到单次网络调用。这些方法大致可以分为两类:
1)轨迹蒸馏(trajectory distillation) [9, 23, 28, 55, 59],它将扩散模型中从噪声到数据的确定性 ODE 映射蒸馏到一步学生模型;
2)分布匹配方法(distribution matching approaches) [71, 72, 79, 82],它们忽略了确定性映射,而是让一步学生生成的样本分布匹配预训练教师扩散模型的分布。在这两类方法中,后一种方法通常在实践中表现更优,因为从噪声到数据的确定性映射较为复杂,难以学习。
自然地,分布匹配中的散度选择起着关键作用,因为它决定了如何让学生分布匹配教师分布。现有研究 [71, 72] 通常采用变分评分蒸馏(variational score distillation) [64],其通过最小化**逆 KL 散度(reverse-KL divergence)来匹配学生与教师的分布。然而,逆 KL 散度是已知的模式寻优(mode-seeking)**散度 [2],可能会忽略扩散模型学习到的多样性模式。

在本工作中,我们提出了一种基于 f f f-散度的新型分布匹配蒸馏框架,称为 f f f-distill。 f f f-散度是一个广义的散度家族,包括逆 KL(reverse-KL)、前向 KL(forward-KL)、Jensen-Shannon(JS)、平方 Hellinger 散度等。这些散度在惩罚学生忽略教师分布中的模式方面具有不同的权衡,并且可以通过**蒙特卡洛采样(Monte Carlo sampling)**进行估计和优化。在我们的框架中,我们基于这些属性评估了不同的 f f f-散度,并观察到不同的权衡。例如:

  • 前向 KL(forward-KL) 具有更好的模式覆盖能力,但梯度方差较大;
  • Jensen-Shannon(JS) 在训练早期表现出中等程度的模式寻优和梯度饱和现象,但整体梯度方差较低。

我们的分析表明,没有任何单一的 f f f-散度在所有数据集上始终优于其他散度。我们观察到,具有更好模式覆盖能力的散度通常在 CIFAR-10 上表现更佳;然而,在 ImageNet-64 及 Stable Diffusion(SD)文本生成等大规模任务上,梯度方差较低的散度能够取得更优的结果

我们推导了 f f f-散度分布匹配的梯度,其形式如下:
Gradient = ( teacher score − student score ) × weighting function \text{Gradient} = (\text{teacher score} - \text{student score}) \times \text{weighting function} Gradient=(teacher scorestudent score)×weighting function
其中,权重函数(weighting function) 由密度比和所选 f f f-散度决定(这一点是本工作的新贡献),如图1。该密度比可以直接由 GAN 目标中的判别器获得。我们证明了先前的 DMD 方法是本方法的一个特例,即它使用了一个常数权重。我们进一步讨论了新推导出的权重系数如何影响上述权衡,并提出了归一化技术,以稳定梯度方差较大的散度。

图1

图1. f-distill中一步学生的梯度更新。梯度是教师分数和假分数之间的差值与由所选f-散度和密度比确定的加权函数的乘积。密度比可以从辅助GAN目标中的函数中轻松获得。

图2

图 2. 评分差异与加权函数的二维示例。 h h h 是前向-KL(forward-KL)中的加权函数。
可以观察到,在低密度区域,教师评分和伪评分之间通常存在较大差异(底部左图中的深色区域表示较大的评分差异),这些区域的评分估计误差较大。 在 f f f-distill 的梯度更新过程中,加权函数会降低这些区域的权重(底部右图中的浅色区域)。

如图 2 所示,我们观察到:对于不太倾向于模式寻优的 f f f-散度,权重系数会在教师分布密度较低的区域降低评分差异的影响。这一现象与最近的研究相一致,即低密度区域的评分估计通常不准确 [18]。因此,我们的方法可以自适应地减少在这些区域与教师模型评分匹配的依赖。

经验验证:
我们在多个图像生成任务上验证了 f f f-distill 框架。定量实验表明,在 f f f-distill 中,较少模式寻优的散度一致优于以往的变分评分蒸馏方法。值得注意的是,通过最小化Jensen-Shannon 散度(一种较少模式寻优且梯度方差较低的散度), f f f-distill 在 ImageNet-64 一步生成任务和 MS-COCO 零样本文本生成任务(使用 SD v1.5)上达到了新的 SOTA(state-of-the-art)性能。此外,我们的实验进一步证实了权重函数能够有效地为具有较大评分差异的区域分配较小的权重。

贡献:
(i) 我们提出了一种基于 f f f-散度的分布匹配蒸馏的新泛化框架,使得学生分布可以更加灵活地匹配教师分布。
(ii) 我们讨论了不同 f f f-散度在模式寻优梯度饱和方差上的不同权衡。
(iii) 我们提供了减少梯度方差和高效估计目标函数中不同项的实用指南。
(iv) 通过实验,我们证明了 f f f-distill 在 ImageNet-64 和 MS-COCO 文本生成基准上达到了最先进的 FID 评分,实现了一步生成的新 SOTA 性能。

2. Background

2.1. Diffusion models


f f f-distill 的目标是加速预训练(连续时间)扩散模型(DMs)的生成过程 [12, 58]。在本文中,我们遵循流行的 EDM 框架 [17] 进行符号表示,并定义前向/后向过程。扩散模型(DMs)在固定的前向过程中使用方差 σ 2 ( t ) \sigma^2(t) σ2(t) 的高斯噪声来扰动干净数据 x 0 ∼ p data \mathbf{x}_0 \sim p_{\text{data}} x0pdata,其中 x 0 ∈ R d \mathbf{x}_0 \in \mathbb{R}^d x0Rd t t t 表示扩散过程中的时间。得到的中间分布记作 p t ( x t ) p_t\left(\mathbf{x}_t\right) pt(xt),其中 x t ∈ R d \mathbf{x}_t \in \mathbb{R}^d xtRd。为简化符号表示,除非另有说明,我们将在全文中用 x x x 代替 x t \mathbf{x}_t xt。当 σ max ⁡ \sigma_{\max} σmax 充分大时,该分布几乎与纯高斯噪声相同。扩散模型利用这一观察,首先对初始噪声进行采样 ϵ max ⁡ ∼ N ( 0 , σ max ⁡ 2 I ) \epsilon_{\max} \sim \mathcal{N}\left(\mathbf{0}, \sigma_{\max}^2 \boldsymbol{I}\right) ϵmaxN(0,σmax2I),然后通过求解以下后向 ODE/SDE 迭代去噪,确保当 σ ( 0 ) = 0 \sigma(0) = 0 σ(0)=0 时,最终 x \mathbf{x} x 服从数据分布 p data p_{\text{data}} pdata

d x = − σ ˙ ( t ) σ ( t ) ∇ x log ⁡ p t ( x ) d t ⏟ 概率流 ODE(Probability Flow ODE) − β ( t ) σ 2 ( t ) ∇ x log ⁡ p t ( x ) d t + 2 β ( t ) σ ( t ) d ω t ⏟ 朗之万扩散 SDE(Langevin Diffusion SDE) \begin{aligned} & d \mathbf{x}=\underbrace{-\dot{\sigma}(t) \sigma(t) \nabla_{\mathbf{x}} \log p_t(\mathbf{x}) d t}_{\text{概率流 ODE(Probability Flow ODE)}} \\ & \underbrace{-\beta(t) \sigma^2(t) \boldsymbol{\nabla}_{\mathbf{x}} \log p_t(\mathbf{x}) d t+\sqrt{2 \beta(t)} \sigma(t) d \omega_t}_{\text{朗之万扩散 SDE(Langevin Diffusion SDE)}} \end{aligned} dx=概率流 ODEProbability Flow ODE σ˙(t)σ(t)xlogpt(x)dt朗之万扩散 SDELangevin Diffusion SDE β(t)σ2(t)xlogpt(x)dt+2β(t) σ(t)dωt

其中, ω t \omega_t ωt 是标准维纳过程(Wiener process), ∇ x log ⁡ p t ( x ) \nabla_{\mathbf{x}} \log p_t(\mathbf{x}) xlogpt(x) 是中间分布 p t ( x ) p_t(\mathbf{x}) pt(x) 的评分函数(score function)。评分函数由一个神经网络 s ϕ ( x ; σ ( t ) ) s_\phi(\mathbf{x} ; \sigma(t)) sϕ(x;σ(t)) 通过去噪评分匹配(denoising score matching)目标进行训练 [56, 62]。

在方程 (1) 中:

  • 第一项是概率流 ODE(Probability Flow ODE),它引导样本从高噪声水平向低噪声水平演化。
  • 第二项是朗之万扩散 SDE(Langevin Diffusion SDE),它在不同噪声水平 σ ( t ) \sigma(t) σ(t) 之间充当平衡采样器(equilibrium sampler),有效地优化样本并在采样过程中修正误差 [17, 67]。该项可以由时间相关参数 β ( t ) \beta(t) β(t) 进行缩放。设定 β ( t ) = 0 \beta(t) = 0 β(t)=0 将导致纯 ODE 生成。

然而,求解扩散 ODE 和 SDE 通常需要大量迭代(通常为几十或几百步),这对扩散模型的实际部署提出了重大挑战。尽管已经提出了多种加速扩散 ODE [17, 29, 54] 和 SDE [15, 17, 67] 采样的方法,但在实际应用中,它们通常仍然需要 > 20 >20 >20 步才能生成高质量的样本。

2.2. Variational score distillation


最近的一系列研究 [71,72] 旨在通过变分评分蒸馏(Variational Score Distillation, VSD),将教师扩散模型 s ϕ s_\phi sϕ 蒸馏到单步生成器 G θ G_\theta Gθ 中。VSD 最初被提出用于 3D 物体的测试时优化(test-time optimization)[64]。其目标是让学生模型 G θ G_\theta Gθ 直接将噪声 z \mathbf{z} z 从先验分布 p ( z ) = N ( z ; 0 , I ) p(\mathbf{z}) = \mathcal{N}(\mathbf{z} ; \mathbf{0}, \boldsymbol{I}) p(z)=N(z;0,I) 映射到 σ = 0 \sigma=0 σ=0 处的干净样本 x 0 \mathbf{x}_0 x0,即通过 x 0 = G θ ( z ) \mathbf{x}_0 = G_\theta(\mathbf{z}) x0=Gθ(z) 实现一步生成,从而绕过传统的迭代采样过程。

p ϕ p_\phi pϕ 表示将预训练的扩散模型 s ϕ ( x ; σ ( t ) ) s_\phi(\mathbf{x} ; \sigma(t)) sϕ(x;σ(t)) 代入方程 (1) 所得到的分布, q θ q_\theta qθ 表示一步生成器 G θ G_\theta Gθ 的输出分布(为了符号简洁,下文省略 p ϕ p_\phi pϕ q θ q_\theta qθ 的下标)。那么,生成器的梯度更新可表示为:

E t , z , ϵ [ ( s ϕ ( x ; σ ( t ) ) − ∇ x log ⁡ q θ ( x ; σ ( t ) ) ) ∇ θ G θ ( z ) ] \mathbb{E}_{t, \mathbf{z}, \epsilon}\left[\left(s_\phi(\mathbf{x} ; \sigma(t))-\nabla_{\mathbf{x}} \log q_\theta(\mathbf{x} ; \sigma(t))\right) \nabla_\theta G_\theta(\mathbf{z})\right] Et,z,ϵ[(sϕ(x;σ(t))xlogqθ(x;σ(t)))θGθ(z)]

其中 x = G θ ( z ) + σ ( t ) ϵ \mathbf{x} = G_\theta(\mathbf{z}) + \sigma(t) \epsilon x=Gθ(z)+σ(t)ϵ,且 ϵ ∼ N ( 0 , I ) \epsilon \sim \mathcal{N}(\mathbf{0}, \boldsymbol{I}) ϵN(0,I)

直观上,该梯度鼓励生成器在数据分布的高密度区域生成样本。这是通过教师评分项 s ϕ ( x ; σ ( t ) ) s_\phi(\mathbf{x} ; \sigma(t)) sϕ(x;σ(t)) 实现的,该项引导生成样本向教师模型赋予高概率的区域靠拢。为了防止模式崩塌(mode collapse),梯度还包含一个项,用于阻止生成器仅集中在教师分布的某个单一高密度点上。这是通过减去学生分布的评分项 ∇ x log ⁡ q θ ( x ; σ ( t ) ) \nabla_{\mathbf{x}} \log q_\theta(\mathbf{x} ; \sigma(t)) xlogqθ(x;σ(t)) 实现的。

已证明,该梯度更新可通过最小化逆 KL 散度(reverse-KL divergence) 来执行分布匹配 [39, 72],从而让学生分布更接近教师分布。

估计学生分布的评分函数
为了估计学生分布的评分函数,先前的研究 [64, 72] 引入了伪评分网络(fake score network) s ψ ( x , σ ( t ) ) s_\psi(\mathbf{x}, \sigma(t)) sψ(x,σ(t)) 来近似 ∇ x log ⁡ q θ ( x ; σ ( t ) ) \nabla_{\mathbf{x}} \log q_\theta(\mathbf{x} ; \sigma(t)) xlogqθ(x;σ(t))。该伪评分网络 s ψ ( x , σ ( t ) ) s_\psi(\mathbf{x}, \sigma(t)) sψ(x,σ(t)) 通过标准的去噪评分匹配损失(denoising score matching loss)进行动态更新,在训练过程中,“干净”样本来自生成器 G θ G_\theta Gθ。因此,VSD 训练过程在生成器更新伪评分更新之间交替进行,并采用双时间尺度(two time-scale)更新规则来稳定训练 [71]。

结合 GAN 损失
此外,为了进一步缩小一步生成器与多步教师扩散模型之间的差距,VSD 训练管道中引入了GAN 损失 [71],其中一个轻量级 GAN 分类器以伪评分网络的中间特征作为输入,以提高生成样本的质量。

2.3. f f f-divergence


在概率论中, f f f-散度( f f f-divergence)[41] 用于量化两个概率密度函数 p p p q q q 之间的差异。具体来说,当 p p p 关于 q q q 绝对连续时, f f f-散度定义如下:

D f ( p ∣ ∣ q ) = ∫ q ( x ) f ( p ( x ) q ( x ) ) d x D_f(p || q) = \int q(\mathbf{x}) f\left(\frac{p(\mathbf{x})}{q(\mathbf{x})}\right) d\mathbf{x} Df(p∣∣q)=q(x)f(q(x)p(x))dx

其中, f f f 是定义在 ( 0 , + ∞ ) (0, +\infty) (0,+) 上的凸函数,并满足 f ( 1 ) = 0 f(1) = 0 f(1)=0。该散度具有多个重要性质,包括非负性数据处理不等式

许多常见的散度可以通过选择合适的 f f f 函数作为 f f f-散度的特例来表示。这些特例包括:

  • 前向 KL 散度(forward-KL divergence)
  • 逆向 KL 散度(reverse-KL divergence)
  • Hellinger 距离(Hellinger distance)
  • Jensen-Shannon(JS)散度

这些散度在表 1 中列出。

在生成式学习(generative learning)中, f f f-散度已被广泛应用于多个流行的生成模型,包括:

  • 生成对抗网络(GANs)[37]
  • 变分自编码器(VAEs)[63]
  • 能量模型(energy-based models)[73]
  • 扩散模型(diffusion models)[60]

表1

表 1. 不同 f f f-散度的比较,它们作为似然比 r : = p ( x ) / q ( x ) r := p(\mathbf{x}) / q(\mathbf{x}) r:=p(x)/q(x) 的函数

3. Method: general f f f-divergence minimization

在本节中,我们介绍了一种基于 f f f-散度最小化的通用蒸馏框架,称为 f f f-distill,其目标是最小化教师分布和学生分布之间的 f f f-散度。由于学生分布 q q q 由一步生成器 G θ G_\theta Gθ 诱导的推前测度(push-forward measure)定义,因此它隐式依赖于生成器的参数 θ \theta θ。由于这种隐式依赖关系,直接计算 f f f-散度 D f ( p ∥ q ) D_f(p \| q) Df(pq) 关于 θ \theta θ 的梯度存在挑战。然而,以下定理提供了该梯度的解析表达式,表明该梯度可以表示为**变分评分蒸馏(VSD)**中梯度的加权版本。值得注意的是,这些权重由生成样本的密度比(density ratio)决定。

我们以更一般的形式给出该定理,即提供 p t p_t pt q t q_t qt 之间的梯度,其中 p t p_t pt 是教师分布 p p p 通过扩散前向过程扰动后的分布,即:
p t = p 0 ∗ N ( 0 , σ 2 ( t ) I ) p_t=p_0 * \mathcal{N}\left(\mathbf{0}, \sigma^2(t) \boldsymbol{I}\right) pt=p0N(0,σ2(t)I)
学生分布 q t q_t qt 采用相同的定义。

定理 1:

p p p 为教师生成分布, q q q 为通过可微映射 G θ G_\theta Gθ 从先验分布 p ( z ) p(\mathbf{z}) p(z) 诱导出的分布。假设 f f f二阶连续可微的,则 p t p_t pt q t q_t qt 之间的 f f f-散度关于 θ \theta θ 的梯度为:
∇ θ D f ( p t ∥ q t ) = E z , ϵ − [ f ′ ′ ( p t ( x ) q t ( x ) ) ( p t ( x ) q t ( x ) ) 2 ( ∇ x log ⁡ p t ( x ) ⏟ 教师评分(teacher score) − ∇ x log ⁡ q t ( x ) ⏟ 伪评分(fake score) ) ∇ θ G θ ( z ) ] \begin{array}{r} \nabla_\theta D_f\left(p_t \| q_t\right)=\mathbb{E}_{\mathbf{z}, \epsilon}-\left[f^{\prime \prime}\left(\frac{p_t(\mathbf{x})}{q_t(\mathbf{x})}\right)\left(\frac{p_t(\mathbf{x})}{q_t(\mathbf{x})}\right)^2\right. \\ (\underbrace{\nabla_{\mathbf{x}} \log p_t(\mathbf{x})}_{\text{教师评分(teacher score)}}-\underbrace{\nabla_{\mathbf{x}} \log q_t(\mathbf{x})}_{\text{伪评分(fake score)}}) \nabla_\theta G_\theta(\mathbf{z})] \end{array} θDf(ptqt)=Ez,ϵ[f′′(qt(x)pt(x))(qt(x)pt(x))2(教师评分(teacher score xlogpt(x)伪评分(fake score xlogqt(x))θGθ(z)]
其中 z ∼ p ( z ) , ϵ ∼ N ( 0 , I ) \mathbf{z} \sim p(\mathbf{z}), \epsilon \sim \mathcal{N}(\mathbf{0}, \boldsymbol{I}) zp(z),ϵN(0,I),且 x = G θ ( z ) + σ ( t ) ϵ \mathbf{x}=G_\theta(\mathbf{z})+\sigma(t) \epsilon x=Gθ(z)+σ(t)ϵ

证明思路:

为了简化,我们在正文中仅证明 t = 0 t=0 t=0 的情况,类似的推导可适用于任意 t > 0 t>0 t>0

∇ θ D f ( p ( x ) ∥ q ( x ) ) = ∇ θ ∫ q ( x ) f ( p ( x ) q ( x ) ) d x = ∫ ∇ θ q ( x ) f ( p ( x ) q ( x ) ) d x ⏟ I − ∫ ∇ θ q ( x ) f ′ ( p ( x ) q ( x ) ) p ( x ) q ( x ) d x ⏟ I I \begin{aligned} & \nabla_\theta D_f(p(\mathbf{x}) \| q(\mathbf{x}))=\nabla_\theta \int q(\mathbf{x}) f\left(\frac{p(\mathbf{x})}{q(\mathbf{x})}\right) d \mathbf{x} \\ & =\underbrace{\int \nabla_\theta q(\mathbf{x}) f\left(\frac{p(\mathbf{x})}{q(\mathbf{x})}\right) d \mathbf{x}}_I-\underbrace{\int \nabla_\theta q(\mathbf{x}) f^{\prime}\left(\frac{p(\mathbf{x})}{q(\mathbf{x})}\right) \frac{p(\mathbf{x})}{q(\mathbf{x})} d \mathbf{x}}_{II} \end{aligned} θDf(p(x)q(x))=θq(x)f(q(x)p(x))dx=I θq(x)f(q(x)p(x))dxII θq(x)f(q(x)p(x))q(x)p(x)dx

可以证明,项 (I) 和 (II) 分别来源于 f f f 关于 x \mathbf{x} x q q q 的偏导数。我们可以使用如下恒等式来处理这两个项:
∫ ∇ θ q ( x ) g ( x ) d x = ∫ p ( z ) ∇ x g ( x ) ∇ θ G θ ( z ) d z \int \nabla_\theta q(\mathbf{x}) g(\mathbf{x}) d \mathbf{x}=\int p(\mathbf{z}) \nabla_{\mathbf{x}} g(\mathbf{x}) \nabla_\theta G_\theta(\mathbf{z}) d \mathbf{z} θq(x)g(x)dx=p(z)xg(x)θGθ(z)dz
该恒等式的证明见附录。

利用该恒等式,我们可以简化 (I) 和 (II) 为:
I = ∫ p ( z ) f ′ ( p ( x ) q ( x ) ) ∇ x p ( x ) q ( x ) ∇ θ G θ ( z ) d z I I = ∫ p ( z ) f ′ ′ ( p ( x ) q ( x ) ) p ( x ) q ( x ) ∇ x p ( x ) q ( x ) ∇ θ G θ ( z ) d z + ∫ p ( z ) f ′ ( p ( x ) q ( x ) ) ∇ x p ( x ) q ( x ) ∇ θ G θ ( z ) d z \begin{aligned} I &=\int p(\mathbf{z}) f^{\prime}\left(\frac{p(\mathbf{x})}{q(\mathbf{x})}\right) \nabla_{\mathbf{x}} \frac{p(\mathbf{x})}{q(\mathbf{x})} \nabla_\theta G_\theta(\mathbf{z}) d \mathbf{z} \\ II &=\int p(\mathbf{z}) f^{\prime \prime}\left(\frac{p(\mathbf{x})}{q(\mathbf{x})}\right) \frac{p(\mathbf{x})}{q(\mathbf{x})} \nabla_{\mathbf{x}} \frac{p(\mathbf{x})}{q(\mathbf{x})} \nabla_\theta G_\theta(\mathbf{z}) d \mathbf{z} \\ & +\int p(\mathbf{z}) f^{\prime}\left(\frac{p(\mathbf{x})}{q(\mathbf{x})}\right) \nabla_{\mathbf{x}} \frac{p(\mathbf{x})}{q(\mathbf{x})} \nabla_\theta G_\theta(\mathbf{z}) d \mathbf{z} \end{aligned} III=p(z)f(q(x)p(x))xq(x)p(x)θGθ(z)dz=p(z)f′′(q(x)p(x))q(x)p(x)xq(x)p(x)θGθ(z)dz+p(z)f(q(x)p(x))xq(x)p(x)θGθ(z)dz

将 (I) 和 (II) 代入方程 (3),我们得到:
∇ θ D f = − ∫ p ( z ) f ′ ′ ( p ( x ) q ( x ) ) p ( x ) q ( x ) ∇ x p ( x ) q ( x ) ∇ θ G θ ( z ) d z = − ∫ p ( z ) f ′ ′ ( p ( x ) q ( x ) ) ( p ( x ) q ( x ) ) 2 [ ∇ x log ⁡ p ( x ) − ∇ x log ⁡ q ( x ) ] ∇ θ G θ ( z ) d z \begin{aligned} \nabla_\theta D_f&=-\int p(\mathbf{z}) f^{\prime \prime}\left(\frac{p(\mathbf{x})}{q(\mathbf{x})}\right) \frac{p(\mathbf{x})}{q(\mathbf{x})} \nabla_{\mathbf{x}} \frac{p(\mathbf{x})}{q(\mathbf{x})} \nabla_\theta G_\theta(\mathbf{z}) d \mathbf{z} \\ &=-\int p(\mathbf{z}) f^{\prime \prime}\left(\frac{p(\mathbf{x})}{q(\mathbf{x})}\right)\left(\frac{p(\mathbf{x})}{q(\mathbf{x})}\right)^2 \\ & \left[\nabla_{\mathbf{x}} \log p(\mathbf{x})-\nabla_{\mathbf{x}} \log q(\mathbf{x})\right] \nabla_\theta G_\theta(\mathbf{z}) d \mathbf{z} \end{aligned} θDf=p(z)f′′(q(x)p(x))q(x)p(x)xq(x)p(x)θGθ(z)dz=p(z)f′′(q(x)p(x))(q(x)p(x))2[xlogp(x)xlogq(x)]θGθ(z)dz

最后一步使用了对数导数技巧(log derivative trick)

完整的证明见附录 A。尽管学生的生成分布 q q q 依赖于参数 θ \theta θ,但定理 1 提供了教师分布与学生分布之间 f f f-散度梯度的解析表达式。该梯度由教师评分与学生评分的差值表示,并加权一个时间相关系数
f ′ ′ ( p t ( x t ) / q t ( x t ) ) ( p t ( x t ) / q t ( x t ) ) 2 f^{\prime \prime}\left(p_t\left(\mathbf{x}_t\right) / q_t\left(\mathbf{x}_t\right)\right)\left(p_t\left(\mathbf{x}_t\right) / q_t\left(\mathbf{x}_t\right)\right)^2 f′′(pt(xt)/qt(xt))(pt(xt)/qt(xt))2
该权重由所选的 f f f-散度以及密度比共同决定。值得注意的是,该定理中的所有项都是可计算的,因此可以通过一般的 f f f-散度最小化来优化分布匹配。

为了方便表示,我们定义:
h ( r ) : = f ′ ′ ( r ) r 2 , r t ( x ) : = p t ( x ) / q t ( x ) h(r):=f^{\prime \prime}(r) r^2, \quad r_t(\mathbf{x}):=p_t(\mathbf{x}) / q_t(\mathbf{x}) h(r):=f′′(r)r2,rt(x):=pt(x)/qt(x)
其中 h ( r ) h(r) h(r) 为加权函数, r t ( x ) r_t(\mathbf{x}) rt(x) 为时间 t t t 处的密度比。

值得注意的是,变分评分蒸馏(VSD)的梯度(方程 (2))可以看作我们框架的特例,即在方程 (3) 中设定 h ( r ) ≡ 1 h(r) \equiv 1 h(r)1,这对应于最小化逆 KL 散度 f ( r ) = − log ⁡ r f(r)=-\log r f(r)=logr)。

[57] 也研究了 f f f-散度与评分差异的关系,并在其定理 2 中将 f f f-散度表示为评分差异的时间积分。然而,我们的公式在两方面不同:

  1. 他们的梯度需要计算 Jacobian-向量积,计算成本较高。
  2. 他们的 f f f-散度是评分差异的时间积分,而我们的方法仅依赖于时间 t t t 处的加权和评分。

在接下来的命题中,我们进一步证明,如果加权函数 h h h ( 0 , + ∞ ) (0,+\infty) (0,+) 上是连续且非负的,则其与评分差异的乘积是某些 f f f-散度的梯度:

命题 1:
对于任意在 ( 0 , + ∞ ) (0,+\infty) (0,+) 上连续且非负的函数 h h h,期望值
E z , ϵ [ h ( r t ( x ) ) ( ∇ x log ⁡ p t ( x ) − ∇ x log ⁡ q t ( x ) ) ∇ θ G θ ( z ) ] \mathbb{E}_{\mathbf{z}, \epsilon} \left[h\left(r_t(\mathbf{x})\right)\left(\nabla_{\mathbf{x}} \log p_t(\mathbf{x})-\nabla_{\mathbf{x}} \log q_t(\mathbf{x})\right) \nabla_\theta G_\theta(\mathbf{z})\right] Ez,ϵ[h(rt(x))(xlogpt(x)xlogqt(x))θGθ(z)]
对应于某个 f f f-散度的梯度。

尽管本文的研究范围局限于 f f f-散度的标准形式,命题 1 允许我们将任何连续且非负的标量函数用作 h h h

在实践中,[72] 建议在整个扩散过程中始终执行分布匹配,因为在较小时间 t t t 处,教师模型和学生模型之间的差异较大,导致优化困难。我们遵循这一设置,并在整个时间范围上最小化 f f f-散度,即:
L ( θ ) = ∫ 0 T w t D f ( p t ∥ q t ) d t \mathcal{L}(\theta)=\int_0^T w_t D_f\left(p_t \| q_t\right) d t L(θ)=0TwtDf(ptqt)dt
其中 w t w_t wt 是时间相关权重,用于平衡不同时刻的梯度幅度。

最终, f f f-distill 的目标函数如下:
L f -distill  ( θ ) = E t , x [ sg ⁡ ( w t h ( r t ( x ) ) ( ∇ x log ⁡ p t ( x ) − ∇ x log ⁡ q t ( x ) ) ) T x ] \begin{aligned} \mathcal{L}_{f \text{-distill }}(\theta) &=\mathbb{E}_{t, \mathbf{x}} \left[ \operatorname{sg} \left(w_t h\left(r_t(\mathbf{x})\right) \right.\right. \\ &\quad \left.\left. \left(\nabla_{\mathbf{x}} \log p_t(\mathbf{x})-\nabla_{\mathbf{x}} \log q_t(\mathbf{x})\right)\right)^T \mathbf{x} \right] \end{aligned} Lf-distill (θ)=Et,x[sg(wth(rt(x))(xlogpt(x)xlogqt(x)))Tx]

其中:
x = G θ ( z ) + σ ( t ) ϵ , z ∼ p ( z ) , ϵ ∼ N ( 0 , I ) \mathbf{x}=G_\theta(\mathbf{z})+\sigma(t) \epsilon, \quad \mathbf{z} \sim p(\mathbf{z}), \quad \epsilon \sim \mathcal{N}(\mathbf{0}, \boldsymbol{I}) x=Gθ(z)+σ(t)ϵ,zp(z),ϵN(0,I)
sg ⁡ \operatorname{sg} sg 代表 停止梯度(stop gradient)

方程 (4) 的梯度等于定理 1(方程 (3))中 f f f-散度梯度的时间积分。在实际操作中,学生分布的评分函数 ∇ x log ⁡ q t ( x t ) \nabla_{\mathbf{x}} \log q_t\left(\mathbf{x}_t\right) xlogqt(xt)在线扩散模型(online diffusion model) s ψ ( x , σ ( t ) ) s_\psi(\mathbf{x}, \sigma(t)) sψ(x,σ(t)) 近似。

结合 GAN 目标:
[71] 在变分评分蒸馏损失的基础上加入了 GAN 目标,以进一步提升性能。这一策略的动机在于,变分评分蒸馏完全依赖于教师模型的评分函数,因此会受到教师模型能力的限制。通过引入 GAN 目标,学生生成器 G θ G_\theta Gθ 可以通过利用真实数据训练判别器 D λ D_\lambda Dλ,突破教师模型的局限性,其目标函数为:
L G A N ( λ ) = E t , x ∼ p data , ϵ 1 [ log ⁡ D λ ( x + σ ( t ) ϵ 1 ) ] + E t , z , ϵ 2 [ log ⁡ ( 1 − D λ ( G θ ( z ) + σ ( t ) ϵ 2 ) ) ] \begin{aligned} \mathcal{L}_{\mathrm{GAN}}(\lambda) &=\mathbb{E}_{t, \mathbf{x} \sim p_{\text{data}}, \epsilon_1}\left[\log D_\lambda\left(\mathbf{x}+\sigma(t) \epsilon_1\right)\right] \\ &\quad + \mathbb{E}_{t, \mathbf{z}, \epsilon_2}\left[\log \left(1-D_\lambda\left(G_\theta(\mathbf{z})+\sigma(t) \epsilon_2\right)\right)\right] \end{aligned} LGAN(λ)=Et,xpdata,ϵ1[logDλ(x+σ(t)ϵ1)]+Et,z,ϵ2[log(1Dλ(Gθ(z)+σ(t)ϵ2))]
其中:
z ∼ p ( z ) , ϵ 1 , ϵ 2 ∼ N ( 0 , I ) \mathbf{z} \sim p(\mathbf{z}), \quad \epsilon_1, \epsilon_2 \sim \mathcal{N}(\mathbf{0}, \boldsymbol{I}) zp(z),ϵ1,ϵ2N(0,I)

我们参考了先前研究的辅助 GAN 目标,并利用其提供的额外优势,即:
GAN 判别器 D λ D_\lambda Dλ 能直接提供密度比 r ( x t ) r\left(\mathbf{x}_t\right) r(xt) 的估计值,该密度比在方程 (4) 的加权函数计算中是必需的。具体来说,该密度比可通过以下近似计算:
r ( x t ) = p t ( x ) q t ( x ) ≈ p data , t q t ( x ) = D λ ( x t , t ) 1 − D λ ( x t , t ) r\left(\mathbf{x}_t\right)=\frac{p_t(\mathbf{x})}{q_t(\mathbf{x})} \approx \frac{p_{\text{data}, t}}{q_t(\mathbf{x})}=\frac{D_\lambda\left(\mathbf{x}_t, t\right)}{1-D_\lambda\left(\mathbf{x}_t, t\right)} r(xt)=qt(x)pt(x)qt(x)pdata,t=1Dλ(xt,t)Dλ(xt,t)

本质上,GAN 判别器 D λ D_\lambda Dλ 直接提供了密度比的估计,从而简化了加权函数的计算

4. Comparing properties of f f f-divergence

本节概述:
在本节中,我们比较了 f f f-散度家族中不同距离度量的性质,特别是在扩散蒸馏(diffusion distillation)中的表现。我们考察了以下三个性质:

  1. 模式寻优(Mode-seeking)
  2. 饱和(Saturation)
  3. 训练方差(Variance)

我们在表 1 中总结了不同 f f f-散度的比较以及它们对应的加权函数 h h h


模式寻优(Mode-seeking)
模式寻优散度 [2, 24],如逆 KL 散度(reverse-KL),倾向于使生成分布 q q q 仅覆盖数据分布的部分模式,并避免在数据密度较低的区域分配概率质量。当优化 min ⁡ q D f ( p ∥ q ) = ∫ q f ( p / q ) d x \min _q D_f(p \| q)=\int q f(p / q) d \mathbf{x} minqDf(pq)=qf(p/q)dx 时,这种行为可能会导致生成模型的模式丢失(mode collapse),从而降低生成样本的多样性。

这一现象也出现在 DMD 方法 [71, 72] 所使用的变分评分蒸馏损失 [28] 中,该方法对应于在 f f f-distill 框架下最小化逆 KL 散度
衡量模式寻优行为的一种方法是研究极限:
lim ⁡ r → ∞ f ( r ) r \lim _{r \rightarrow \infty} \frac{f(r)}{r} rlimrf(r)
较低的增长率表明较弱的模式寻优(详细讨论见附录 C)。

  • 逆 KL 散度(reverse-KL)JS 散度(Jensen-Shannon) 都具有有限的极限,其中 JS 的增长率更低(模式寻优性更弱)
  • 前向 KL 散度(forward-KL) 具有无穷大的极限,反映了其模式覆盖(mode-covering)特性。

在表 1 中,我们还观察到:模式寻优程度较强的散度,其加权函数 h ( r ) h(r) h(r) r → ∞ r \to \infty r 处增长较慢
这是因为:
f ′ ′ ( r ) = h ( r ) r 2 f^{\prime \prime}(r) = \frac{h(r)}{r^2} f′′(r)=r2h(r)
较慢的增长使得较大的密度比 p / q p/q p/q 更容易被容忍(即允许 q q q 忽略 p p p 中的某些模式),最终导致更强的模式寻优行为 [52]。
例如:

  • JS 和前向 KL 的 h h h 是递增函数
  • 逆 KL 的 h h h 保持恒定

因此,模式寻优程度较弱的散度,其加权函数往往会在教师分布的低密度区域降低样本权重


饱和(Saturation)
使用 f f f-散度的生成模型(如 GANs [10])面临的一个挑战是饱和问题(saturation)
在训练的早期阶段,生成分布和数据分布的匹配程度较差,导致:

  • p p p 的样本在 q q q 下的概率极低
  • 反之亦然

这会使密度比 p / q p/q p/q 变得极大或极小,从而在梯度极小的区域导致优化问题
从图 3a 可以看出:

  • 平方 Hellinger 散度(squared Hellinger)JS 散度 在极端值处的梯度较小。

然而,在扩散蒸馏(diffusion distillation) 过程中,该问题可以通过初始化学生模型的权重为预训练扩散模型的参数来缓解 [71, 72]。

图3

图 3. 不同 f f f-散度中 f ′ f^{\prime} f 的绝对值 (a) 及加权函数 h ( r ) h(r) h(r) (b) 的变化情况。


训练方差(Variance)
目标函数中加权函数 h h h 的方差对小批量训练的稳定性至关重要。
我们使用归一化方差来衡量不同 f f f-散度的方差:
Var ⁡ q ( f ′ ′ ( p / q ) ( p / q ) 2 E q [ f ′ ′ ( p / q ) ( p / q ) 2 ] ) \operatorname{Var}_q\left(\frac{f^{\prime \prime}(p / q)(p / q)^2}{\mathbb{E}_q\left[f^{\prime \prime}(p / q)(p / q)^2\right]}\right) Varq(Eq[f′′(p/q)(p/q)2]f′′(p/q)(p/q)2)
这样可以确保尺度不变性(scale-invariant comparison)。

从图 4a 可以看出:

  • 前向 KL(forward-KL)Jefferys 散度 的方差随着高斯分布之间的距离增加而显著增加。
  • Jensen-Shannon(JS)散度平方 Hellinger 散度 的方差相对稳定。

方差的稳定性有助于低方差的 JS 散度在实验中的优越表现

图4

图4.(a)归一化方差与两个高斯模型之间的平均差。(B)前向KL w/和w/o归一化的训练损失。


实际考虑(Practical Considerations)
为了减少在模式寻优程度较弱的散度(如图 3b)中常见的高方差问题,我们提出了两阶段归一化方案

第一阶段:时间依赖的密度比归一化

我们利用如下事实:
E q t [ r t ] = 1 \mathbb{E}_{q_t}[r_t] = 1 Eqt[rt]=1
r t r_t rt 的期望应为 1。
为了强制满足该性质,我们将时间范围离散化为多个区间,并在每个区间内使用其均值归一化 r t r_t rt

第二阶段:批次内加权函数归一化

我们直接将加权函数 h h h 归一化为其 mini-batch 内的均值

归一化的重要性
  • 训练过程涉及 f f f-distill 目标函数GAN 目标函数
  • 归一化确保 权重的尺度不变性,因为不同的 f f f-散度的尺度变化可能较大
  • 维持 f f f-distill 目标相对于 GAN 目标的相对重要性,从而提供稳定性和一致性

图 4b 说明,在 ImageNet-64 上使用 前向 KL 散度 进行训练时,上述归一化显著降低了目标函数的方差

我们在 算法 1(Alg 1) 中提供了完整的算法框架。

算法1

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐