这篇文章想从随机微分方程(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 是怎么从“把扩散过程反过来”自然出现的。
先给一个大图景:
- SDE 描述单个样本点如何随机运动。
- FPE 描述所有样本形成的概率密度如何随时间变化。
- 扩散模型的前向过程是从数据分布 \(p_0\) 走向噪声分布 \(p_T\)。
- 生成过程要从 \(p_T\) 走回 \(p_0\),所以需要把前向 SDE 反过来。
- 把 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 到扩散模型,可以按下面的顺序理解:
- 前向 SDE 描述加噪过程。
- FPE 描述概率密度 \(p_t\) 如何演化。
- 生成过程要反向走,所以引入 \(\tau=T-t\)。
- 把前向 FPE 改写成反向时间下的密度演化方程。
- 将它配凑回标准 FPE 形式。
- 从标准 FPE 读出逆向 SDE。
- 逆向漂移中自然出现 \(\nabla_{\mathbf{x}}\log p_t(\mathbf{x})\),也就是 score。
扩散模型的核心不是凭空从噪声生成数据,而是学习每个噪声强度下概率密度上升最快的方向,然后沿这个方向场把样本一步步推回数据分布。