Skip to content
huc
Go back

得分扩散的 SDE 与概率流 ODE

目录 1 / 7

《生成扩散模型漫谈(五):一般框架之SDE篇》 —— 苏剑林

《生成扩散模型漫谈(六):一般框架之ODE篇》 —— 苏剑林

连续时间把过程与离散步数分开

离散 DDPM 需要事先指定 TT,连续时间框架则先定义一个前向随机微分方程:

dx=ft(x)dt+gtdw.(1)dx=f_t(x)\,dt+g_t\,dw. \tag{1}

ft(x)f_t(x) 是 drift,决定确定性运动方向;gtg_t 是 diffusion coefficient,ww 是标准 Wiener 过程。理解公式 (1) 时,可以先看它的 Euler–Maruyama 离散形式:

xt+Δtxt=ft(xt)Δt+gtΔtε,εN(0,I).(2)x_{t+\Delta t}-x_t =f_t(x_t)\Delta t +g_t\sqrt{\Delta t}\,\varepsilon, \qquad \varepsilon\sim\mathcal N(0,I). \tag{2}

随机项必须是 O(Δt)\mathcal O(\sqrt{\Delta t})。若它也是 O(Δt)\mathcal O(\Delta t),把长度固定的时间区间分成 nn 段后,独立噪声和的方差大约是 n(1/n)2=1/nn(1/n)^2=1/n,极限会消失;取 Δt\sqrt{\Delta t} 时,方差约为 n(1/n)2=1n(1/\sqrt n)^2=1,随机效应才能留在连续极限中。

因此 TT 不再是模型定义的一部分,而是数值求解精度的一部分。理论上分析同一条连续轨迹,实现时再选择时间网格和求解器。

从局部贝叶斯推到反向 SDE

公式 (2) 对应的单步转移核是

p(xt+Δtxt)=N(xt+Δt;xt+ft(xt)Δt,gt2ΔtI).(3)p(x_{t+\Delta t}\mid x_t) =\mathcal N\left( x_{t+\Delta t}; x_t+f_t(x_t)\Delta t, g_t^2\Delta t\,I \right). \tag{3}

反向局部核由贝叶斯公式给出:

p(xtxt+Δt)=p(xt+Δtxt)exp[logpt(xt)logpt+Δt(xt+Δt)]×C,p(x_t\mid x_{t+\Delta t}) =p(x_{t+\Delta t}\mid x_t) \exp\left[ \log p_t(x_t)-\log p_{t+\Delta t}(x_{t+\Delta t}) \right] \times C,

其中 CCxtx_t 无关。由于高斯局部核只在 xt+Δtxt=O(Δt)x_{t+\Delta t}-x_t=\mathcal O(\sqrt{\Delta t}) 时有显著密度,可以在 (xt,t)(x_t,t) 附近展开:

logpt+Δt(xt+Δt)logpt(xt)+(xt+Δtxt)xlogpt(xt)+Δttlogpt(xt).(4)\begin{aligned} \log p_{t+\Delta t}(x_{t+\Delta t}) \approx{}&\log p_t(x_t) +(x_{t+\Delta t}-x_t)^\top\nabla_x\log p_t(x_t)\\ &+\Delta t\,\partial_t\log p_t(x_t). \end{aligned} \tag{4}

时间偏导不能省略,因为密度本身也随扩散时间变化。把公式 (4) 代回局部贝叶斯表达式并对 xtx_t 配方,保留决定高斯均值和方差的最低阶项,反向一步的均值相对 xt+Δtx_{t+\Delta t} 多出 [ftgt2xlogpt]Δt-[f_t-g_t^2\nabla_x\log p_t]\Delta t。连续极限就是

dx=[ft(x)gt2xlogpt(x)]dt+gtdwˉ,(5)dx=\left[f_t(x)-g_t^2\nabla_x\log p_t(x)\right]dt +g_t\,d\bar w, \tag{5}

其中公式沿 t:T0t:T\to0 积分,dtdt 为负;wˉ\bar w 表示反向时间的 Wiener 过程。前向与反向使用相同的瞬时扩散强度,额外出现的 gt2xlogpt(x)-g_t^2\nabla_x\log p_t(x) 把样本推向当时的高密度区域。

这也说明网络真正缺失的量不是“噪声”这个特定参数化,而是边缘分布的 score:

s(x,t)=xlogpt(x).s^*(x,t)=\nabla_x\log p_t(x).

条件 score 为什么能学到边缘 score

边缘分布是数据分布经过前向核后的混合:

pt(xt)=p(xtx0)p~(x0)dx0.p_t(x_t)=\int p(x_t\mid x_0)\tilde p(x_0)\,dx_0.

xtx_t 求梯度并除以 pt(xt)p_t(x_t)

xtlogpt(xt)=p(xtx0)p~(x0)xtlogp(xtx0)dx0pt(xt)=Ep(x0xt)[xtlogp(xtx0)].(6)\begin{aligned} \nabla_{x_t}\log p_t(x_t) &=\frac{\int p(x_t\mid x_0)\tilde p(x_0) \nabla_{x_t}\log p(x_t\mid x_0)\,dx_0} {p_t(x_t)}\\ &=\mathbb E_{p(x_0\mid x_t)} \left[\nabla_{x_t}\log p(x_t\mid x_0)\right]. \end{aligned} \tag{6}

也就是说,给定带噪样本 xtx_t,条件 score 的后验均值恰好等于难以直接计算的边缘 score。平方损失的最优回归函数是条件均值,因此训练

LDSM=Ex0,t,xtp(xtx0)[sθ(xt,t)xtlogp(xtx0)2](7)L_{\mathrm{DSM}} =\mathbb E_{x_0,t,x_t\sim p(x_t\mid x_0)} \left[ \left\|s_\theta(x_t,t) -\nabla_{x_t}\log p(x_t\mid x_0) \right\|^2 \right] \tag{7}

就能得到 sθ(xt,t)xtlogpt(xt)s_\theta(x_t,t)\approx\nabla_{x_t}\log p_t(x_t)。公式 (6) 是这个结论的关键条件:训练标签是条件 score,模型输入却只有 xt,tx_t,t,平方回归自动对不可见的 x0x_0 做后验平均。

若选择高斯扰动核

xt=αˉtx0+βˉtε,εN(0,I),x_t=\bar\alpha_t x_0+\bar\beta_t\varepsilon, \qquad \varepsilon\sim\mathcal N(0,I),

则条件 score 有解析式

xtlogp(xtx0)=xtαˉtx0βˉt2=εβˉt.(8)\nabla_{x_t}\log p(x_t\mid x_0) =-\frac{x_t-\bar\alpha_t x_0}{\bar\beta_t^2} =-\frac{\varepsilon}{\bar\beta_t}. \tag{8}

sθ(xt,t)=εθ(xt,t)/βˉts_\theta(x_t,t)=-\varepsilon_\theta(x_t,t)/\bar\beta_t,公式 (7) 就变成带权噪声预测:

LDSM=E[1βˉt2εεθ(xt,t)2].L_{\mathrm{DSM}} =\mathbb E\left[ \frac{1}{\bar\beta_t^2} \left\|\varepsilon-\varepsilon_\theta(x_t,t)\right\|^2 \right].

去掉 1/βˉt21/\bar\beta_t^2 后得到 DDPM 常用的 simple loss,但不同噪声强度的相对权重已经改变。

先指定扰动核,再反推线性 SDE

从公式 (1) 出发,未必容易求出 p(xtx0)p(x_t\mid x_0)。更适合训练的路线是先设计一个可直接采样、条件 score 可解析的高斯扰动核,再求与之匹配的 SDE。

p(xtx0)=N(xt;αˉtx0,βˉt2I),p(x_t\mid x_0) =\mathcal N(x_t;\bar\alpha_t x_0,\bar\beta_t^2I),

并寻找线性过程 dx=ftxdt+gtdwdx=f_t x\,dt+g_t\,dw。短时间转移为

xt+Δt=(1+ftΔt)xt+gtΔtε2.x_{t+\Delta t}=(1+f_t\Delta t)x_t +g_t\sqrt{\Delta t}\,\varepsilon_2.

xt=αˉtx0+βˉtε1x_t=\bar\alpha_t x_0+\bar\beta_t\varepsilon_1 代入,并分别匹配 xt+Δtx_{t+\Delta t} 的条件均值与方差:

αˉt+Δt=(1+ftΔt)αˉt,βˉt+Δt2=(1+ftΔt)2βˉt2+gt2Δt.\begin{aligned} \bar\alpha_{t+\Delta t} &=(1+f_t\Delta t)\bar\alpha_t,\\ \bar\beta_{t+\Delta t}^2 &=(1+f_t\Delta t)^2\bar\beta_t^2+g_t^2\Delta t. \end{aligned}

Δt0\Delta t\to0,得到

ft=ddtlogαˉt,gt2=αˉt2ddt(βˉt2αˉt2).(9)f_t=\frac{d}{dt}\log\bar\alpha_t, \qquad g_t^2=\bar\alpha_t^2 \frac{d}{dt}\left( \frac{\bar\beta_t^2}{\bar\alpha_t^2} \right). \tag{9}

αˉt1\bar\alpha_t\equiv1 时只有方差增长,对应 VE-SDE;当 αˉt2+βˉt2=1\bar\alpha_t^2+\bar\beta_t^2=1 时总体尺度保持稳定,对应 VP-SDE。Noise Schedule 不只是离散超参数,它完整决定了连续 drift 和 diffusion coefficient。

Fokker–Planck 方程只追踪边缘密度

反向 SDE 仍然是随机路径。如果只关心每个时刻的边缘分布 pt(x)p_t(x),需要把路径方程转成密度演化方程。

用 Dirac 函数表示密度:

pt(x)=E[δ(xxt)].p_t(x)=\mathbb E[\delta(x-x_t)].

把公式 (2) 代入 δ(xxt+Δt)\delta(x-x_{t+\Delta t}),围绕 xxtx-x_t 做二阶展开。对噪声求期望后,一阶随机项因 Eε=0\mathbb E\varepsilon=0 消失,二阶项由 E[εε]=I\mathbb E[\varepsilon\varepsilon^\top]=I 留下。利用 Dirac 导数在期望中的转移关系,可以得到

tpt(x)=x[ft(x)pt(x)]+12gt2Δxpt(x).(10)\partial_t p_t(x) =-\nabla_x\cdot\left[f_t(x)p_t(x)\right] +\frac12g_t^2\Delta_x p_t(x). \tag{10}

第一项描述 drift 搬运概率质量,第二项描述噪声导致的密度扩散。Fokker–Planck 方程把大量随机样本路径压缩成一个确定性的边缘密度演化规律。

同一组边缘分布对应一族随机过程

利用 pt=ptlogpt\nabla p_t=p_t\nabla\log p_t,对任意满足 0σt2gt20\le\sigma_t^2\le g_t^2 的函数,都可以把公式 (10) 改写为

tpt=[(ft12(gt2σt2)logpt)pt]+12σt2Δpt.(11)\begin{aligned} \partial_t p_t ={}&-\nabla\cdot\left[ \left(f_t-\frac12(g_t^2-\sigma_t^2) \nabla\log p_t\right)p_t \right] +\frac12\sigma_t^2\Delta p_t. \end{aligned} \tag{11}

因此下面整族前向 SDE 都产生完全相同的边缘分布 ptp_t

dx=[ft(x)12(gt2σt2)xlogpt(x)]dt+σtdw.(12)dx=\left[ f_t(x)-\frac12(g_t^2-\sigma_t^2) \nabla_x\log p_t(x) \right]dt +\sigma_t\,dw. \tag{12}

对公式 (12) 应用反向 SDE 公式,drift 中原有的修正项再减去 σt2logpt\sigma_t^2\nabla\log p_t,得到反向族:

dx=[ft(x)12(gt2+σt2)xlogpt(x)]dt+σtdwˉ.(13)dx=\left[ f_t(x)-\frac12(g_t^2+\sigma_t^2) \nabla_x\log p_t(x) \right]dt +\sigma_t\,d\bar w. \tag{13}

这里的“等价”仅指同一时刻的边缘分布相同;不同 σt\sigma_t 对应不同联合路径、不同样本相关性和不同数值性质。这正是连续时间版本的 DDIM 自由度。

概率流 ODE 与 DDIM

σt=0\sigma_t=0,公式 (12)(13) 都退化为同一个确定性 ODE:

dxdt=ft(x)12gt2xlogpt(x).(14)\frac{dx}{dt} =f_t(x)-\frac12g_t^2\nabla_x\log p_t(x). \tag{14}

它叫 probability flow ODE。虽然单条轨迹不再注入噪声,但其边缘密度仍与原 SDE 完全一致。用 sθs_\theta 代替真实 score 后,正向积分把数据确定性映射到先验,反向积分把先验确定性映射回数据。

这带来三层作用:

  • 可以使用成熟的自适应、高阶 ODE 求解器减少离散误差;
  • 确定性可逆轨迹提供稳定的 latent representation 与编辑路径;
  • 通过 ODE 的瞬时变量变换公式,可以像 continuous normalizing flow 一样计算似然。

在线性高斯扰动核下,把公式 (9)sθ=εθ/βˉts_\theta=-\varepsilon_\theta/\bar\beta_t 代入公式 (14),整理可得

ddt(xtαˉt)=εθ(xt,t)ddt(βˉtαˉt).(15)\frac{d}{dt}\left(\frac{x_t}{\bar\alpha_t}\right) =\varepsilon_\theta(x_t,t) \frac{d}{dt}\left( \frac{\bar\beta_t}{\bar\alpha_t} \right). \tag{15}

这正是确定性 DDIM 离散更新的连续极限。因此关系不是“DDIM 恰好像一个 ODE”,而是:DDIM 在高斯线性扰动路径上离散化了 probability flow ODE;原始随机采样则离散化了 reverse-time SDE。二者共享同一个 score 网络和同一组边缘分布,但沿着不同路径把噪声运回数据。