从入门到入土 3:Score Matching 与 Guidance——从学习梯度到控制生成

在前两篇文章中,我们搭建了生成模型的平台:用微分方程(ODE/SDE)将噪声转化为数据,并用 Flow Matching 训练了向量场。事实上扩散模型还有另一条路径:学习数据分布的对数梯度(即分数函数),并利用它构造更灵活的采样过程。 本文将介绍分数匹配(Score Matching)与引导技术(Guidance),带你理解如何用“梯度”驱动生成,并精确控制生成内容。


1. 分数函数:概率密度的“方向”

在 Flow Matching 中,我们学习的是向量场 ut(x)u_t(x)ut(x),它告诉我们粒子在空间中的“速度”。而在扩散模型中,另一种核心对象是分数函数(Score Function):给定一个概率密度 p(x)p(x)p(x),其分数函数定义为

∇log⁡p(x)=∇p(x)p(x). \nabla \log p(x) = \frac{\nabla p(x)}{p(x)}. logp(x)=p(x)p(x).

(p.s. 这里用到了∇log⁡x=1x\nabla \log x=\frac{1}{x}logx=x1这个trick,且后面将多次用到)
这个向量场指向概率密度增加最快的方向,即“哪里更可能”。例如,对于高斯分布,分数函数是从当前点指向均值的向量,大小与距离成正比。

为什么需要分数函数?
它为我们提供了一种新的方式去描述分布,并且可以方便地构造出随机微分方程(SDE),使得采样过程更加多样化,甚至可以在生成时引入“引导”来强化条件控制。

1.1 条件分数与边际分数

类似于 Flow Matching 中的概率路径,我们可以定义条件概率路径 pt(x∣z)p_t(x|z)pt(xz)边际概率路径 pt(x)p_t(x)pt(x),以及相应的分数函数:

∇log⁡pt(x∣z),∇log⁡pt(x). \nabla \log p_t(x|z), \quad \nabla \log p_t(x). logpt(xz),logpt(x).

对于高斯条件路径 pt(x∣z)=N(αtz,βt2Id)p_t(x|z) = \mathcal{N}(\alpha_t z, \beta_t^2 I_d)pt(xz)=N(αtz,βt2Id),其条件分数函数有简洁的解析形式:

∇log⁡pt(x∣z)=−x−αtzβt2. \nabla \log p_t(x|z) = -\frac{x - \alpha_t z}{\beta_t^2}. logpt(xz)=βt2xαtz.

这意味着,如果我们能训练一个神经网络 stθ(x)s_t^\theta(x)stθ(x) 来逼近边际分数 ∇log⁡pt(x)\nabla \log p_t(x)logpt(x),那么我们就拥有了描述整个概率路径的“方向场”。


2. 学习分数:Denoising Score Matching

与 Flow Matching 类似,我们无法直接获得边际分数,但可以通过条件分数来间接学习。这就是分数匹配(Score Matching) 的基本思想。

2.1 损失函数

我们定义条件分数匹配损失

LCSM(θ)=Et,z,x∼pt(⋅∣z)[∥stθ(x)−∇log⁡pt(x∣z)∥2]. \mathcal{L}_{\text{CSM}}(\theta) = \mathbb{E}_{t, z, x \sim p_t(\cdot|z)} \left[ \| s_t^\theta(x) - \nabla \log p_t(x|z) \|^2 \right]. LCSM(θ)=Et,z,xpt(z)[stθ(x)logpt(xz)2].

这个损失是可行的,因为右边第二项是我们已知的(如高斯情况下的公式)。而边际分数匹配损失

LSM(θ)=Et,x∼pt[∥stθ(x)−∇log⁡pt(x)∥2] \mathcal{L}_{\text{SM}}(\theta) = \mathbb{E}_{t, x \sim p_t} \left[ \| s_t^\theta(x) - \nabla \log p_t(x) \|^2 \right] LSM(θ)=Et,xpt[stθ(x)logpt(x)2]

则无法直接计算。然而,与 Flow Matching 类似,我们可以证明:

LSM(θ)=LCSM(θ)+C, \mathcal{L}_{\text{SM}}(\theta) = \mathcal{L}_{\text{CSM}}(\theta) + C, LSM(θ)=LCSM(θ)+C,

其中 CCCθ\thetaθ 无关。因此,最小化条件分数匹配损失等价于最小化边际分数匹配损失

2.2 高斯情况下的简化

将高斯条件分数代入,得到:

LCSM(θ)=Et,z,ϵ[∥stθ(αtz+βtϵ)+ϵβt∥2]. \mathcal{L}_{\text{CSM}}(\theta) = \mathbb{E}_{t, z, \epsilon} \left[ \left\| s_t^\theta(\alpha_t z + \beta_t \epsilon) + \frac{\epsilon}{\beta_t} \right\|^2 \right]. LCSM(θ)=Et,z,ϵ[ stθ(αtz+βtϵ)+βtϵ 2].

在扩散模型(如 DDPM)中,通常用一个噪声预测网络 ϵtθ\epsilon_t^\thetaϵtθ 来重参数化:

ϵtθ(x)=−βtstθ(x), \epsilon_t^\theta(x) = -\beta_t s_t^\theta(x), ϵtθ(x)=βtstθ(x),

那么损失就简化为:

LDDPM(θ)=Et,z,ϵ[∥ϵtθ(αtz+βtϵ)−ϵ∥2]. \mathcal{L}_{\text{DDPM}}(\theta) = \mathbb{E}_{t, z, \epsilon} \left[ \| \epsilon_t^\theta(\alpha_t z + \beta_t \epsilon) - \epsilon \|^2 \right]. LDDPM(θ)=Et,z,ϵ[ϵtθ(αtz+βtϵ)ϵ2].

这个形式非常直观:网络学习预测添加到数据上的噪声。这也是“去噪扩散”名称的来源。


3. 用分数函数进行采样:SDE 的扩展

一旦我们训练好了分数网络 stθ≈∇log⁡pts_t^\theta \approx \nabla \log p_tstθlogpt,就可以用它来构造一个随机微分方程,使得其轨迹仍然保持相同的边际分布 ptp_tpt。这一核心结果由 SDE 扩展技巧 给出:

对于任意扩散系数 σt≥0\sigma_t \ge 0σt0,定义 SDE

dXt=[uttarget(Xt)+σt22∇log⁡pt(Xt)]dt+σtdWt, dX_t = \left[ u_t^{\text{target}}(X_t) + \frac{\sigma_t^2}{2} \nabla \log p_t(X_t) \right] dt + \sigma_t dW_t, dXt=[uttarget(Xt)+2σt2logpt(Xt)]dt+σtdWt,

Xt∼ptX_t \sim p_tXtpt 对所有 ttt 成立。特别地,X1∼pdataX_1 \sim p_{\text{data}}X1pdata

这个结果非常强大:它允许我们在已学习到的向量场 uttargetu_t^{\text{target}}uttarget(或分数函数)基础上,通过添加一个可调节的噪声项来改变采样过程,而仍然保持正确的边际分布。如果我们将 uttargetu_t^{\text{target}}uttarget 替换为通过学习得到的 utθu_t^\thetautθ 或直接用分数表示,就得到了实际可用的采样算法。

3.1 高斯情况下的 SDE

利用高斯路径中向量场与分数函数的关系:

uttarget(x)=at∇log⁡pt(x)+btx, u_t^{\text{target}}(x) = a_t \nabla \log p_t(x) + b_t x, uttarget(x)=atlogpt(x)+btx,

其中 at,bta_t, b_tat,bt 由噪声调度器决定,我们可以将上述 SDE 重写为仅依赖分数的形式:

dXt=[(at+σt22)∇log⁡pt(Xt)+btXt]dt+σtdWt. dX_t = \left[ \left(a_t + \frac{\sigma_t^2}{2}\right) \nabla \log p_t(X_t) + b_t X_t \right] dt + \sigma_t dW_t. dXt=[(at+2σt2)logpt(Xt)+btXt]dt+σtdWt.

这意味着,我们可以直接使用训练好的分数网络 stθs_t^\thetastθ 来模拟 SDE,生成样本。

关于 Langevin 动力学
如果取 ut=0u_t = 0ut=0σt\sigma_tσt 为常数,上述 SDE 变为经典的 Langevin 动力学:
dXt=σ22∇log⁡p(Xt)dt+σdWt, dX_t = \frac{\sigma^2}{2} \nabla \log p(X_t) dt + \sigma dW_t, dXt=2σ2logp(Xt)dt+σdWt,
它会在平衡时收敛到 ppp。这是 MCMC 采样的基础,也是扩散模型能够生成多样样本的根源。

(p.s. 当研究对象的质量趋近于0阻尼系数趋近于无穷大时,朗之万动力学就退化为我们熟悉的布朗运动。)


4. 条件生成:如何控制你想要的输出

在实际应用中,我们往往希望生成特定类别的图像或遵循文本提示,即从条件分布 pdata(x∣y)p_{\text{data}}(x|y)pdata(xy) 中采样。最简单的方法(vanilla guidance)是将条件信息 yyy 直接作为输入传入网络,即训练一个条件分数网络 stθ(x∣y)s_t^\theta(x|y)stθ(xy),使得

stθ(x∣y)≈∇log⁡pt(x∣y). s_t^\theta(x|y) \approx \nabla \log p_t(x|y). stθ(xy)logpt(xy).

训练时,只需要将条件 yyy 作为额外输入,其他与无条件训练相同。

然而,单纯的条件训练往往导致模型对条件信息的“关注度”不足,生成结果与提示的匹配度不高。为了强制模型更严格地遵循条件,研究者提出了Classifier-Free Guidance(CFG),它成为当前所有大型生成模型(如 Stable Diffusion 3、DALL·E 3)的标配。

4.1 从分类器引导到无分类器引导

回顾条件分数的贝叶斯分解:

∇log⁡pt(x∣y)=∇log⁡pt(x)+∇log⁡pt(y∣x). \nabla \log p_t(x|y) = \nabla \log p_t(x) + \nabla \log p_t(y|x). logpt(xy)=logpt(x)+logpt(yx).

第二项 ∇log⁡pt(y∣x)\nabla \log p_t(y|x)logpt(yx) 相当于一个“分类器”的梯度,指导生成向满足条件的方向移动。如果我们希望加强这一指导,可以引入一个缩放因子 w>1w > 1w>1

u~t(x∣y)=uttarget(x)+w⋅at∇log⁡pt(y∣x). \tilde{u}_t(x|y) = u_t^{\text{target}}(x) + w \cdot a_t \nabla \log p_t(y|x). u~t(xy)=uttarget(x)+watlogpt(yx).

这就是分类器引导(Classifier Guidance)。但训练一个额外的分类器 pt(y∣x)p_t(y|x)pt(yx) 并不容易,尤其是对于文本条件。

CFG 巧妙地将上述操作转化为无条件与条件模型的线性组合

u~t(x∣y)=(1−w)uttarget(x∣∅)+w uttarget(x∣y), \tilde{u}_t(x|y) = (1-w) u_t^{\text{target}}(x|\emptyset) + w \, u_t^{\text{target}}(x|y), u~t(xy)=(1w)uttarget(x∣∅)+wuttarget(xy),

其中 uttarget(x∣∅)u_t^{\text{target}}(x|\emptyset)uttarget(x∣∅) 是无条件模型(用空标签训练得到),uttarget(x∣y)u_t^{\text{target}}(x|y)uttarget(xy) 是条件模型。通过调整权重 www,我们可以控制条件影响的程度。当 w=1w = 1w=1 时,退化为标准条件生成;当 w>1w > 1w>1 时,相当于放大了条件梯度,从而产生更贴合提示的结果。

训练技巧:为了用一个网络同时表示无条件与条件模型,我们在训练时随机以概率 η\etaη 将条件 yyy 替换为空标签 ∅\emptyset,这样网络就学会了同时处理有条件和无条件输入。推理时,我们分别计算 utθ(x∣∅)u_t^\theta(x|\emptyset)utθ(x∣∅)utθ(x∣y)u_t^\theta(x|y)utθ(xy),然后按上述公式组合。

4.2 实际效果

CFG 能显著提升生成内容与文本提示的对齐度。例如,在文本到图像的生成中,使用 CFG(w>1w > 1w>1)可以让生成的图像更准确地包含提示中的细节,如“一只戴着红帽子的白色小狗”等复杂组合。


5. 小结:从分数到引导的完整路径

本文我们走完了另一条通往生成模型的道路:

  1. 分数函数 提供了概率密度的“方向信息”,可通过 Denoising Score Matching 从数据中学习。
  2. 利用分数函数,我们可以构造 SDE 采样器,实现灵活的生成过程(包括 Langevin 动力学)。
  3. 为了条件生成,我们可以直接训练条件分数网络,但 Classifier-Free Guidance 通过线性组合无条件和条件模型,大幅提升了生成与提示的对齐度。

Flow Matching 和 Score Matching 本质上是两种等价的理论框架——一个学习“速度”,一个学习“梯度”。在下一篇文章中,我们将把理论付诸实践,看看如何利用这些技术构建大规模图像与视频生成模型,包括如何设计神经网络架构、在潜空间训练,以及剖析像 Stable Diffusion 3 和 Meta Movie Gen 这样的实际系统。

如果你对如何将抽象的数学变成真实的像素感到好奇,敬请期待第四篇!

Logo

AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。

更多推荐