DDPM
这篇主要记公式,不展开太多概念解释。主线只有三件事:
- 正向过程怎样逐步把数据加成高斯噪声。
- 反向过程为什么要学习 pθ(xt−1∣xt)。
- 训练目标最后为什么会落到预测噪声 ϵ。
正向加噪
最直接的想法是每一步都往前一个状态上加一点高斯噪声:
xt=xt−1+βϵ
相当于均值为 xt−1,方差为 β2 的分布中采样。但这样写不太好,因为均值始终围绕前一个的样本点,噪声强度的控制也不够自然。DDPM 的做法是给每一步定义一个噪声系数序列:
0<β1<β2<⋯<βT<1
常见会把 βt 设得很小,例如从 0.0001 到 0.02。
于是正向过程写成:
q(xt∣xt−1)=N(xt;1−βtxt−1,βtI)
也就是:
xt=1−βtxt−1+βtϵt,ϵt∼N(0,I)
记
αt=1−βt
则上式可以写成更常用的形式:
q(xt∣xt−1)=N(xt;αtxt−1,(1−αt)I)
xt=αtxt−1+1−αtϵt
继续展开:
x1=α1x0+1−α1ϵ1
x2=α2x1+1−α2ϵ2
⋯
记累计乘积
αˉt=s=1∏tαs
则可以得到闭式:
q(xt∣x0)=N(xt;αˉtx0,(1−αˉt)I)
也就是:
xt=αˉtx0+1−αˉtϵ,ϵ∼N(0,I)
所以正向过程的目标很明确:随着 t 变大,xt 会逐渐接近均值为 0、方差为 I 的标准高斯分布。
反向过程
如果正向过程是把 x0 一步步加噪到 xT,那生成时就希望反过来:从高斯噪声开始,一步步还原出数据。
理想情况下,我们想要的是真实反向后验:
q(xt−1∣xt)=q(xt)q(xt∣xt−1)q(xt−1)
但这个分布通常不好直接求,所以更常见的做法是考虑:
q(xt−1∣xt,x0)=q(xt∣x0)q(xt∣xt−1,x0)q(xt−1∣x0)
由于 Markov 性质,
q(xt∣xt−1,x0)=q(xt∣xt−1)
因此:
q(xt−1∣xt,x0)=q(xt∣x0)q(xt∣xt−1)q(xt−1∣x0)
这一项是可以算的,因为分子分母全都是高斯分布。
于是训练时就让神经网络去拟合这个真实后验:
pθ(xt−1∣xt)≈q(xt−1∣xt,x0)
生成时只要从
xT∼N(0,I)
开始,不断采样
pθ(xt−1∣xt)
就可以一步步走回数据分布。
训练目标
给定样本 x1,x2,…,xn,目标仍然是最大化训练数据的似然:
argθmaxi∏pθ(xi)⟺argθmaxi∑logpθ(xi)⟺argθmini∑−logpθ(xi)
正向过程定义为:
q(x1:T∣x0)=t=1∏Tq(xt∣xt−1)
反向生成过程定义为:
pθ(x0:T)=p(xT)t=1∏Tpθ(xt−1∣xt)
对单个样本 x0,从负对数似然开始:
−logpθ(x0)=−log∫pθ(x0:T)dx1:T
乘除同一个 q(x1:T∣x0):
−logpθ(x0)=−log∫q(x1:T∣x0)q(x1:T∣x0)pθ(x0:T)dx1:T=−logEq(x1:T∣x0)[q(x1:T∣x0)pθ(x0:T)]
由于 −log 是凸函数,由 Jensen 不等式可得:
−logpθ(x0)≤Eq(x1:T∣x0)[logpθ(x0:T)q(x1:T∣x0)]
后面继续展开会得到一个很长的式子,这里直接省略中间整理过程:
−logpθ(x0)≤Eq[logp(xT)q(xT∣x0)]+Eq[t=2∑Tlogpθ(xt−1∣xt)q(xt−1∣xt,x0)]−Eq[logpθ(x0∣x1)]
这里三项里:
- 第一项没有可训练参数。
- 第二项是每一步反向分布的拟合误差。
- 第三项通常可以单独处理,或者在这类笔记里先略掉。
所以训练时最核心的是最小化第二项:
Eq[t=2∑Tlogpθ(xt−1∣xt)q(xt−1∣xt,x0)]
把求和和期望拆开,省略推导过程,此项可写成:
t=2∑TEq(xt∣x0)[DKL(q(xt−1∣xt,x0)∥pθ(xt−1∣xt))]
所以 DDPM 的训练,本质上是在每个时间步都让模型学会把真实后验
q(xt−1∣xt,x0)
拟合成模型的反向分布
pθ(xt−1∣xt)
后验分布的均值和方差
因为
q(xt−1∣xt,x0)=q(xt∣x0)q(xt∣xt−1)q(xt−1∣x0)
而右边都是高斯分布,所以这个后验本身也一定是高斯分布:
q(xt−1∣xt,x0)=N(xt−1;μ~t(xt,x0),β~tI)
中间配方过程比较长,这里直接记结果。
方差是:
β~t=1−αˉt1−αˉt−1βt=1−αˉt(1−αt)(1−αˉt−1)
均值可以写成:
μ~t(xt,x0)=1−αˉtαˉt−1βtx0+1−αˉtαt(1−αˉt−1)xt
又因为
xt=αˉtx0+1−αˉtϵ
所以
x0=αˉt1(xt−1−αˉtϵ)
代回去之后,可以把均值改写成更常用的形式:
μ~t(xt,ϵ)=αt1(xt−1−αˉt1−αtϵ)
这一步很关键,因为它说明了:
只要网络能预测出噪声 ϵ,就能反推出从 xt 到 xt−1 所需的均值。
最后落到预测噪声
上面已经看到,反向分布的方差 β~t 可以直接由 schedule 算出来;而均值可以通过预测噪声来得到。
所以实际做法通常是让网络输出:
ϵθ(xt,t)
然后用它代替真实噪声 ϵ,得到反向过程的均值:
μθ(xt,t)=αt1(xt−1−αˉt1−αtϵθ(xt,t))
于是模型分布写成:
pθ(xt−1∣xt)=N(xt−1;μθ(xt,t),β~tI)
再进一步,DDPM 常用的简化训练目标就是直接做噪声回归:
Lsimple=Ex0,ϵ,t[∥ϵ−ϵθ(xt,t)∥2]
其中
xt=αˉtx0+1−αˉtϵ
这就是最后最常见的训练形式。
换句话说,DDPM 看起来是在学习“反向去噪分布”,但落到实现上,通常就是让网络学会:
给定某一步的 noisy sample xt 和时间步 t,预测当时被加进去的噪声 ϵ。