这篇文章想从随机微分方程(Stochastic Differential Equation, SDE)出发,推到扩散模型里常见的逆向 SDE:

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

这个式子看起来突然出现了一个 \(\nabla_{\mathbf{x}}\log p_t(\mathbf{x})\),也就是 score。本文的主线就是解释:这个 score 是怎么从“把扩散过程反过来”自然出现的。

先给一个大图景:

  1. SDE 描述单个样本点如何随机运动。
  2. FPE 描述所有样本形成的概率密度如何随时间变化。
  3. 扩散模型的前向过程是从数据分布 \(p_0\) 走向噪声分布 \(p_T\)。
  4. 生成过程要从 \(p_T\) 走回 \(p_0\),所以需要把前向 SDE 反过来。
  5. 把 FPE 的时间反过来并整理成标准形式,就会得到逆向 SDE,其中自然出现 score。

1. 从 ODE 到 SDE

普通微分方程(ODE)描述确定性系统:

\[\frac{d\mathbf{x}}{dt}=f(\mathbf{x},t)\]

给定初始状态之后,轨迹是确定的。很多高阶微分方程也可以写成一阶向量形式。例如二阶系统可以引入速度 \(v=\dot{x}\),改写为:

\[\frac{d}{dt} \begin{bmatrix} x \\ v \end{bmatrix} = \begin{bmatrix} v \\ F/m \end{bmatrix}\]

但现实系统里经常存在随机扰动,比如传感器噪声、建模误差、环境扰动。此时只写确定性漂移不够,需要加入随机过程项:

\[d\mathbf{x}_t = \mathbf{f}(\mathbf{x}_t,t)dt {}+ g(\mathbf{x}_t,t)d\mathbf{W}_t\]

这里:

  • \(\mathbf{f}(\mathbf{x}_t,t)dt\) 是漂移项,表示平均意义下系统往哪里走。
  • \(g(\mathbf{x}_t,t)d\mathbf{W}_t\) 是扩散项,表示随机扰动有多强。
  • \(\mathbf{W}_t\) 是标准布朗运动,也叫 Wiener 过程。

Wiener 过程的增量满足:

\[d\mathbf{W}_t = \mathbf{W}_{t+dt}-\mathbf{W}_t \sim \mathcal{N}(\mathbf{0},\mathbf{I}dt)\]

因此:

\[\mathbb{E}[d\mathbf{W}_t]=\mathbf{0}, \qquad \mathbb{E}[d\mathbf{W}_t d\mathbf{W}_t^\mathsf{T}] = \mathbf{I}dt\]

这说明噪声增量的标准差是 \(\sqrt{dt}\) 量级,而不是 \(dt\) 量级。于是一个非常重要的规则出现了:

\[(dW_t)^2 \sim dt\]

普通微积分里二阶小量通常可以忽略,但 SDE 里 \((dW_t)^2\) 会留下 \(dt\) 级别的贡献。这是后面 Itô 引理里多出二阶项的根源。

一个直观例子是 IMU bias 的随机游走:

\[d\mathbf{b}_t = \sigma_b d\mathbf{W}_t\]

离散化后:

\[\mathbf{b}_{t+\Delta t} = \mathbf{b}_t + \mathbf{w}_{b,t}, \qquad \mathbf{w}_{b,t} \sim \mathcal{N}(\mathbf{0},\sigma_b^2\Delta t\,\mathbf{I})\]

也就是说,时间越长,bias 的方差越大;标准差按 \(\sqrt{t}\) 增长。这就是“扩散”的直觉。

2. Itô 引理:为什么会出现二阶项

先看一维 SDE:

\[dX_t = a(X_t,t)dt + b(X_t,t)dW_t\]

如果有一个函数 \(h(t,X_t)\),我们希望知道 \(h\) 如何变化。对 \(h\) 做泰勒展开:

\[dh = \frac{\partial h}{\partial t}dt {}+ \frac{\partial h}{\partial X}dX_t {}+ \frac{1}{2}\frac{\partial^2 h}{\partial X^2}(dX_t)^2 {}+ \cdots\]

关键是计算 \((dX_t)^2\)。代入 SDE:

\[\begin{aligned} (dX_t)^2 &= \left(a\,dt+b\,dW_t\right)^2 \\ &= a^2(dt)^2 + 2ab\,dt\,dW_t + b^2(dW_t)^2 \\ &= b^2dt \end{aligned}\]

这里使用了 Itô 乘法规则:

\(dt\)\(dW_t\)
\(dt\)\(0\)\(0\)
\(dW_t\)\(0\)\(dt\)

所以 Itô 引理为:

\[dh = \left( \frac{\partial h}{\partial t} + a\frac{\partial h}{\partial X} + \frac{1}{2}b^2\frac{\partial^2 h}{\partial X^2} \right)dt {}+ b\frac{\partial h}{\partial X}dW_t\]

这个式子后面用来从 SDE 推导概率密度的演化方程。

3. 从 SDE 推导 FPE

SDE 描述的是单个随机变量 \(X_t\) 的运动:

\[dX_t = a(X_t,t)dt + b(X_t,t)dW_t\]

但扩散模型更关心整个分布 \(p(x,t)\) 如何变化。这个变化由福克-普朗克方程(Fokker-Planck Equation, FPE)描述。

为了推导它,我们引入任意光滑测试函数 \(h(x)\)。它的期望为:

\[\mathbb{E}[h(X_t)] = \int_{-\infty}^{\infty}h(x)p(x,t)dx\]

从概率密度的角度,对时间求导:

\[\frac{d}{dt}\mathbb{E}[h(X_t)] = \frac{d}{dt}\int h(x)p(x,t)dx = \int h(x)\frac{\partial p(x,t)}{\partial t}dx \tag{1}\]

另一方面,根据 Itô 引理,因为 \(h\) 不显式依赖 \(t\),有:

\[dh(X_t) = \left[ a\frac{\partial h}{\partial x} + \frac{1}{2}b^2\frac{\partial^2 h}{\partial x^2} \right]dt {}+ b\frac{\partial h}{\partial x}dW_t\]

对两边取期望,随机项的期望为零:

\[\mathbb{E}[dh(X_t)] = \mathbb{E}\left[ a\frac{\partial h}{\partial x} + \frac{1}{2}b^2\frac{\partial^2 h}{\partial x^2} \right]dt\]

因此:

\[\frac{d}{dt}\mathbb{E}[h(X_t)] = \int \left[ a(x,t)\frac{\partial h}{\partial x} + \frac{1}{2}b^2(x,t)\frac{\partial^2 h}{\partial x^2} \right]p(x,t)dx \tag{2}\]

比较式 (1) 和式 (2):

\[\int h(x)\frac{\partial p}{\partial t}dx = \int \left[ a\frac{\partial h}{\partial x}p + \frac{1}{2}b^2\frac{\partial^2 h}{\partial x^2}p \right]dx\]

我们的目标是把右边 \(h\) 的导数转移到 \(a,b,p\) 上,这样才能得到关于 \(p\) 的偏微分方程。这一步靠分部积分。

3.1 漂移项

先处理:

\[\int a\frac{\partial h}{\partial x}p\,dx = \int (ap)\frac{\partial h}{\partial x}dx\]

分部积分:

\[\int (ap)\frac{\partial h}{\partial x}dx = \left[(ap)h\right]_{-\infty}^{\infty} - \int h\frac{\partial(ap)}{\partial x}dx\]

假设边界项为零,则:

\[\int a\frac{\partial h}{\partial x}p\,dx = -\int h\frac{\partial(ap)}{\partial x}dx\]

3.2 扩散项

再处理:

\[\int \frac{1}{2}b^2\frac{\partial^2 h}{\partial x^2}p\,dx = \int \frac{1}{2}b^2p\frac{\partial^2 h}{\partial x^2}dx\]

需要做两次分部积分。第一次:

\[\int \frac{1}{2}b^2p\frac{\partial^2 h}{\partial x^2}dx = -\int \frac{\partial}{\partial x}\left(\frac{1}{2}b^2p\right) \frac{\partial h}{\partial x}dx\]

第二次:

\[-\int \frac{\partial}{\partial x}\left(\frac{1}{2}b^2p\right) \frac{\partial h}{\partial x}dx = \int h \frac{\partial^2}{\partial x^2} \left(\frac{1}{2}b^2p\right)dx\]

所以扩散项变成:

\[\int \frac{1}{2}b^2\frac{\partial^2 h}{\partial x^2}p\,dx = \int h \frac{1}{2}\frac{\partial^2}{\partial x^2} \left(b^2p\right)dx\]

3.3 得到 FPE

把漂移项和扩散项代回:

\[\int h(x)\frac{\partial p}{\partial t}dx = \int h(x) \left[ -\frac{\partial}{\partial x}(ap) + \frac{1}{2}\frac{\partial^2}{\partial x^2}(b^2p) \right]dx\]

因为这个等式对任意测试函数 \(h(x)\) 都成立,所以被积函数相等:

\[\frac{\partial p(x,t)}{\partial t} = -\frac{\partial}{\partial x}\left[a(x,t)p(x,t)\right] {}+ \frac{1}{2}\frac{\partial^2}{\partial x^2}\left[b^2(x,t)p(x,t)\right] \tag{3}\]

这就是一维 FPE。它的直观含义是:

  • 漂移项 \(a\) 负责搬运概率密度。
  • 扩散项 \(b\) 负责摊开概率密度。

多维情况下,如果扩散系数只依赖时间:

\[d\mathbf{x}_t = \mathbf{f}(\mathbf{x}_t,t)dt {}+ g(t)d\mathbf{w}_t\]

对应的 FPE 是:

\[\frac{\partial p(\mathbf{x},t)}{\partial t} = -\nabla\cdot\left[\mathbf{f}(\mathbf{x},t)p(\mathbf{x},t)\right] {}+ \frac{1}{2}g^2(t)\Delta p(\mathbf{x},t) \tag{4}\]

其中 \(\Delta p=\nabla\cdot(\nabla p)\)。

4. 扩散模型里的前向 SDE

扩散模型的前向过程是连续加噪:

\[p_0 \longrightarrow p_t \longrightarrow p_T\]

\(p_0\) 是数据分布,\(p_T\) 通常接近标准高斯噪声。我们写前向 SDE:

\[d\mathbf{x}_t = \mathbf{f}(\mathbf{x}_t,t)dt {}+ g(t)d\mathbf{w}_t \tag{5}\]

它对应的 FPE 是:

\[\frac{\partial p_t}{\partial t} = -\nabla\cdot(\mathbf{f}p_t) {}+ \frac{1}{2}g^2(t)\nabla\cdot(\nabla p_t) \tag{6}\]

这里为了书写简洁,把 \(p(\mathbf{x},t)\) 写成 \(p_t\),把 \(\mathbf{f}(\mathbf{x},t)\) 写成 \(\mathbf{f}\)。

生成过程则是反过来:

\[p_T \longrightarrow p_t \longrightarrow p_0\]

只把时间反过来还不够,因为前向扩散会丢失信息。反向过程必须知道“当前位置应该往哪个高概率区域移动”。这个方向由 score 给出:

\[\nabla_{\mathbf{x}}\log p_t(\mathbf{x})\]

神经网络通常学习的就是这个量:

\[s_\theta(\mathbf{x},t) \approx \nabla_{\mathbf{x}}\log p_t(\mathbf{x})\]

下面推导逆向 SDE,看看这个 score 是如何出现的。

5. 引入反向时间变量

生成过程从 \(t=T\) 走回 \(t=0\)。为了把它写成一个“随新时间正向演化”的过程,定义:

\[\tau = T-t, \qquad t=T-\tau\]

当 \(\tau:0\to T\) 时,原时间 \(t:T\to 0\)。为了避免一边写 \(p_t\) 一边对 \(\tau\) 求导,我们定义反向时间上的密度:

\[q_\tau(\mathbf{x}) \triangleq p_{T-\tau}(\mathbf{x}) = p_t(\mathbf{x})\]

于是:

\[q_0=p_T, \qquad q_T=p_0\]

链式法则给出:

\[\frac{\partial q_\tau}{\partial \tau} = \frac{\partial p_t}{\partial t}\frac{\partial t}{\partial \tau} = -\frac{\partial p_t}{\partial t}, \qquad t=T-\tau\]

把前向 FPE 式 (6) 代进去:

\[\frac{\partial q_\tau}{\partial \tau} = \nabla\cdot(\mathbf{f}q_\tau) - \frac{1}{2}g^2(t)\nabla\cdot(\nabla q_\tau), \qquad t=T-\tau \tag{7}\]

式 (7) 描述了概率密度随反向时间 \(\tau\) 的变化。但它还不是一个标准 FPE 的形式。

6. 把反向密度方程配凑成标准 FPE

任意 SDE:

\[d\mathbf{x} = \tilde{\mathbf{f}}(\mathbf{x},\tau)d\tau {}+ \tilde{g}(\tau)d\tilde{\mathbf{w}}_\tau\]

对应的标准 FPE 是:

\[\frac{\partial q_\tau}{\partial \tau} = -\nabla\cdot(\tilde{\mathbf{f}}q_\tau) {}+ \frac{1}{2}\tilde{g}^2\nabla\cdot(\nabla q_\tau) \tag{8}\]

对比式 (7),可以看到扩散项前面是负号:

\[\frac{\partial q_\tau}{\partial \tau} = \nabla\cdot(\mathbf{f}q_\tau) - \frac{1}{2}g^2(t)\nabla\cdot(\nabla q_\tau)\]

这里最容易让人困惑:为什么可以“配凑”?

原因是这一步没有改变方程,只是在做恒等变形。标准 FPE 里扩散项必须是正的,因为它来自 SDE 中噪声强度的平方 \(\tilde{g}^2\),不可能把一个负扩散系数直接解释成合法的 SDE。因此,我们要把反向时间方程里的“负扩散”拆成两部分:一部分保留下来作为标准的正扩散项,另一部分吸收到漂移项里。

代数上就是:

\[-\frac{1}{2}A = -A + \frac{1}{2}A\]

其中:

\[A = g^2(t)\nabla\cdot(\nabla q_\tau)\]

所以,为了凑出标准形式,在右侧同时加上并减去同一项:

\[\begin{aligned} \frac{\partial q_\tau}{\partial \tau} &= \nabla\cdot(\mathbf{f}q_\tau) - \frac{1}{2}g^2(t)\nabla\cdot(\nabla q_\tau) - \frac{1}{2}g^2(t)\nabla\cdot(\nabla q_\tau) {}+ \frac{1}{2}g^2(t)\nabla\cdot(\nabla q_\tau) \\ &= \nabla\cdot(\mathbf{f}q_\tau) - g^2(t)\nabla\cdot(\nabla q_\tau) {}+ \frac{1}{2}g^2(t)\nabla\cdot(\nabla q_\tau) \end{aligned} \tag{9}\]

现在保留最后一项作为标准扩散项,集中处理前两项:

\[\nabla\cdot(\mathbf{f}q_\tau) - g^2(t)\nabla\cdot(\nabla q_\tau) = -\nabla\cdot\left[-\mathbf{f}q_\tau+g^2(t)\nabla q_\tau\right] \tag{10}\]

这里用到对数导数恒等式:

\[\nabla\log q_\tau(\mathbf{x}) = \frac{\nabla q_\tau(\mathbf{x})}{q_\tau(\mathbf{x})}, \qquad \nabla q_\tau = q_\tau\nabla\log q_\tau \tag{11}\]

代入式 (10):

\[\begin{aligned} -\nabla\cdot\left[-\mathbf{f}q_\tau+g^2(t)\nabla q_\tau\right] &= -\nabla\cdot\left[-\mathbf{f}q_\tau+g^2(t)q_\tau\nabla\log q_\tau\right] \\ &= -\nabla\cdot\left[\left(-\mathbf{f}+g^2(t)\nabla\log q_\tau\right)q_\tau\right] \end{aligned} \tag{12}\]

把式 (12) 代回式 (9),得到:

\[\frac{\partial q_\tau}{\partial \tau} = -\nabla\cdot\left[ \left(-\mathbf{f}+g^2(t)\nabla\log q_\tau\right)q_\tau \right] {}+ \frac{1}{2}g^2(t)\nabla\cdot(\nabla q_\tau) \tag{13}\]

现在式 (13) 已经是标准 FPE 形式。直观地说,时间反过来以后,原来的扩散不可能真的变成“负噪声”,所以它必须表现为一个额外的漂移修正。这个修正项正是 \(g^2(t)\nabla\log q_\tau\),也就是后面逆向 SDE 里的 score 项。

和式 (8) 对比,可以读出:

\[\tilde{\mathbf{f}} = -\mathbf{f}(\mathbf{x}_\tau,t) {}+ g^2(t)\nabla_{\mathbf{x}}\log q_\tau(\mathbf{x}_\tau), \qquad \tilde{g}=g(t)\]

因此,反向时间 \(\tau\) 下的 SDE 是:

\[d\mathbf{x}_\tau = \left[ -\mathbf{f}(\mathbf{x}_\tau,t) {}+ g^2(t)\nabla_{\mathbf{x}}\log q_\tau(\mathbf{x}_\tau) \right]d\tau {}+ g(t)d\mathbf{w}_\tau, \qquad t=T-\tau \tag{14}\]

因为 \(q_\tau(\mathbf{x})=p_t(\mathbf{x})\),所以这里的 score 也可以写成:

\[\nabla_{\mathbf{x}}\log q_\tau(\mathbf{x}) = \nabla_{\mathbf{x}}\log p_t(\mathbf{x})\]

这就是 score 项出现的地方。

7. 写回原时间变量 \(t\)

工程实现和论文里通常继续使用原时间变量 \(t\),只是让它从 \(T\) 递减到 \(0\)。因为:

\[d\tau=-dt\]

将式 (14) 中的 \(q_\tau=p_t\)、\(d\tau=-dt\) 代入:

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

把负号分配进去:

\[d\mathbf{x} = \left[ \mathbf{f}(\mathbf{x},t) - g^2(t)\nabla_{\mathbf{x}}\log p_t(\mathbf{x}) \right]dt {}+ g(t)d\bar{\mathbf{w}} \tag{15}\]

这就是 Yang Song et al. 论文中常见的逆向 SDE 形式。这里的 \(dt\) 可以理解为沿原时间轴反向走的负时间步长。

如果令 \(\Delta t=T/N>0\),从 \(t\) 走到 \(t-\Delta t\),则 Euler-Maruyama 离散化为:

\[\mathbf{x}_{t-\Delta t} = \mathbf{x}_t - \left[ \mathbf{f}(\mathbf{x}_t,t) - g^2(t)s_\theta(\mathbf{x}_t,t) \right]\Delta t {}+ g(t)\sqrt{\Delta t}\mathbf{z}, \qquad \mathbf{z}\sim\mathcal{N}(\mathbf{0},\mathbf{I}) \tag{16}\]

其中:

\[s_\theta(\mathbf{x}_t,t) \approx \nabla_{\mathbf{x}}\log p_t(\mathbf{x}_t)\]

对应的采样代码:

import torch


def euler_maruyama_sampler(model, shape, T, steps, f_func, g_func, device="cuda"):
    """Euler-Maruyama sampler for the reverse SDE."""
    delta_t = T / steps
    t_sequence = torch.linspace(T, 1e-3, steps, device=device)
    x = torch.randn(shape, device=device)

    for t in t_sequence:
        t_batch = torch.full((shape[0],), t, device=device)

        score = model(x, t_batch)
        f_t = f_func(x, t_batch)
        g_t = g_func(t_batch)

        reverse_drift = f_t - (g_t ** 2) * score
        noise = torch.randn_like(x)
        diffusion = g_t * torch.sqrt(torch.tensor(delta_t, device=device)) * noise

        x = x - reverse_drift * delta_t + diffusion

    return x

代码和公式的对应关系:

  • score 对应 \(s_\theta(\mathbf{x}_t,t)\)。
  • reverse_drift 对应 \(\mathbf{f}-g^2s_\theta\)。
  • x = x - reverse_drift * delta_t + diffusion 对应从 \(t\) 更新到 \(t-\Delta t\)。
  • diffusion 对应 \(g(t)\sqrt{\Delta t}\mathbf{z}\)。

8. 和 DDPM、DDIM、DPM-Solver 的关系

上面的推导是连续时间视角。DDPM、DDIM、DPM-Solver 可以看成这个框架下不同的离散化或求解方式。

DDPM 是随机采样。它每一步都会注入噪声,对应逆向 SDE 的离散版本。

DDIM 更接近确定性采样。当随机噪声项被减弱甚至去掉时,采样轨迹更稳定,也可以用更少步数生成。

DPM-Solver 则把扩散采样看成微分方程求解问题,用更高阶的数值方法减少采样步数。

粗略地说:

\[\text{连续时间 SDE/ODE 框架} \quad\Longrightarrow\quad \text{不同采样器选择不同数值解法}\]

总结

从 SDE 到扩散模型,可以按下面的顺序理解:

  1. 前向 SDE 描述加噪过程。
  2. FPE 描述概率密度 \(p_t\) 如何演化。
  3. 生成过程要反向走,所以引入 \(\tau=T-t\)。
  4. 把前向 FPE 改写成反向时间下的密度演化方程。
  5. 将它配凑回标准 FPE 形式。
  6. 从标准 FPE 读出逆向 SDE。
  7. 逆向漂移中自然出现 \(\nabla_{\mathbf{x}}\log p_t(\mathbf{x})\),也就是 score。

扩散模型的核心不是凭空从噪声生成数据,而是学习每个噪声强度下概率密度上升最快的方向,然后沿这个方向场把样本一步步推回数据分布。


参考文献