Skip to content
huc
Go back

扩散模型的反向方差估计

目录 1 / 7

《生成扩散模型漫谈(七):最优扩散方差估计(上)》 —— 苏剑林

《生成扩散模型漫谈(八):最优扩散方差估计(下)》 —— 苏剑林

固定 x0x_0 的随机性不是全部随机性

设正向扰动核为

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),

DDIM 给出的条件反向核可以写成

p(xt1xt,x0)=N(xt1;ctxt+γtx0,σt2I),(1)p(x_{t-1}\mid x_t,x_0) =\mathcal N\left( x_{t-1}; c_t x_t+\gamma_t x_0, \sigma_t^2I \right), \tag{1}

其中

ct=βˉt12σt2βˉt,γt=αˉt1ctαˉt.c_t=\frac{\sqrt{\bar\beta_{t-1}^2-\sigma_t^2}}{\bar\beta_t}, \qquad \gamma_t=\bar\alpha_{t-1}-c_t\bar\alpha_t.

DDIM 通常用单点估计 μˉt(xt)\bar\mu_t(x_t) 替换 x0x_0。这会保留公式 (1) 中显式的方差 σt2I\sigma_t^2I,却丢掉“给定 xtx_t 后,真实 x0x_0 仍不确定”带来的随机性。

严格的边缘化应是

p(xt1xt)=p(xt1xt,x0)p(x0xt)dx0.(2)p(x_{t-1}\mid x_t) =\int p(x_{t-1}\mid x_t,x_0)p(x_0\mid x_t)\,dx_0. \tag{2}

用各向同性高斯近似未知后验

p(x0xt)N(x0;μˉt(xt),σˉt2I),(3)p(x_0\mid x_t) \approx\mathcal N(x_0;\bar\mu_t(x_t),\bar\sigma_t^2I), \tag{3}

并写成 x0=μˉt(xt)+σˉtε2x_0=\bar\mu_t(x_t)+\bar\sigma_t\varepsilon_2。将它代入条件采样式:

xt1=ctxt+γtx0+σtε1ctxt+γtμˉt(xt)+γtσˉtε2+σtε1.\begin{aligned} x_{t-1} &=c_t x_t+\gamma_t x_0+\sigma_t\varepsilon_1\\ &\approx c_t x_t+\gamma_t\bar\mu_t(x_t) +\gamma_t\bar\sigma_t\varepsilon_2 +\sigma_t\varepsilon_1. \end{aligned}

ε1\varepsilon_1ε2\varepsilon_2 相互独立,所以两部分随机性的方差相加:

p(xt1xt)N(xt1;ctxt+γtμˉt(xt),(σt2+γt2σˉt2)I).(4)p(x_{t-1}\mid x_t) \approx\mathcal N\left( x_{t-1}; c_t x_t+\gamma_t\bar\mu_t(x_t), \left(\sigma_t^2+\gamma_t^2\bar\sigma_t^2\right)I \right). \tag{4}

公式 (4) 中的 γt2σˉt2\gamma_t^2\bar\sigma_t^2 是 Analytic-DPM 的核心修正。即便显式取 σt=0\sigma_t=0,只要 p(x0xt)p(x_0\mid x_t) 没有退化成单点,边缘化后的反向过程仍有非零方差。确定性 DDIM 相当于额外忽略了这部分后验不确定性。

均值仍由噪声预测器给出

平方损失下,后验均值是给定 xtx_t 后预测 x0x_0 的最优函数:

μˉt(xt)=E[x0xt]=argminμ(xt)E[x0μ(xt)2].(5)\bar\mu_t(x_t) =\mathbb E[x_0\mid x_t] =\arg\min_{\mu(x_t)} \mathbb E\left[ \|x_0-\mu(x_t)\|^2 \right]. \tag{5}

使用扩散模型常见的参数化

μˉt(xt)=1αˉt[xtβˉtεθ(xt,t)](6)\bar\mu_t(x_t) =\frac{1}{\bar\alpha_t} \left[x_t-\bar\beta_t\varepsilon_\theta(x_t,t)\right] \tag{6}

后,公式 (5) 就是噪声预测目标。Analytic-DPM 不改变这个均值网络,而是要在它训练完成后估计 p(x0xt)p(x_0\mid x_t) 中尚未被条件均值解释的方差。

用全方差公式估计各向同性后验

先假设公式 (6) 精确等于真实条件均值。对任意常向量 μ0\mu_0,条件协方差满足

Σ(xt)=E[(x0μˉt)(x0μˉt)xt]=E[(x0μ0)(x0μ0)xt](μˉtμ0)(μˉtμ0).(7)\begin{aligned} \Sigma(x_t) &=\mathbb E\left[ (x_0-\bar\mu_t)(x_0-\bar\mu_t)^\top \mid x_t \right]\\ &=\mathbb E\left[ (x_0-\mu_0)(x_0-\mu_0)^\top \mid x_t \right] -(\bar\mu_t-\mu_0)(\bar\mu_t-\mu_0)^\top. \end{aligned} \tag{7}

为了得到只依赖时间 tt 的单个标量方差,先对 xtx_t 求平均,再取协方差矩阵的迹并除以维度 dd

σˉt2=1dEx0μ021dEμˉt(xt)μ02.(8)\bar\sigma_t^2 =\frac1d\mathbb E\|x_0-\mu_0\|^2 -\frac1d\mathbb E\|\bar\mu_t(x_t)-\mu_0\|^2. \tag{8}

μ0=E[x0]\mu_0=\mathbb E[x_0] 时,第一项就是数据的平均逐维方差。公式 (8) 是全方差公式的迹版本:总方差等于“条件均值之间的方差”与“条件内剩余方差”之和。

这条估计直观,但需要统计数据方差和去噪均值。还可以把 μˉt\bar\mu_t 的噪声参数化直接代入,得到只依赖噪声网络输出的 Analytic-DPM 形式。

从噪声网络输出得到解析方差

由公式 (6),在 Perfect Mean 假设下有

Σ(xt)=1αˉt2E[(xtαˉtx0)(xtαˉtx0)xt]βˉt2αˉt2εθ(xt,t)εθ(xt,t).(9)\begin{aligned} \Sigma(x_t) ={}&\frac{1}{\bar\alpha_t^2} \mathbb E\left[ (x_t-\bar\alpha_t x_0)(x_t-\bar\alpha_t x_0)^\top \mid x_t \right]\\ &-\frac{\bar\beta_t^2}{\bar\alpha_t^2} \varepsilon_\theta(x_t,t)\varepsilon_\theta(x_t,t)^\top. \end{aligned} \tag{9}

xtx_t 求平均后,公式第一项中的嵌套期望可以换回 x0p~(x0)x_0\sim\tilde p(x_0)xtp(xtx0)x_t\sim p(x_t\mid x_0) 的联合采样。由于 xtαˉtx0=βˉtεx_t-\bar\alpha_t x_0=\bar\beta_t\varepsilon,其二阶矩就是 βˉt2I\bar\beta_t^2I,于是

Ext[Σ(xt)]=βˉt2αˉt2[IExt[εθ(xt,t)εθ(xt,t)]].(10)\mathbb E_{x_t}[\Sigma(x_t)] =\frac{\bar\beta_t^2}{\bar\alpha_t^2} \left[ I-\mathbb E_{x_t} \left[ \varepsilon_\theta(x_t,t) \varepsilon_\theta(x_t,t)^\top \right] \right]. \tag{10}

取迹并除以 dd,得到各向同性估计:

σˉt2=βˉt2αˉt2(11dExtεθ(xt,t)2)(11)\boxed{ \bar\sigma_t^2 =\frac{\bar\beta_t^2}{\bar\alpha_t^2} \left( 1-\frac1d\mathbb E_{x_t} \|\varepsilon_\theta(x_t,t)\|^2 \right) } \tag{11}

它可以在均值模型训练完成后,对每个 tt 采样一批 xtx_t 离线统计,不需要重新训练扩散模型。理论上 σˉt2βˉt2/αˉt2\bar\sigma_t^2\le\bar\beta_t^2/\bar\alpha_t^2;有限样本、模型误差或数值误差可能让括号中的估计越界,实际实现需要保证方差非负。

如果不再强迫所有维度共享同一方差,只取公式 (10) 的对角线,就得到向量方差:

σˉt2=βˉt2αˉt2(1Ext[εθ(xt,t)2]),(12)\bar{\boldsymbol\sigma}_t^2 =\frac{\bar\beta_t^2}{\bar\alpha_t^2} \left( \boldsymbol 1- \mathbb E_{x_t}[\varepsilon_\theta(x_t,t)^2] \right), \tag{12}

平方为逐元素平方。完整 d×dd\times d 协方差虽然也有公式,但图像维度下存储、预测和采样成本过高,对角近似是更现实的折中。

Imperfect Mean 下应估计残差二阶矩

上面的“总方差减去已解释方差”依赖 μˉt(xt)=E[x0xt]\bar\mu_t(x_t)=\mathbb E[x_0\mid x_t]。真实网络并不精确,若均值固定但可能有偏,最合适的高斯方差应直接从似然求出。

对各向同性近似 N(x0;μˉt(xt),σˉt2I)\mathcal N(x_0;\bar\mu_t(x_t),\bar\sigma_t^2I),平均负对数似然中与方差有关的部分是

L(σˉt2)=Ex0μˉt(xt)22σˉt2+d2logσˉt2.\mathcal L(\bar\sigma_t^2) =\frac{\mathbb E\|x_0-\bar\mu_t(x_t)\|^2} {2\bar\sigma_t^2} +\frac d2\log\bar\sigma_t^2.

σˉt2\bar\sigma_t^2 求导并令其为零:

σˉt2=1dEx0,xtx0μˉt(xt)2.(13)\bar\sigma_t^2 =\frac1d\mathbb E_{x_0,x_t} \|x_0-\bar\mu_t(x_t)\|^2. \tag{13}

xt=αˉtx0+βˉtεx_t=\bar\alpha_t x_0+\bar\beta_t\varepsilon 和公式 (6) 代入,可得

σˉt2=βˉt2αˉt2dEx0,εεεθ(αˉtx0+βˉtε,t)2(14)\boxed{ \bar\sigma_t^2 =\frac{\bar\beta_t^2}{\bar\alpha_t^2d} \mathbb E_{x_0,\varepsilon} \left\| \varepsilon- \varepsilon_\theta( \bar\alpha_t x_0+\bar\beta_t\varepsilon,t ) \right\|^2 } \tag{14}

Perfect Mean 时,残差只表示不可约的后验不确定性;Imperfect Mean 时,它还包含均值模型的系统误差。因此公式 (14) 估计的是“相对于当前固定均值,最优高斯近似应使用的残差二阶矩”,不能再解释成纯粹的真实条件方差。

逐维版本只需去掉维度平均:

σˉt2=βˉt2αˉt2Ex0,ε[(εεθ(xt,t))2].(15)\bar{\boldsymbol\sigma}_t^2 =\frac{\bar\beta_t^2}{\bar\alpha_t^2} \mathbb E_{x_0,\varepsilon} \left[ (\varepsilon-\varepsilon_\theta(x_t,t))^2 \right]. \tag{15}

条件方差需要学习后验残差

无条件的 σˉt2\bar{\boldsymbol\sigma}_t^2 对所有 xtx_t 做了平均。若不同带噪输入有不同的不确定性,应保留 xtx_t 依赖:

σˉt2(xt)=βˉt2αˉt2E[(εtεθ(xt,t))2xt],(16)\bar{\boldsymbol\sigma}_t^2(x_t) =\frac{\bar\beta_t^2}{\bar\alpha_t^2} \mathbb E\left[ (\varepsilon_t-\varepsilon_\theta(x_t,t))^2 \mid x_t \right], \tag{16}

其中 εt=(xtαˉtx0)/βˉt\varepsilon_t=(x_t-\bar\alpha_t x_0)/\bar\beta_t。条件期望无法为每个 xtx_t 重复采样真实后验,但仍可利用平方回归:训练网络 gϕ(xt,t)g_\phi(x_t,t) 预测逐元素目标

rt=(εtεθ(xt,t))2,r_t=(\varepsilon_t-\varepsilon_\theta(x_t,t))^2,

并最小化

Ex0,t,εt[rtgϕ(xt,t)2].(17)\mathbb E_{x_0,t,\varepsilon_t} \left[ \|r_t-g_\phi(x_t,t)\|^2 \right]. \tag{17}

平方损失的最优解是 gϕ(xt,t)=E[rtxt]g_\phi(x_t,t)=\mathbb E[r_t\mid x_t],乘上 βˉt2/αˉt2\bar\beta_t^2/\bar\alpha_t^2 后就是公式 (16)

为什么采用两阶段训练

均值决定反向一步往哪里走,方差只是决定在这个方向附近注入多少随机性。若从头联合训练,持续变化的均值会让方差目标本身不断移动,而可学习方差又会改变负对数似然对均值误差的加权,两者容易相互干扰。

Extended-Analytic-DPM 因此先用固定方差训练好均值/噪声网络,再冻结它,离线统计解析方差或训练条件方差头。这样做允许复用已有模型,也让公式 (14) 中的残差目标保持稳定。

方差修正在少步采样时通常更重要:跨步越大,单步条件分布越宽,忽略 x0xtx_0\mid x_t 的不确定性造成的误差越明显;接近完整细网格时,均值误差和数值离散误差往往仍是主要因素,精细方差更像次级改进。

还要保留一个经验上的边界:原文提到基于 Perfect Mean 假设的方案在实验中反而优于显式面向 Imperfect Mean 的残差方案。这不等于均值网络数学上真的完美,只说明估计偏差、有限样本方差和训练目标的组合可能让更严格的公式未必给出更好的生成指标。