《生成扩散模型漫谈(七):最优扩散方差估计(上)》 —— 苏剑林
《生成扩散模型漫谈(八):最优扩散方差估计(下)》 —— 苏剑林
固定 x0 的随机性不是全部随机性
设正向扰动核为
p(xt∣x0)=N(xt;αˉtx0,βˉt2I),
DDIM 给出的条件反向核可以写成
p(xt−1∣xt,x0)=N(xt−1;ctxt+γtx0,σt2I),(1)
其中
ct=βˉtβˉt−12−σt2,γt=αˉt−1−ctαˉt.
DDIM 通常用单点估计 μˉt(xt) 替换 x0。这会保留公式
(1) 中显式的方差 σt2I,却丢掉“给定 xt 后,真实 x0 仍不确定”带来的随机性。
严格的边缘化应是
p(xt−1∣xt)=∫p(xt−1∣xt,x0)p(x0∣xt)dx0.(2)
用各向同性高斯近似未知后验
p(x0∣xt)≈N(x0;μˉt(xt),σˉt2I),(3)
并写成 x0=μˉt(xt)+σˉtε2。将它代入条件采样式:
xt−1=ctxt+γtx0+σtε1≈ctxt+γtμˉt(xt)+γtσˉtε2+σtε1.
ε1、ε2 相互独立,所以两部分随机性的方差相加:
p(xt−1∣xt)≈N(xt−1;ctxt+γtμˉt(xt),(σt2+γt2σˉt2)I).(4)
公式 (4) 中的
γt2σˉt2 是 Analytic-DPM 的核心修正。即便显式取
σt=0,只要 p(x0∣xt) 没有退化成单点,边缘化后的反向过程仍有非零方差。确定性 DDIM 相当于额外忽略了这部分后验不确定性。
均值仍由噪声预测器给出
平方损失下,后验均值是给定 xt 后预测 x0 的最优函数:
μˉt(xt)=E[x0∣xt]=argμ(xt)minE[∥x0−μ(xt)∥2].(5)
使用扩散模型常见的参数化
μˉt(xt)=αˉt1[xt−βˉtεθ(xt,t)](6)
后,公式 (5) 就是噪声预测目标。Analytic-DPM 不改变这个均值网络,而是要在它训练完成后估计
p(x0∣xt) 中尚未被条件均值解释的方差。
用全方差公式估计各向同性后验
先假设公式 (6) 精确等于真实条件均值。对任意常向量 μ0,条件协方差满足
Σ(xt)=E[(x0−μˉt)(x0−μˉt)⊤∣xt]=E[(x0−μ0)(x0−μ0)⊤∣xt]−(μˉt−μ0)(μˉt−μ0)⊤.(7)
为了得到只依赖时间 t 的单个标量方差,先对 xt 求平均,再取协方差矩阵的迹并除以维度 d:
σˉt2=d1E∥x0−μ0∥2−d1E∥μˉt(xt)−μ0∥2.(8)
取 μ0=E[x0] 时,第一项就是数据的平均逐维方差。公式
(8) 是全方差公式的迹版本:总方差等于“条件均值之间的方差”与“条件内剩余方差”之和。
这条估计直观,但需要统计数据方差和去噪均值。还可以把
μˉt 的噪声参数化直接代入,得到只依赖噪声网络输出的 Analytic-DPM 形式。
从噪声网络输出得到解析方差
由公式 (6),在 Perfect Mean 假设下有
Σ(xt)=αˉt21E[(xt−αˉtx0)(xt−αˉtx0)⊤∣xt]−αˉt2βˉt2εθ(xt,t)εθ(xt,t)⊤.(9)
对 xt 求平均后,公式第一项中的嵌套期望可以换回
x0∼p~(x0)、xt∼p(xt∣x0) 的联合采样。由于
xt−αˉtx0=βˉtε,其二阶矩就是
βˉt2I,于是
Ext[Σ(xt)]=αˉt2βˉt2[I−Ext[εθ(xt,t)εθ(xt,t)⊤]].(10)
取迹并除以 d,得到各向同性估计:
σˉt2=αˉt2βˉt2(1−d1Ext∥εθ(xt,t)∥2)(11)
它可以在均值模型训练完成后,对每个 t 采样一批 xt 离线统计,不需要重新训练扩散模型。理论上
σˉt2≤βˉt2/αˉt2;有限样本、模型误差或数值误差可能让括号中的估计越界,实际实现需要保证方差非负。
如果不再强迫所有维度共享同一方差,只取公式
(10) 的对角线,就得到向量方差:
σˉt2=αˉt2βˉt2(1−Ext[εθ(xt,t)2]),(12)
平方为逐元素平方。完整 d×d 协方差虽然也有公式,但图像维度下存储、预测和采样成本过高,对角近似是更现实的折中。
Imperfect Mean 下应估计残差二阶矩
上面的“总方差减去已解释方差”依赖
μˉt(xt)=E[x0∣xt]。真实网络并不精确,若均值固定但可能有偏,最合适的高斯方差应直接从似然求出。
对各向同性近似
N(x0;μˉt(xt),σˉt2I),平均负对数似然中与方差有关的部分是
L(σˉt2)=2σˉt2E∥x0−μˉt(xt)∥2+2dlogσˉt2.
对 σˉt2 求导并令其为零:
σˉt2=d1Ex0,xt∥x0−μˉt(xt)∥2.(13)
把 xt=αˉtx0+βˉtε 和公式
(6) 代入,可得
σˉt2=αˉt2dβˉt2Ex0,εε−εθ(αˉtx0+βˉtε,t)2(14)
Perfect Mean 时,残差只表示不可约的后验不确定性;Imperfect Mean 时,它还包含均值模型的系统误差。因此公式
(14) 估计的是“相对于当前固定均值,最优高斯近似应使用的残差二阶矩”,不能再解释成纯粹的真实条件方差。
逐维版本只需去掉维度平均:
σˉt2=αˉt2βˉt2Ex0,ε[(ε−εθ(xt,t))2].(15)
条件方差需要学习后验残差
无条件的 σˉt2 对所有 xt 做了平均。若不同带噪输入有不同的不确定性,应保留 xt 依赖:
σˉt2(xt)=αˉt2βˉt2E[(εt−εθ(xt,t))2∣xt],(16)
其中
εt=(xt−αˉtx0)/βˉt。条件期望无法为每个 xt 重复采样真实后验,但仍可利用平方回归:训练网络 gϕ(xt,t) 预测逐元素目标
rt=(εt−εθ(xt,t))2,
并最小化
Ex0,t,εt[∥rt−gϕ(xt,t)∥2].(17)
平方损失的最优解是
gϕ(xt,t)=E[rt∣xt],乘上
βˉt2/αˉt2 后就是公式
(16)。
为什么采用两阶段训练
均值决定反向一步往哪里走,方差只是决定在这个方向附近注入多少随机性。若从头联合训练,持续变化的均值会让方差目标本身不断移动,而可学习方差又会改变负对数似然对均值误差的加权,两者容易相互干扰。
Extended-Analytic-DPM 因此先用固定方差训练好均值/噪声网络,再冻结它,离线统计解析方差或训练条件方差头。这样做允许复用已有模型,也让公式
(14) 中的残差目标保持稳定。
方差修正在少步采样时通常更重要:跨步越大,单步条件分布越宽,忽略 x0∣xt 的不确定性造成的误差越明显;接近完整细网格时,均值误差和数值离散误差往往仍是主要因素,精细方差更像次级改进。
还要保留一个经验上的边界:原文提到基于 Perfect Mean 假设的方案在实验中反而优于显式面向 Imperfect Mean 的残差方案。这不等于均值网络数学上真的完美,只说明估计偏差、有限样本方差和训练目标的组合可能让更严格的公式未必给出更好的生成指标。