前言

本篇接上篇扩散模型系列--NCSN原理解析、示例、代码

与其他博客主要区别在于:

  • 详细的公式推导 尝试将公式详细拆分推导,而不仅限于使用公式。包括其他博客都会跳过的逆向SDE漂移项和分数函数推导。
  • 举例说明 用具体的数值演化,模拟SDE训练和生成过程,帮助理解。

1.SDE简介

SDE是NCSN 后来的改进版本。基于随机微分方程(SDE, Stochastic Differential Equation)的分数生成模型,这个方法进一步推广了 NCSN,使得加噪和去噪过程更加连续化,从而提高了生成质量和稳定性。

1.1 为什么要引入SDE

NCSN 采用的是 分层噪声建模,即在多个固定的噪声水平 \sigma_t上训练神经网络 s_{\theta}(x_t, \sigma_t) 来预测分数函数 \nabla_{x_t} \log p(x_t)
但这样做有两个问题:

  1. 离散噪声水平:NCSN 需要手动选择一组固定的噪声水平,噪声之间的平滑性无法保证,训练和生成可能会受到影响。
  2. 采样过程有限制:在 NCSN 里,我们用 Langevin Dynamics 进行去噪生成,但它本质上是一个马尔科夫链采样方法,可能会收敛得较慢。

因此,基于 SDE 的方法 被提出,它提供了一个更自然的 连续加噪 & 去噪过程

1.2 SDE 版本的核心思想

SDE 版本的分数生成模型使用 一个连续的随机微分方程 来表示数据加噪过程:

dx = f(x, t) dt + g(t)

其中:

  • f(x, t) 是漂移项(drift term),决定数据如何随时间演化。
  • g(t) 是扩散项(diffusion term),决定噪声的大小。
  • dW 是标准布朗运动(Wiener 过程),引入随机噪声。

在这个框架下:

  • 加噪过程:数据 x_0​ 经过一个 SDE 的演化,逐渐变成噪声分布(如高斯分布)。
  • 去噪(生成)过程:反转 SDE,即沿着 SDE 逆向求解,恢复数据分布。

关键点:我们仍然需要一个神经网络 s_{\theta}(x, t) 来估计分数函数 \nabla_{x} \log p(x_t),但现在训练时的噪声水平 t 是连续变化的,而不是像 NCSN 那样用一个离散集合 \{\sigma_i\} 采样。

1.3 具体的加噪和去噪过程

我们用 正向 SDE 来描述加噪:

dx = f(x, t) dt + g(t) dW

常见的两种加噪 SDE:

  • VP-SDE(Variance Preserving):它的扩散过程与 DDPM(扩散模型)类似:

dx = -\frac{1}{2} \beta(t) x dt + \sqrt{\beta(t)} dW

其中 \beta(t) 控制噪声随时间的增加。(之后主要以此为例)

  • VE-SDE(Variance Exploding):它的噪声随时间增加更快:

dx = \sqrt{d\sigma^2(t)/dt} dW

这里的 \sigma(t) 直接决定噪声水平。

在训练时,我们仍然学习 分数函数

s_{\theta}(x, t) \approx \nabla_x \log p(x_t)

但这次 t连续的,使得模型可以在整个时间尺度上有效预测梯度。

去噪(生成)阶段:用 反向 SDE 进行采样:

dx = \left[f(x, t) - g(t)^2 \nabla_x \log p(x_t) \right] dt + g(t) d\bar{W}

其中 d\bar{W} 是一个反向 Wiener 过程。

我们通过数值求解这个逆 SDE,从纯噪声 x_T 逐步得到高质量的样本。

1.4 关键改进点

相较于 NCSN,SDE 版本的改进点有:

  • 连续噪声尺度 t:避免了 NCSN 需要预设噪声水平的限制,提高了训练和生成的平滑性。
  • 更稳定的生成过程:Langevin Dynamics 需要手动设定步长,而 SDE 的数值解法更加自然,效果更稳定。
  • 更好的多尺度建模:可以使用不同的 SDE 形式(如 VP-SDE、VE-SDE)来适配不同的数据特性,提高生成质量。
SDE(Score-Based SDE)NCSN(Score Matching)
分数定义s_{\theta}(x, t) \approx \nabla_x \log p_{t}(x)s_{\theta}(x, \sigma ) \approx \nabla_x \log p_{\sigma }(x)
噪声建模显式使用 SDE 正向过程直接用高斯噪声加到数据上
训练目标基于贝叶斯估计的去噪方法直接用 分数匹配 方法
时间演化x_{t}​ 由 SDE 生成,每个时间点都有 p_t(x)x_{\sigma }​ 由高斯核平滑生成,分布取决于 \sigma

2. 详细的公式推导

2.0 随机微分方程

对这块不熟悉请查看基于分数的生成模型(Score-based generative models)4.2.1. 微分方程 和 4.2.2. 随机微分方程 ,该博主的博客这两部分的解释足够阅读接下来的内容。如果基础薄弱也没关系,这一章会不吝重复的详细拆解

2.1 从最基本的加噪过程推导 SDE

回顾 扩散模型(DDPM),它的加噪过程是:

x_t = \sqrt{\alpha_t} x_0 + \sqrt{1 - \alpha_t} \epsilon

其中:

  • x_0​ 是原始数据。
  • x_t​ 是在时间 t 时刻的加噪数据。
  • \epsilon \sim \mathcal{N}(0, I) 是高斯噪声。

这表示我们用 离散的步长 加噪,但如果我们改用 连续时间的加噪方式 呢?

2.1.1 用连续时间描述加噪

我们希望找一个微分方程,使得 数据点 x 在时间 t 逐步演化到一个高斯分布,即:

数据演化:\quad x_0 \to x_t \to x_T

其中:

  • x_0​ 是原始数据分布 p(x)
  • x_T​ 是一个 标准高斯分布 \mathcal{N}(0, I)(完全加噪)。
  • x_t​ 是中间的加噪状态。

如果 x 经过一个 连续的时间演化过程 变成高斯分布,我们可以把它建模为一个 随机微分方程(SDE)

dx = f(x, t) dt + g(t) dW

这就是我们要推导的公式!那么为什么能写成这样呢?

2.1.2 SDE 的物理意义

公式:

dx = f(x, t) dt + g(t) dW

成分解析:

  • dx:表示数据 x 在时间 t 的变化。
  • dt:一个无穷小的时间间隔。
  • f(x, t) dt:一个 确定性 漂移项,表示数据随时间的平稳变化趋势(比如,数据慢慢向零点收缩)。
  • g(t) dW:一个 随机项,其中:
    • dW 是一个 标准布朗运动(Wiener 过程)。
    • g(t) 控制噪声的大小。

这个方程的物理意义:

  1. 确定性部分(f(x, t) dt
    • 让数据逐渐收缩到一个均值,比如 f(x, t) = -\frac{1}{2} \beta(t) x,表示数据朝着 0 方向收缩。
  2. 随机部分(g(t) dW
    • 让数据被随机扰动,使得数据慢慢变成一个高斯分布。

这样,我们就得到了一个 连续的加噪过程,它可以把数据 x_0​ 逐渐变成一个高斯分布 x_T​。

2.2 不同类型的SDE

不同的 SDE 形式可以决定不同的加噪方式,以下是两种常见的 SDE:

2.2.1 VP-SDE(Variance Preserving SDE)

dx = -\frac{1}{2} \beta(t) x dt + \sqrt{\beta(t)} dW

特点:

  • 漂移项 (x, t) = -\frac{1}{2} \beta(t) xx 指数衰减,趋近于 0。
  • 扩散项 g(t) = \sqrt{\beta(t)} 增加噪声,使数据模糊化,让数据逐渐变成一个高斯分布。

这个形式和 扩散模型(DDPM) 很像!

这个方程描述的是一个 均值衰减 + 添加噪声 的过程。

你可以把它理解成:

  • 在每个时间步 dt 内,我们都 往数据里加一点随机噪声,同时 数据还会被拉回原点
  • 最终,经过足够长的时间 T,数据就会完全变成一个 标准高斯分布

2.2.2 VE-SDE(Variance Exploding SDE)

dx = \sqrt{\frac{d\sigma^2(t)}{dt}} dW

特点:

  • 没有漂移项,只有一个扩散项。
  • 噪声会快速增加,适用于数据分布扩散较快的情况。

2.2.3 不同类型 SDE 的参数设置

不同类型的 SDE 会有不同的 \beta(t) 设定,例如:


(1)VP-SDE(Variance Preserving SDE)

  • \beta(t) 通常是一个线性或指数增长的函数,如:

\beta(t) = \beta_{\text{min}} + (\beta_{\text{max}} - \beta_{\text{min}}) t

  • 这里,\beta_{\text{min}}\beta_{\text{max}}​ 是超参数,控制最小和最大的噪声强度。

(2)VE-SDE(Variance Exploding SDE)

  • 这里 \beta(t) 通常采用指数增长形式,例如:

\beta(t) = \sigma_{\text{min}}^2 (\frac{\sigma_{\text{max}}}{\sigma_{\text{min}}})^{2t}

  • 这种形式使得数据在高噪声水平下变化更剧烈。

2.3 反向SDE进行去噪

现在,我们已经知道如何用 SDE 进行 加噪,那么如何 去噪 呢?

我们可以 反转 SDE,也就是沿着 SDE 逆向求解

dx = \left[f(x, t) - g(t)^2 \nabla_x \log p_t(x) \right] dt + g(t) d\bar{W}

其中:

  • \nabla_x \log p_t(x) 是我们训练神经网络要估计的 分数函数,也就是数据分布的梯度。
  • d\bar{W} 是一个新的 Wiener 过程。

通过数值求解这个逆 SDE,我们可以从纯噪声 x_T​ 逐渐恢复到数据 x_0​。

在训练时,我们用神经网络 s_{\theta}(x_t, t) 来逼近 \nabla_x \log p_t(x),即:

s_{\theta}(x_t, t) \approx \nabla_x \log p_t(x)

从而使得:

- g^2(t) s_{\theta}(x_t, t) \approx - g^2(t) \nabla_x \log p_t(x)

这里在训练时的s_\theta(x_t, t)计算与NSCN类似:

s_\theta(x_t, t) \approx \frac{x_t - x_0}{g^2(t)}

这就意味着,我们训练的目标是 让神经网络 s_{\theta}(x_t, t) 预测出数据在噪声污染状态下的最优调整方向,这样在反向去噪的过程中,我们可以使用这个神经网络来指导数据回归到真实分布。

换句话说:

训练的目标是最小化分数匹配损失,即:

\mathbb{E}_{p_t(x)} \left[ \| s_{\theta}(x_t, t) - \nabla_x \log p_t(x) \|^2 \right]

这样就能让神经网络学习到最优的分数函数估计。

该小节接下来部分将推导这个SDE逆向求解公式,这是其他博客中没有的。如果不需要了解,可以跳过。

2.3.1 不同的推导背景和时间尺度下的分数函数

与NCSN不同,这个SDE 逆向求解公式明显不是通过SDE公式简单移项获得的,到底时怎么推导出来的?在此之前,让我们先分析下s_{\theta}(x_t, t)


第一种表达式:

s_\theta(x_t, t) \approx \frac{x_t - x_{t-1} - f(x_{t-1}, t-1) \Delta t}{g^2(t)\Delta t}

适用场景:离散时间 SDE 近似,基于局部增量推导

  • 这一表达式基于相邻时间步 t-1t 之间的差分估计分数函数。
  • 从正向 SDE 角度来看,前向过程的增量可以写为:

x_t = x_{t-1} + f(x_{t-1}, t-1) \Delta t + g(t) \sqrt{\Delta t} \epsilon

其中 \epsilon \sim \mathcal{N}(0, I) 是高斯噪声。

  • 如果希望估计分数函数(即后验梯度):

s_\theta(x_t, t) = \nabla_{x_t} \log p_t(x_t)

  • 可以利用相邻时间步的变化率来估计,得到:

s_\theta(x_t, t) \approx \frac{x_t - x_{t-1} - f(x_{t-1}, t-1) \Delta t}{g^2(t)\Delta t}

  • 直观理解:它描述了局部时间增量如何影响噪声项的估计,是一种离散化近似

第二种表达式:

s_\theta(x_t, t) \approx \frac{x_t - x_0}{g^2(t)}

适用场景:全局轨迹建模,基于初始分布的推导

  • 这个表达式通常用于全局时间尺度的分数估计,特别是在扩散模型(如 DDPM)或SDE 解析解的情况下。
  • 从连续 SDE 角度来看,如果我们假设 x_t 可以从初始状态 x_0 逐步演化而来,那么其均值变化可以近似为:

x_t = x_0 + \int_0^t f(x_s, s) ds + \int_0^t g(s) dW_s

        在简单情形下(如线性扩散模型),可以得到:

s_\theta(x_t, t) \approx \frac{x_t - x_0}{g^2(t)}

  • 直观理解:它描述了全局时间尺度上,数据如何偏离初始状态 x_0​,从而估计其后验分布的梯度。

两者的关系

  • 第一种表达式 适用于离散时间建模,尤其是基于相邻时间步的局部梯度估计。
  • 第二种表达式 适用于连续时间建模,尤其是基于全局时间积分的推导。
  • 当时间步足够小\Delta t \to 0),第一种表达式可以在期望意义上收敛到第二种表达式的形式。

总结

这两种表达式本质上是 时间尺度不同 的近似,没有矛盾

  • 局部增量推导(基于相邻时间步)s_\theta(x_t, t) \approx \frac{x_t - x_{t-1} - f(x_{t-1}, t-1) \Delta t}{g^2(t)\Delta t}
  • 全局轨迹建模(基于初始状态)s_\theta(x_t, t) \approx \frac{x_t - x_0}{g^2(t)}

你可以根据具体的推导背景选择合适的公式,比如:

  • 训练过程中,基于 SDE 离散化的损失函数,通常使用局部增量公式(第一种)。
  • 解析求解 SDE,或理解扩散过程整体行为,通常使用全局公式(第二种)。

如果你的问题涉及的是SDE 逆向推导中的分数函数训练,那么第一种公式更常见。如果你要理解扩散模型的整体生成行为,第二种公式更直观。


到这里可能发现,不但SDE逆向求解公式没弄明白,还不知道s_{\theta}(x_t, t)公式从哪推导出来的,没关系,接着推导!很快就能拨云见日!

2.3.2 为什么正向和逆向过程不能直接通过移项建立联系?

在扩散模型中,正向过程的加噪可以表示为:

x_{t+1} = x_t + z

去噪过程是:

x_t = x_{t+1} - z

可以直接通过移项来表示一个去噪的逆过程。真实扩散模型中的噪声增长是有时间相关性的(受 \sigma_t 控制),这里简写。

但在SDE中,正向和逆向的关系是 通过Fokker-Planck方程(FP方程)建立的,即:

  1. 正向SDE控制的是数据分布的演化,描述了数据如何被加噪成噪声分布。
  2. 逆向SDE是基于FP方程得到的,它并不是简单的数学移项,而是通过对时间反转和分数函数(score function)的利用推导出来的。

直观理解:

  • 在正向SDE中,数据在 漂移项扩散项 的作用下演化,逐渐变得像噪声分布。
  • 逆向SDE不是简单地把正向SDE的等式变换一下,而是通过 学习数据分布的梯度 来引导数据回归真实分布。

因此,虽然表面上它们的形式相似,但逆向SDE的关键在于额外的 梯度修正项g^2(t) \nabla_x \log p_t(x),这个项是通过分数函数学习得到的,而不是直接从正向SDE的表达式移项得到的。

2.3.3 逆向SDE公式

我们首先写出正向 SDE:

dx = f(x, t) dt + g(t) dW

这个公式描述了数据在时间 t 的演化方式。我们现在想要找到逆向过程,即:

dx = \tilde{f}(x,t) dt + g(t) dW

其中 \tilde{f}(x,t) 是逆向 SDE 的漂移项。

已有的研究表明,逆向 SDE 的漂移项 \tilde{f}(x,t) 与正向漂移项 f(x,t) 之间的关系是:

\tilde{f}(x, t) = f(x, t) - g^2(t) \nabla_x \log p_t(x)

其中,\nabla_x \log p_t(x) 就是我们要学习的分数函数。

这个公式的直觉是:

  • f(x,t) 描述的是数据如何向前传播
  • \nabla_x \log p_t(x) 是数据分布的梯度(即分数)。
  • - g^2(t) \nabla_x \log p_t(x) 是我们在逆向过程中应该抵消的项,确保数据在逆向过程中遵循正确的轨迹。

这个公式的推导涉及到 福克-普朗克方程时间反演第一个重头戏在2.4节展开推演。

2.4 逆向 SDE 的漂移项与正向漂移项关系式推导

参考论文链接:

宋飏论文:Score-Based Generative Modeling through Stochastic Differential Equations

Anderson (1982) 时间反转SDE的经典理论:Reverse-time diffusion equation models

2.4.1 问题设定

我们有一个前向 SDE

dx = f(x,t) dt + g(t) dW_t

其中:

  • f(x,t) 是漂移项(drift term)。
  • g(t) 是扩散系数(diffusion coefficient)。
  • dW_t​ 是布朗运动(Wiener process)。

我们想知道如果逆向时间(从 T 逆推到 0),数据会如何演化?


2.4.2 轨迹时间反演的定义

在随机过程里,时间反演指的是我们观察一个过程 X_t​,但以逆向时间 s = T - t 来看它,得到新的过程:

Y_s = X_{T-s}

我们的目标是找到Y_s 满足的 SDE,即找到逆向 SDE 的漂移项。

我们需要解决两个核心问题:

  1. 如何求出逆过程的漂移项?
  2. 如何与福克-普朗克方程(Fokker-Planck)联系起来?

2.4.3 正向 SDE 的概率流方程

SDE 的轨迹分布由福克-普朗克方程控制:

\frac{\partial p_t(x)}{\partial t} = - \nabla_x \cdot (f(x,t) p_t(x)) + \frac{1}{2} g^2(t) \nabla_x^2 p_t(x)

这里:

  • 第一项 -\nabla_x \cdot (f p_t)漂移项,描述了 f(x,t) 对分布 p_t(x) 的影响。
  • 第二项 \frac{1}{2} g^2 \nabla_x^2 p_t(x)扩散项,由噪声 g(t) dW_t 引入。

这个方程描述了数据点的分布如何随时间演化。

但我们要找的是单个样本点的逆向轨迹如何变化。

说直白点,为了找到逆向 SDE,我们需要找到逆过程的漂移项 \tilde{f}(x,t)


2.4.4 逆过程的漂移项

根据 Anderson (1982)Haussmann & Pardoux (1986),逆向 SDE 具有以下形式:

dx = \tilde{f}(x,t) dt + g(t) d\tilde{W}_t

其中:

  • \tilde{W}_t 是新的布朗运动,方向相反
  • \tilde{f}(x,t) 是新的漂移项,我们需要推导其表达式。

时间反演的关键是找到当我们从时间 T 开始向 t=0 逆向演化时,数据如何变化。关键定理(Anderson 1982)指出,逆向 SDE 的漂移项\tilde{f}(x,t) 满足::

\tilde{f}(x,t) = f(x,t) - g^2(t) \nabla_x \log p_t(x)

这表示:

  • 逆向漂移项不仅由原来的漂移项 f(x,t) 决定,还需要一个额外的校正-g^2(t) \nabla_x \log p_t(x),它与数据分布的梯度 \nabla_x \log p_t(x) 相关。
  • \nabla_x \log p_t(x) 被称为 分数函数(Score Function),表示在时刻 t 处的数据分布梯度
  • 额外的项 -g^2(t) \nabla_x \log p_t(x)  调整了原始的正向漂移项 f(x,t),确保逆向过程能够正确地还原数据。

这个额外项的来源是 Ito 公式+福克-普朗克方程+时间反演 的结合,我们继续推导。


2.4.5 详细推导

首先,我们考虑正向 SDE 生成的概率流:

J_t(x) = f(x,t) p_t(x) - \frac{1}{2} g^2(t) \nabla_x p_t(x)

这个 J_t(x) 叫做 概率流(Probability Current),它衡量了样本在分布 p_t(x) 内的流动情况。

在时间反演下,概率流需要满足:

\tilde{J}_t(x) = -J_t(x)

这意味着在逆向过程中,数据应该逆向流动

我们代入 J_t(x) 的定义:

-\big( f(x,t) p_t(x) - \frac{1}{2} g^2(t) \nabla_x p_t(x) \big) = \tilde{f}(x,t) p_t(x) - \frac{1}{2} g^2(t) \nabla_x p_t(x)

整理得:

\tilde{f}(x,t) p_t(x) = f(x,t) p_t(x) - g^2(t) \nabla_x p_t(x)

两边同时除以 p_t(x)(假设 p_t(x) > 0):

\tilde{f}(x,t) = f(x,t) - g^2(t) \nabla_x \log p_t(x)

这就是逆向 SDE 的漂移项公式!


如果有认真推导,可能发现整理的式子的右边存在符号问题。实际整理的结果应该是

\tilde{f}(x,t) = - f(x,t) + g^2(t) \nabla_x \log p_t(x)

这里的关键是漂移项 f(x,t) 的方向反演

在时间反演的过程中,我们实际上在推导 反演后的系统,其中时间是倒流的,所以 f(x,t) 这个项本身也要进行变换。具体来说:

  • 正向时间的漂移项是 f(x,t)
  • 反演后,我们得到的公式的符号是 - f(x,t)
  • 但我们通常会改写公式,使得漂移项仍然保持正向的解释(即,写成 f(x,t) 形式,而不是 - f(x,t))。
  • 这样,我们通过调整符号,把 - f(x,t) 重新解释成 f(x,t) 的方向,而修正项 g^2(t) \nabla_x \log p_t(x) 的符号相应调整。

所以最终我们写作:

\tilde{f}(x,t) = f(x,t) - g^2(t) \nabla_x \log p_t(x)

这一步是一个数学上的等价变换,同时也使得公式更加符合直觉。

这里可以跳转 2.5 时间反演的离散化步骤方向 ,单开了一节分析。


2.4.6 直观解释

  • f(x,t) 代表数据的正向流动趋势
  • \nabla_x \log p_t(x) 代表数据分布的梯度,它指向数据密度更高的方向。
  • 乘上 -g^2(t) 作为修正项,让数据朝着更符合真实分布的方向回溯。

直观上:

  • 如果 x_t 是一个噪声数据,逆向过程中我们希望它慢慢向真实数据靠近。
  • 这个校正项 -g^2(t) \nabla_x \log p_t(x) 就像是“引导力”,使得数据回到它原本应该在的地方。

2.4.7 结论

逆向 SDE 为:

dx = (f(x,t) - g^2(t) \nabla_x \log p_t(x)) dt + g(t) d\tilde{W}_t

  • 第一项 f(x,t) - g^2(t) \nabla_x \log p_t(x) 调整了数据流动方向,使其更符合数据分布。
  • 第二项 g(t) d\tilde{W}_t​ 仍然是随机噪声,但在逆向过程中其分布有所调整。

这个推导结合了:

  1. 福克-普朗克方程(描述分布演化)
  2. 概率流守恒(保证数据逆向合理)
  3. 时间反演理论(推导出修正项)

最终,我们得到了这个关键公式,并用于扩散模型、SDE-GAN、NCSN生成式建模中。

2.5 时间反演的离散化步骤方向

分析结论写在前面:

先不考虑扩散项等问题,简化思路,可以把SDE正向过程看作x_{t+1}=x_t + f(x,t)

那么从时间反演理论或者福克普朗克方程的角度考虑的逆向过程是x_{t}=x_{t+1} + \tilde{f}(x,t)

在这种情况下\tilde{f}(x,t) = g^2(t) \nabla_x \log p_t(x) - f(x,t)称之为逆向漂移项。

而宋飏的论文里正向过程同上,在逆向操作的时候,应该是x_{t}=x_{t+1} - \tilde{f}(x,t)

如果分析没错,就意味着宋飏论文里的逆向漂移项应该再求个反,也就对上了!

2.5.1 离散化视角下的正向与逆向过程

  • 正向过程离散化(显式欧拉法):

    x_{t+1} = x_t + f(x_t,t)\Delta t(忽略扩散项)
  • 这对应于连续SDE:

    dx=f(x,t)dt
  • 逆向过程离散化(时间反演):
    若想从 x_{t+1}​ 反推 x_t,理论上应写作:

    x_{t}=x_{t+1} - \tilde{f}(x,t)\Delta t

    这里的关键是 逆向步的漂移项符号与正向相反,即:

    \tilde{f}(x,t)\Delta t = -f(x,t)\Delta t(若仅考虑漂移项反转)

2.5.2 加入扩散项后的修正

当考虑扩散项 g(t) 时,逆向过程需额外补偿概率流。此时逆向漂移项应修正为:

\tilde{f}(x,t)=-f(x,t)+g^2(t)\bigtriangledown_{x}logp_{t}(x)​​

后面为扩散补偿项

然而,宋飏的论文中符号与该推导相反,这是因为:

  • 时间反演的连续视角
    在连续时间下,时间反演 \tau =T-t 会导致时间导数符号变化:

    \frac{\partial}{\partial \tau } = -\frac{\partial}{\partial t }

    这会进一步影响FPE中的对流项(漂移项)符号,最终推导出的逆向漂移项为:

    \tilde{f}(x,t)=f(x,t)-g^2(t)\bigtriangledown_{x}logp_{t}(x)

2.5.3 符号矛盾的根源

该推导与宋飏论文的差异源于以下两点:

离散化方向的定义

  • 该推导假设逆向过程为 x_t = x_{t+1}+\tilde{f}\Delta t(显式加号),而宋飏的连续时间推导隐含了逆向步的负号调整(dx=\tilde{f}dt 对应 x_t = x_{t+1}-\tilde{f}\Delta t)。

概率梯度项的补偿方向

  • 扩散项 g^2(t) \nabla_x \log p_t(x) 的符号由FPE的匹配条件决定,需抵消正向扩散的概率流,与漂移项符号无关。


2.5.4 结论

推导逻辑是正确的,但符号差异源于:

  • 离散化步骤的方向定义:显式(加号) vs 隐式(减号)。

  • 宋飏论文的约定:直接使用连续时间推导,漂移项符号已包含时间反演调整。

修正建议

  • 若离散化步骤定义为 x_t = x_{t+1}-\tilde{f}\Delta t,则逆向漂移项应为:

    \tilde{f}(x,t)=f(x,t)-g^2(t)\bigtriangledown_{x}logp_{t}(x)
  • 若定义为 x_t = x_{t+1}+\tilde{f}\Delta t,则需取 \tilde{f}(x,t)=-f(x,t)+g^2(t)\bigtriangledown_{x}logp_{t}(x),但此约定与宋飏论文不一致。


2.5.5 扩散项的符号不变性

顺便提一下扩散项的符号问题。

正向的SDE离散化形式可以写成:

x_{t+1} = x_t + f(x_t,t) \Delta t + g(t) \sqrt{\Delta t} \cdot \epsilon, \quad \epsilon \sim \mathcal{N}(0,1)

在推导逆向SDE时,可能会推导出离散化形式:

x_t = x_{t+1} - [ f(x_{t+1}, t+1) - g^2(t+1) \nabla_x \log p_{t+1}(x_{t+1}) ] dt - g(t+1) \epsilon

扩散项中的噪声 \tilde{\epsilon }是独立采样的标准正态随机变量,其分布满足对称性:\tilde{\epsilon }~-\tilde{\epsilon }。因此,无论是 +g(t)\sqrt{\Delta t}\tilde{\epsilon } 还是 -g(t)\sqrt{\Delta t}\tilde{\epsilon },其统计性质完全相同。符号的差异会被噪声的对称性吸收,对实际采样结果没有影响。

因此在逆向SDE的离散化步骤中,扩散项(噪声项)的符号不需要额外添加负号,即正确形式为:

x_t = x_{t+1} - [ f(x_{t+1}, t+1) - g^2(t+1) \nabla_x \log p_{t+1}(x_{t+1}) ] dt + g(t+1) \epsilon

2.6 分数函数推导

接下来就只剩一个问题,即分数函数s_{\theta}(x_t, t)如何表达,为什么这样表达?

第二个重头戏开始了。

2.6.1 用条件期望表达分数函数

我们现在的目标是找到:

\nabla_x \log p_t(x_t)

它衡量了在时间 t 下,数据分布如何变化

首先,考虑正向 SDE 的离散化:

x_{t+1} = x_t + f(x_t, t) \Delta t + g(t) \sqrt{\Delta t} \cdot \epsilon, \quad \epsilon \sim \mathcal{N}(0,1)

如果我们想反推 x_t 的分布,我们可以写出贝叶斯反向更新:

p(x_t | x_{t+1}) \approx \mathcal{N}(x_t; \mu_t, \Sigma_t)

其中:

\mu_t = x_{t+1} - f(x_{t+1}, t+1) \Delta t

\Sigma_t = g^2(t+1) \Delta t

由于分数函数是:

\nabla_x \log p_t(x_t) = \mathbb{E}_{x_{t-1} | x_t} \left[ \nabla_x \log p(x_{t-1} | x_t) \right]

而我们知道 p(x_{t-1} | x_t) 也是高斯分布,因此:

\nabla_x \log p_t(x_t) \approx \mathbb{E}_{x_{t-1} | x_t} \left[ \frac{x_t - x_{t-1} - f(x_{t-1}, t-1) \Delta t}{g^2(t)\Delta t} \right]

这个公式的参数可以理解为:

  • x_t - x_{t-1} 是观测到的变化。
  • f(x_{t-1}, t-1) \Delta t 是理论上应该有的漂移部分。
  • g^2(t) 是扩散项,决定了噪声强度。

这个公式的意思是:SDE 逆向过程中,我们无法直接观察到 x_{t-1},但可以通过贝叶斯估计推测它,从而得到分数函数的近似值。

那这个贝叶斯公式是如何来的,均值和方差又是如何得到的?我们接着推导。

2.6.2 贝叶斯反向更新推导

先解决贝叶斯公式的问题。


正向 SDE 的离散化

我们先考虑一个正向 SDE(假设无漂移项f(x,t),仅考虑扩散项 g(t)):

dx = g(t) dW

离散化后的版本:

x_{t+1} = x_t + g(t) \sqrt{\Delta t} \cdot \epsilon, \quad \epsilon \sim \mathcal{N}(0, I)

这表明在每个时间步,我们给 x_t​ 加上一个高斯噪声。

这个过程是个马尔可夫链,意味着:

p(x_{t+1} | x_t) = \mathcal{N}(x_{t+1} ; x_t, g^2(t) \Delta t \cdot I)

即,x_{t+1}​ 在已知 x_t​ 时服从均值为 x_t ,方差为 g^2(t) \Delta t 的正态分布。


计算逆向分布 

我们想知道,如果已知 x_{t+1}​,那么 x_t​ 的分布 p(x_t | x_{t+1}) (逆向分布)是什么?

根据贝叶斯定理:

p(x_t | x_{t+1}) = \frac{p(x_{t+1} | x_t) p(x_t)}{p(x_{t+1})}

但计算这个很困难。幸运的是,由于 正向过程是高斯的,逆向过程也应该是高斯的。因此,我们直接推导出其均值和方差。

从正向条件分布:

p(x_{t+1} | x_t) = \mathcal{N}(x_{t+1} ; x_t, g^2(t) \Delta t \cdot I)

由于高斯分布的逆向公式,我们可以写出:

p(x_t | x_{t+1}) = \mathcal{N}(x_t ; \mu_t, \Sigma_t)

其中:

\mu_t = x_{t+1} - g^2(t) \nabla_x \log p_t(x_{t+1}) \Delta t

\Sigma_t = g^2(t) \Delta t \cdot I

这个公式的直觉是:

  • x_t 的均值是 x_{t+1}​ 减去一个项 g^2(t) \nabla_x \log p_t(x_{t+1}),这个项纠正了 SDE 的噪声效应。
  • 方差仍然是 g^2(t) \Delta t,即和正向过程一致。

这就是贝叶斯反向更新的来源!


计算分数项

我们需要估计:

\nabla_x \log p_t(x_t)

使用贝叶斯定理,我们知道:

\nabla_x \log p_t(x_t) = \mathbb{E}_{x_{t-1} | x_t} \left[ \nabla_x \log p(x_{t-1} | x_t) \right]

其中 p(x_{t-1} | x_t) 也服从高斯:

p(x_{t-1} | x_t) = \mathcal{N}(x_{t-1} ; \mu_{t-1}, \Sigma_{t-1})

从前面的结果:

\mu_{t-1} = x_t - g^2(t-1) \nabla_x \log p_{t-1}(x_t) \Delta t

所以:

\nabla_x \log p_t(x_t) \approx \frac{x_t - x_{t-1} - f(x_{t-1}, t-1) \Delta t}{g^2(t)\Delta t}

这意味着我们可以用 x_{t-1}​ 的样本来近似计算分数项

均值和方差的来源,在 2.7 高斯分布的逆向更新推导   2.8 条件协方差推导 进行推导。

2.6.3 从均值和方差推导近似的分数函数公式

我们希望推导出近似的 score function 公式,即:

\nabla_x \log p_t(x_t) \approx \frac{x_t - x_{t-1} - f(x_{t-1}, t-1) \Delta t}{g^2(t)\Delta t}

从给定的公式出发:

\mu_t = x_{t+1} - g^2(t) \nabla_x \log p_t(x_{t+1}) \Delta t

\Sigma_t = g^2(t) \Delta t \cdot I


第一步:利用条件期望求近似

我们从正向 SDE 公式出发:

dx = f(x,t) dt + g(t) dW

对其进行离散化,得到:

x_{t+1} = x_t + f(x_t, t) \Delta t + g(t) \sqrt{\Delta t} \cdot \xi, \quad \xi \sim \mathcal{N}(0, I)

这意味着条件期望为:

\mathbb{E}[x_{t+1} \mid x_t] = x_t + f(x_t, t) \Delta t

而条件协方差为:

\text{Var}[x_{t+1} \mid x_t] = g^2(t) \Delta t \cdot I

即:

p(x_{t+1} \mid x_t) = \mathcal{N}(x_{t+1} \mid x_t + f(x_t, t) \Delta t, g^2(t) \Delta t \cdot I)


第二步:利用高斯分布的得分函数

已知高斯分布的 score function(对数密度的梯度):

\nabla_x \log p(x) = -\Sigma^{-1} (x - \mu)

对于条件分布 p(x_t \mid x_{t+1}),我们可以写出:

\nabla_x \log p_t(x_t) = -\Sigma_t^{-1} (x_t - \mu_t)

利用 已给出的 均值公式:

\mu_t = x_{t+1} - g^2(t) \nabla_x \log p_t(x_{t+1}) \Delta t

我们带入方差:

\Sigma_t = g^2(t) \Delta t \cdot I

计算 score function:

\nabla_x \log p_t(x_t) = -\frac{1}{g^2(t) \Delta t} (x_t - (x_{t+1} - g^2(t) \nabla_x \log p_t(x_{t+1}) \Delta t))

展开:

\nabla_x \log p_t(x_t) = -\frac{x_t - x_{t+1} + g^2(t) \nabla_x \log p_t(x_{t+1}) \Delta t}{g^2(t) \Delta t}

小步长极限下,可以用 Euler 近似 近似 x_{t+1}​:

x_{t+1} \approx x_t + f(x_t, t) \Delta t + g(t) \sqrt{\Delta t} \xi

所以:

x_t \approx x_{t-1} + f(x_{t-1}, t-1) \Delta t + g(t) \sqrt{\Delta t} \xi

带入:

\nabla_x \log p_t(x_t) \approx \frac{x_t - x_{t-1} - f(x_{t-1}, t-1) \Delta t}{g^2(t)\Delta t}


结论

我们已经从 条件均值高斯分布的 score function 公式 推导出了目标式子。
这个式子表明 score function 可以用 前一个时刻的状态估计得到,这在扩散模型的训练和采样过程中非常关键!

2.7 高斯分布的逆向更新推导

这实际上是更详细的贝叶斯推导。


2.7.1 设定正向分布

正向扩散过程告诉我们,在已知 x_t 的情况下,x_{t+1}​ 服从高斯分布

p(x_{t+1} | x_t) = \mathcal{N} \left( x_{t+1} \middle| x_t + f(x_t, t) \Delta t, g^2(t) \Delta t I \right)

即:

x_{t+1} \sim \mathcal{N}( x_t + f(x_t, t) \Delta t, g^2(t) \Delta t I)

目标是求出逆向分布

p(x_t | x_{t+1})


2.7.2 设定联合高斯分布

假设我们有一个二维随机向量

\mathbf{X} = \begin{bmatrix} x_t \\ x_{t+1} \end{bmatrix}

已知联合分布 p(x_t, x_{t+1}) 服从二维高斯分布:

\begin{bmatrix} x_t \\ x_{t+1} \end{bmatrix} \sim \mathcal{N} \left( \begin{bmatrix} \mu_t \\ \mu_{t+1} \end{bmatrix}, \begin{bmatrix} \Sigma_{t,t} & \Sigma_{t,t+1} \\ \Sigma_{t+1,t} & \Sigma_{t+1,t+1} \end{bmatrix} \right)

其中:

  • \mu_t\mu_{t+1}​ 是 x_t​ 和 x_{t+1}​ 的均值。
  • \Sigma_{t,t}​ 是 x_t​ 的边际协方差。
  • \Sigma_{t+1,t+1}​ 是 x_{t+1}​ 的边际协方差。
  • \Sigma_{t,t+1} = \Sigma_{t+1,t}x_t​ 和 x_{t+1}​ 之间的协方差(跨时间步协方差)。

我们希望求 条件分布 p(x_t | x_{t+1}) 的均值和方差。


2.7.3 高斯条件分布公式

对于一般的多维高斯分布

\mathbf{X} = \begin{bmatrix} X_1 \\ X_2 \end{bmatrix} \sim \mathcal{N} \left( \begin{bmatrix} \mu_1 \\ \mu_2 \end{bmatrix}, \begin{bmatrix} \Sigma_{11} & \Sigma_{12} \\ \Sigma_{21} & \Sigma_{22} \end{bmatrix} \right)

其中:

  • X_1​ 和 X_2 可能是多个维度的向量。
  • \Sigma_{11}​、\Sigma_{22}​ 是它们各自的协方差矩阵。
  • \Sigma_{12} = \Sigma_{21}^T​ 是它们的交叉协方差。

则高斯分布的条件分布 p(X_1 | X_2) 仍然是高斯分布:

p(X_1 | X_2) = \mathcal{N} \left( \mu_1', \Sigma_1' \right)

其中:

  • 条件均值(后验估计)\mu_1' = \mu_1 + \Sigma_{12} \Sigma_{22}^{-1} (X_2 - \mu_2)
  • 条件协方差\Sigma_1' = \Sigma_{11} - \Sigma_{12} \Sigma_{22}^{-1} \Sigma_{21}

2.7.4 应用于本问题

根据高斯分布的条件分布公式,我们知道:

p(x_t | x_{t+1}) = \mathcal{N}(x_t | \mu_t', \Sigma_t')

其中:

  • 均值(后验估计)

\mu_t' = \mu_t + \Sigma_{t, t+1} \Sigma_{t+1, t+1}^{-1} (x_{t+1} - \mu_{t+1})

  • 协方差

\Sigma_t' = \Sigma_{t, t} - \Sigma_{t, t+1} \Sigma_{t+1, t+1}^{-1} \Sigma_{t+1, t}


2.7.5 应用于SDE逆向过程

在正向过程中:

\mu_{t+1} = x_t + f(x_t, t) \Delta t, \quad \Sigma_{t+1, t+1} = g^2(t) \Delta t I

由于 x_{t+1}​ 仅由 x_t 产生,我们有:

\Sigma_{t, t+1} = g^2(t) \Delta t I

带入高斯逆向公式

\mu_t' = x_t + g^2(t) \Delta t \cdot \frac{1}{g^2(t) \Delta t} (x_{t+1} - \mu_{t+1})

展开:

\mu_t' = x_t + (x_{t+1} - x_t - f(x_t, t) \Delta t)

化简:

\mu_t' = x_{t+1} - f(x_t, t) \Delta t

然而,由于我们不知道 x_t,无法直接计算 f(x_t, t)。我们可以用神经网络 s_\theta(x, t) 估计分数函数

s_\theta(x, t) \approx \nabla_x \log p_t(x)

最终,均值写成:

\mu_t' = x_{t+1} - g^2(t) \nabla_x \log p_t(x_{t+1})

即:

\mu_t' = x_{t+1} - Q \nabla_x \log p_t(x_{t+1})

其中 Q = g^2(t) \Delta t


2.7.6 直观理解

这个公式告诉我们:

  1. 如果我们直接用 x_{t+1}​ 作为 x_t 的估计,会产生偏差
  2. 修正项 Q \nabla_x \log p_t(x_{t+1}) 作用是“往高概率区域推”
    • 如果 x_{t+1}​ 处于低概率区域,修正项大,需要往高概率区推。
    • 如果 x_{t+1} 处于高概率区域,修正项小,不需要太多修正。

这就是为什么逆向过程的均值是:

\mu_t' = x_{t+1} - Q \nabla_x \log p_t(x_{t+1})

    2.8 条件协方差推导

    到此为止其实还是默认了条件协方差近似为g^2(t)\Delta t,这里将进行推导。

    2.8.1 正向SDE的离散化与协方差结构

    考虑线性漂移项的 Ornstein-Uhlenbeck 过程:

    dx = -\beta x dt + g(t) dw

    离散化为时间步长 \Delta t

    x_{t+\Delta t} = x_t (1 - \beta \Delta t) + g(t) \sqrt{\Delta t} \epsilon, \quad \epsilon \sim \mathcal{N}(0,1)

    在稳态下,x_t 的协方差为:

    \Sigma_{t,t} = \frac{g^2(t)}{2\beta}

    这里的Ornstein-Uhlenbeck 过程和之前的VP-SDE的区别在于漂移项和扩散项的系数设计,但二者可通过参数调整相互转换。

    2.8.2 联合协方差矩阵的分块元素

    联合变量 z = \begin{bmatrix} x_t \\ x_{t+\Delta t} \end{bmatrix} 的协方差矩阵分块为:

    \Sigma = \begin{bmatrix} \Sigma_{t,t} & \Sigma_{t,t+\Delta t} \\ \Sigma_{t+\Delta t,t} & \Sigma_{t+\Delta t,t+\Delta t} \end{bmatrix}

    边际协方差 \Sigma_{t+\Delta t,t+\Delta t}

    由于稳态性,

    \Sigma_{t+\Delta t,t+\Delta t} = \Sigma_{t,t} = \frac{g^2(t)}{2\beta}

    跨时间步协方差 \Sigma_{t,t+\Delta t}

    \Sigma_{t,t+\Delta t} = \mathbb{E}[x_t x_{t+\Delta t}] = \mathbb{E} [x_t (x_t (1 - \beta \Delta t) + g(t) \sqrt{\Delta t} \epsilon)]

    由于 x_t\epsilon 独立且 \mathbb{E}[\epsilon] = 0,得:

    \Sigma_{t,t+\Delta t} = \Sigma_{t,t} (1 - \beta \Delta t)

    2.8.3 条件协方差公式代入

    条件协方差公式为:

    \Sigma_{t|t+\Delta t} = \Sigma_{t,t} - \frac{\Sigma_{t,t+\Delta t}^2}{\Sigma_{t+\Delta t,t+\Delta t}}

    代入分块协方差:

    \Sigma_{t|t+\Delta t} = \frac{g^2(t)}{2\beta} - \frac{\left( \frac{g^2(t)}{2\beta} (1 - \beta \Delta t) \right)^2}{\frac{g^2(t)}{2\beta}}

    2.8.4 分子与分母的展开

    展开分子:

    \left( \frac{g^2(t)}{2\beta} (1 - \beta \Delta t) \right)^2 = \left( \frac{g^2(t)}{2\beta} \right)^2 (1 - 2\beta \Delta t + \beta^2 \Delta t^2)

    分母为:

    \frac{g^2(t)}{2\beta}

    因此:

    \Sigma_{t|t+\Delta t} = \frac{g^2(t)}{2\beta} - \frac{\frac{g^4(t)}{4\beta^2} (1 - 2\beta \Delta t + \beta^2 \Delta t^2)}{\frac{g^2(t)}{2\beta}}

    化简:

    \Sigma_{t|t+\Delta t} = \frac{g^2(t)}{2\beta} - \frac{g^2(t)}{2\beta} (1 - 2\beta \Delta t + \beta^2 \Delta t^2)

    2.8.5 保留一阶小量

    展开后:

    \Sigma_{t|t+\Delta t} = \frac{g^2(t)}{2\beta} [1 - (1 - 2\beta \Delta t + \beta^2 \Delta t^2)]

    = \frac{g^2(t)}{2\beta} (2\beta \Delta t - \beta^2 \Delta t^2)

    \Delta t \to 0 时,忽略高阶小项 \beta^2 \Delta t^2

    \Sigma_{t|t+\Delta t} \approx \frac{g^2(t)}{2\beta} \cdot 2\beta \Delta t = g^2(t) \Delta t

    2.8.6 结论

    通过精确展开并保留一阶小量,条件协方差在小时间步长下简化为:

    \Sigma_{t|t+\Delta t} = g^2(t) \Delta t

    推导关键步骤总结:

    1. 线性漂移假设:确保协方差传播可解析计算。
    2. 稳态协方差代入:利用稳态下 \Sigma_{t,t} = \frac{g^2(t)}{2\beta}
    3. 泰勒展开与高阶项忽略:仅保留 \Delta t 的一阶项。

    这一结果验证了分数函数近似公式中分母 g^2(t) \Delta t 的物理意义,即条件协方差。

    3. 举例说明

    设定一个简单的一维SDE

    我们考虑一个一维随机微分方程(SDE):

    dx = f(x,t) dt + g(t) dW

    其中:

    • 漂移项(决定趋势):f(x,t) = -0.5x (让 x 逐渐趋近于 0)
    • 扩散项(决定噪声强度):g(t) = 0.1

    即:

    dx = (-0.5x) dt + 0.1 dW


    设定一个离散时间步的加噪过程

    我们用欧拉-马尔可夫方法对其进行离散化:

    x_{t+1} = x_t + f(x_t, t) \Delta t + g(t) \sqrt{\Delta t} \cdot \epsilon

    其中:

    • \Delta t = 1.0 (假设时间间隔为 1)
    • \epsilon \sim \mathcal{N}(0,1) 是标准正态分布的随机噪声

    假设我们初始时刻 t=0 的数据 x_0​ 采样自某个真实分布,比如 p_0(x) = \mathcal{N}(2, 0.5^2)(均值为2,标准差为0.5)。


    正向加噪示例

    我们从 x_0 = 2.0 开始进行正向加噪。

    tx_t计算过程(加噪)
    02.0x_1 = 2.0 + (-0.5 \times 2.0) \times 1.0 + 0.1 \times \sqrt{1.0} \times \epsilon_0
    11.1设 \epsilon_0 = 0.2,则 x_1 = 2.0 - 1.0 + 0.1 \times 0.2 = 1.1
    20.53\epsilon_1 = -0.1,则 x_2 = 1.1 - 0.55 + 0.1 \times (-0.1) = 0.53
    30.15\epsilon_2 = -0.3,则 x_3 = 0.53 - 0.265 + 0.1 \times (-0.3) = 0.15

    这样,我们用 真实数据 x_0​ 生成了带噪数据 x_1, x_2, x_3, ...


    训练过程

    我们希望学习一个神经网络 s_\theta(x,t) 预测分数函数:

    s_\theta(x_t, t) \approx \nabla_x \log p_t(x_t)

    \nabla_x \log p_t(x_t) 不能直接计算,所以我们使用分数匹配目标

    s_\theta(x_t, t) \approx \frac{x_t - x_{t-1} - f(x_{t-1}, t-1) \Delta t}{g^2(t)}


    训练数据的构造

    对于 t=1

    • 我们有 x_1 = 1.1,希望预测 \nabla_x \log p_1(x_1)
    • 由于 x_0 = 2.0,我们使用:

    s_\theta(x_1,1) \approx \frac{1.1 - 2.0 - (-0.5 \times 2.0 \times 1.0)}{(0.1)^2}

    = \frac{1.1 - 2.0 + 1.0}{0.01} = \frac{0.1}{0.01} = 10

    所以我们希望神经网络 s_\theta(x_1,1) 预测 10

    对于 t=2

    • 我们有 x_2 = 0.53,希望预测 \nabla_x \log p_2(x_2)
    • 由于 x_1 = 1.1,我们计算:

    s_\theta(x_2,2) \approx \frac{0.53 - 1.1 - (-0.5 \times 1.1 \times 1.0)}{(0.1)^2}

    = \frac{0.53 - 1.1 + 0.55}{0.01} = \frac{-0.02}{0.01} = -2

    所以我们希望神经网络 s_\theta(x_2,2) 预测 -2


    损失函数

    神经网络 s_\theta(x,t) 通过均方误差(MSE)损失学习:

    L = \mathbb{E}_t \mathbb{E}_{p_t(x_t)} \left[ \| s_\theta(x_t, t) - \nabla_x \log p_t(x_t) \|^2 \right]

    对于一个训练样本,我们计算:

    L = (s_\theta(1.1, 1) - 10)^2 + (s_\theta(0.53, 2) + 2)^2 + ...


    逆向生成

    一旦神经网络 s_\theta(x,t) 训练完毕,我们可以使用逆向SDE

    dx = [f(x,t) - g^2(t) s_\theta(x,t)] dt + g(t) d\bar{W}

    进行数据生成:

    例如,我们从 x_T = 0.0 开始:

    1. x_3 = 0.15,用 s_\theta(0.15,3) 估计 \nabla_x \log p_3(x_3)
    2. 计算 x_2 = x_3 - f(x_3,3) \Delta t + g^2(3) s_\theta(x_3,3) \Delta t
    3. 计算 x_1 = x_2 - f(x_2,2) \Delta t + g^2(2) s_\theta(x_2,2) \Delta t
    4. 计算 x_0​,得到最终数据

    最终,我们可以生成与真实数据分布 p_0(x) 相似的样本。


    总结

    1. 正向加噪:使用 x_t = x_{t-1} + f(x_{t-1},t-1) \Delta t + g(t) \sqrt{\Delta t} \cdot \epsilon 生成带噪数据。
    2. 训练过程:神经网络 s_\theta(x,t) 通过学习 如何恢复 \nabla_x \log p_t(x_t) 来训练,使其能够预测去噪的方向。
    3. 逆向生成:使用 dx = [f(x,t) - g^2(t) s_\theta(x,t)] dt + g(t) d\bar{W} 进行反向采样,最终生成符合数据分布的样本。

    这样,我们完成了完整的训练和采样流程!

    4. 代码展示

    代码已更新,该代码只用于初步理解SDE的训练和生成。博主找到一个效果非常好的代码博客链接,有需要可以跳转至该链接:随机微分方程的分数扩散模型 (score-based diffusion model) 代码示例

    import torch
    import torch.nn as nn
    import torch.optim as optim
    from torch.utils.data import DataLoader
    from torchvision import datasets, transforms, utils
    import matplotlib.pyplot as plt
    import numpy as np
    
    
    # -----------------------------
    # 定义VP-SDE相关类
    # -----------------------------
    class VPSDE:
        """
        VP-SDE类,用于定义SDE的参数和相关计算。
        前向SDE:dx = -0.5 * beta(t) * x dt + sqrt(beta(t)) dW
        其中 beta(t) 在 [beta_min, beta_max] 内线性变化。
        """
    
        def __init__(self, beta_min=0.1, beta_max=20.0, T=1.0):
            """
            初始化SDE参数
            :param beta_min: beta的最小值
            :param beta_max: beta的最大值
            :param T: SDE终止时间
            """
            self.beta_min = beta_min
            self.beta_max = beta_max
            self.T = T
    
        def beta(self, t):
            """
            根据时间t计算beta值,此处采用线性变化:
            beta(t) = beta_min + t * (beta_max - beta_min)
            :param t: 时间(可以是标量或tensor)
            :return: beta值
            """
            return self.beta_min + t * (self.beta_max - self.beta_min)
    
        def compute_integral_beta(self, t):
            """
            计算从0到t的beta积分:
            ∫_0^t beta(s) ds = beta_min * t + 0.5 * (beta_max - beta_min) * t^2
            :param t: 时间tensor
            :return: 积分结果
            """
            return self.beta_min * t + 0.5 * (self.beta_max - self.beta_min) * (t ** 2)
    
        def marginal_prob(self, x0, t):
            """
            根据初始数据x0和时间t,计算SDE前向过程的边缘分布均值和标准差
            根据公式:x(t) = x0 * exp(-0.5 * ∫ beta ds) + std * noise,
            其中 std = sqrt(1 - exp(-∫ beta ds))
            :param x0: 原始数据(tensor)
            :param t: 时间(tensor,形状为[batch, 1, 1, 1])
            :return: 均值和标准差
            """
            integral_beta = self.compute_integral_beta(t)
            mean_coef = torch.exp(-0.5 * integral_beta)
            std = torch.sqrt(1 - torch.exp(-integral_beta))
            return mean_coef * x0, std
    
    
    # -----------------------------
    # 定义得分网络(Score Network)
    # -----------------------------
    class ScoreNet(nn.Module):
        """
        一个简单的卷积神经网络用于近似数据的得分函数,即 ∇_x log p_t(x)。
        为了让网络感知时间信息,使用了一个MLP对时间变量进行编码,并将其加入到特征中。
        """
    
        def __init__(self):
            super(ScoreNet, self).__init__()
            # 时间嵌入网络:将标量t映射到64维特征
            self.time_embed = nn.Sequential(
                nn.Linear(1, 128),
                nn.ReLU(),
                nn.Linear(128, 512),
                nn.ReLU(),
                nn.Linear(512, 256)
            )
            # 卷积层部分(针对MNIST的单通道图像)
            self.conv1 = nn.Conv2d(1, 64, kernel_size=3, padding=1)
            self.conv2 = nn.Conv2d(64, 256, kernel_size=3, padding=1)
            self.conv3 = nn.Conv2d(256, 64, kernel_size=3, padding=1)
            self.conv4 = nn.Conv2d(64, 1, kernel_size=3, padding=1)
            self.relu = nn.ReLU()
    
        def forward(self, x, t):
            """
            前向传播函数
            :param x: 带噪图像,形状 (N, 1, H, W)
            :param t: 时间变量,形状 (N, 1, 1, 1)
            :return: 网络预测的得分值,形状 (N, 1, H, W)
            """
            # 将时间变量t reshape为(N,1)后经过MLP
            t_input = t.view(x.shape[0], 1)
            t_emb = self.time_embed(t_input)  # 形状: (N, 64)
            # 将时间嵌入扩展到与图像特征同样的空间尺寸
            t_emb = t_emb.view(x.shape[0], 256, 1, 1)
    
            # 卷积部分提取图像特征
            h = self.relu(self.conv1(x))  # 输出形状: (N,32,H,W)
            h = self.relu(self.conv2(h))  # 输出形状: (N,64,H,W)
            # 将时间信息与图像特征相加(广播相加)
            h = h + t_emb  # 输出形状: (N,64,H,W)
            h = self.relu(self.conv3(h))  # 输出形状: (N,32,H,W)
            out = self.conv4(h)  # 输出形状: (N,1,H,W)
            return out
    
    
    # -----------------------------
    # 定义训练函数
    # -----------------------------
    def train(model, sde, train_loader, optimizer, device, num_epochs=10):
        """
        模型训练函数
        :param model: 得分网络模型
        :param sde: VPSDE对象
        :param train_loader: 训练数据加载器
        :param optimizer: 优化器
        :param device: 设备(CPU或GPU)
        :param num_epochs: 训练轮数
        """
        model.train()
        for epoch in range(num_epochs):
            running_loss = 0.0
            for batch_idx, (data, _) in enumerate(train_loader):
                # 将数据移动到指定设备上
                data = data.to(device)
                batch_size = data.size(0)
                optimizer.zero_grad()
                # 随机采样时间t,范围在[1e-5, T]之间,保证t不为0
                t = torch.rand(batch_size, device=device) * (sde.T - 1e-5) + 1e-5
                # 将t reshape为与图像数据匹配的形状 (N,1,1,1)
                t = t.view(batch_size, 1, 1, 1)
                # 根据SDE计算数据x0在时间t时的均值和标准差
                x_t_mean, std = sde.marginal_prob(data, t)
                # 生成与数据形状一致的高斯噪声
                noise = torch.randn_like(data)
                # 根据前向SDE公式,生成带噪数据:x_t = mean + std * noise
                x_t = x_t_mean + std * noise
    
                # 得分网络预测得分函数 s_theta(x_t, t)
                score = model(x_t, t)
                # 构造去噪得分匹配损失:
                # 理想情况下,score应等于 -noise / std
                loss = ((score + noise / std) ** 2).mean()
    
                # 反向传播和参数更新
                loss.backward()
                optimizer.step()
                running_loss += loss.item()
    
            avg_loss = running_loss / len(train_loader)
            print("Epoch [{}/{}], Loss: {:.4f}".format(epoch + 1, num_epochs, avg_loss))
    
    
    # -----------------------------
    # 定义测试函数
    # -----------------------------
    def evaluate(model, sde, test_loader, device):
        """
        模型测试函数,用于评估模型在测试集上的平均损失
        :param model: 得分网络模型
        :param sde: VPSDE对象
        :param test_loader: 测试数据加载器
        :param device: 设备(CPU或GPU)
        """
        model.eval()
        test_loss = 0.0
        with torch.no_grad():
            for batch_idx, (data, _) in enumerate(test_loader):
                data = data.to(device)
                batch_size = data.size(0)
                t = torch.rand(batch_size, device=device) * (sde.T - 1e-5) + 1e-5
                t = t.view(batch_size, 1, 1, 1)
                x_t_mean, std = sde.marginal_prob(data, t)
                noise = torch.randn_like(data)
                x_t = x_t_mean + std * noise
    
                score = model(x_t, t)
                loss = ((score + noise / std) ** 2).mean()
                test_loss += loss.item()
        avg_loss = test_loss / len(test_loader)
        print("Test Loss: {:.4f}".format(avg_loss))
    
    
    # -----------------------------
    # 定义生成函数(基于欧拉-马鲁雅玛方法采样)
    # -----------------------------
    def generate_samples(model, sde, device, num_steps=1000, shape=(64, 1, 28, 28)):
        """
        通过反向SDE生成样本,使用欧拉-马鲁雅玛方法
        :param model: 得分网络模型
        :param sde: VPSDE对象
        :param device: 设备(CPU或GPU)
        :param num_steps: 采样步数,步数越多生成质量越好
        :param shape: 生成样本的形状,例如 (batch_size, 通道, 高, 宽)
        :return: 生成的样本tensor
        """
        model.eval()
        with torch.no_grad():
            # 采样时间步长,转换为 tensor
            dt = sde.T / num_steps
            dt_tensor = torch.tensor(dt, device=device)  # 将dt转换为Tensor
            # 初始化x_T ~ N(0, I)
            x = torch.randn(shape, device=device)
            # 从T到接近0的时间序列
            time_steps = torch.linspace(sde.T, 1e-3, num_steps, device=device)
            for t in time_steps:
                # 构造与当前批次大小对应的时间tensor,形状为 (N,1,1,1)
                t_tensor = torch.ones((x.shape[0], 1, 1, 1), device=device) * t
                # 计算当前时间点的beta值(标量)
                beta_t = sde.beta(t)
                # 根据反向SDE计算漂移项: dx = [ -0.5*beta(t)*x - beta(t)*score(x,t) ] dt + sqrt(beta(t)) dW
                drift = -0.5 * beta_t * x - beta_t * model(x, t_tensor)
                diffusion = torch.sqrt(beta_t)
                # 采样噪声,并进行欧拉-马鲁雅玛更新
                noise = torch.randn_like(x)
                x = x - drift * dt + diffusion * torch.sqrt(dt_tensor) * noise  # dt为正,用显式减法。
            return x
    
    # -----------------------------
    # 主函数:加载数据、训练、测试和生成样本
    # -----------------------------
    def main():
        # 设置设备,如果有GPU则使用GPU
        device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    
        # 超参数设置
        batch_size = 128
        num_epochs = 10  # 根据需求可适当增加训练轮数
        learning_rate = 1e-3
    
        # 数据预处理:将图像转换为Tensor,并归一化到[0,1]
        transform = transforms.Compose([
            transforms.ToTensor()
        ])
        # 加载MNIST训练和测试数据集
        train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transform)
        test_dataset = datasets.MNIST(root='./data', train=False, download=True, transform=transform)
        train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)
        test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)
    
        # 初始化得分网络和SDE模型
        model = ScoreNet().to(device)
        sde = VPSDE(beta_min=0.1, beta_max=20.0, T=1.0)
    
        # 定义优化器(Adam)
        optimizer = optim.Adam(model.parameters(), lr=learning_rate)
    
        # -----------------------------
        # 模型训练
        # -----------------------------
        print("开始训练...")
        train(model, sde, train_loader, optimizer, device, num_epochs=num_epochs)
    
        # -----------------------------
        # 模型测试
        # -----------------------------
        print("开始测试...")
        evaluate(model, sde, test_loader, device)
    
        # -----------------------------
        # 利用训练好的模型生成样本
        # -----------------------------
        print("开始生成样本...")
        # 例如生成64张28x28的图像
        samples = generate_samples(model, sde, device, num_steps=1000, shape=(16, 1, 28, 28))
        # 将生成的样本保存为图片(网格形式)
        grid_img = utils.make_grid(samples, nrow=8, normalize=True)
        plt.figure(figsize=(8, 8))
        plt.imshow(grid_img.permute(1, 2, 0).cpu().numpy())
        plt.axis('off')
        plt.title("Generated Samples")
        plt.savefig("generated_samples_gpt.png")
        plt.show()
    
    
    if __name__ == "__main__":
        main()
    

    Logo

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

    更多推荐