\[ \begin{align}\begin{aligned}\newcommand{\ba}{\boldsymbol{a}} \newcommand{\bb}{\boldsymbol{b}} \newcommand{\be}{\boldsymbol{e}} \newcommand{\bq}{\boldsymbol{q}} \newcommand{\bk}{\boldsymbol{k}} \newcommand{\bw}{\boldsymbol{w}} \newcommand{\bx}{\boldsymbol{x}} \newcommand{\by}{\boldsymbol{y}} \newcommand{\bz}{\boldsymbol{z}} \newcommand{\bd}{\boldsymbol{d}} \newcommand{\bv}{\boldsymbol{v}} \newcommand{\bs}{\boldsymbol{s}}\\\newcommand{\btheta}{\boldsymbol{\theta}} \newcommand{\bbeta}{\boldsymbol{\beta}} \newcommand{\bgamma}{\boldsymbol{\gamma}} \newcommand{\bsigma}{\boldsymbol{\sigma}} \newcommand{\md}{\mbox{d}} \newcommand{\bmu}{\boldsymbol{\mu}} \newcommand{\bone}{\boldsymbol{1}} \newcommand{\bzero}{\boldsymbol{0}} \newcommand{\bepsilon}{\boldsymbol{\epsilon}} \newcommand{\bphi}{\boldsymbol{\phi}} \newcommand{\bh}{\boldsymbol{h}} \newcommand{\bc}{\boldsymbol{c}} \newcommand{\br}{\boldsymbol{r}} \newcommand{\bQ}{\boldsymbol{Q}} \newcommand{\bK}{\boldsymbol{K}} \newcommand{\bV}{\boldsymbol{V}} \newcommand{\bSigma}{\boldsymbol{\Sigma}} \newcommand{\bg}{\boldsymbol{g}} \newcommand{\bxi}{\boldsymbol{\xi}} \newcommand{\bvarepsilon}{\boldsymbol{\varepsilon}} \newcommand{\bdelta}{\boldsymbol{\delta}} \newcommand{\bq}{\boldsymbol{q}} \newcommand{\bk}{\boldsymbol{k}} \newcommand{\bJ}{\boldsymbol{J}} \newcommand{\bp}{\boldsymbol{p}} \newcommand{\bi}{\boldsymbol{i}} \newcommand{\bo}{\boldsymbol{o}} \newcommand{\bE}{\boldsymbol{E}} \newcommand{\bH}{\boldsymbol{H}} \newcommand{\bL}{\boldsymbol{L}} \newcommand{\bu}{\boldsymbol{u}} \newcommand{\bLambda}{\boldsymbol{\Lambda}} \newcommand{\trans}{^{\rm\scriptsize T}} \newcommand{\var}{\mathrm{var}}\\\newcommand{\bA}{\boldsymbol{A}} \newcommand{\bB}{\boldsymbol{B}} \newcommand{\bC}{\boldsymbol{C}} \newcommand{\bD}{\boldsymbol{D}} \newcommand{\bG}{\boldsymbol{G}} \newcommand{\bI}{\boldsymbol{I}} \newcommand{\bM}{\boldsymbol{M}} \newcommand{\bP}{\boldsymbol{P}} \newcommand{\bS}{\boldsymbol{S}} \newcommand{\bU}{\boldsymbol{U}} \newcommand{\bW}{\boldsymbol{W}} \newcommand{\bX}{\boldsymbol{X}} \newcommand{\bY}{\boldsymbol{Y}} \newcommand{\bZ}{\boldsymbol{Z}} \newcommand{\cotp}{\textcolor[RGB]{48,209,88}{TP}} \newcommand{\cotn}{\textcolor[RGB]{100,210,255}{TN}} \newcommand{\cofp}{\textcolor[RGB]{94,92,230}{FP}} \newcommand{\cofn}{\textcolor[RGB]{191,90,242}{FN}}\\\newcommand{\numcotp}{\textcolor[RGB]{48,209,88}{50}} \newcommand{\numcotn}{\textcolor[RGB]{100,210,255}{30}} \newcommand{\numcofp}{\textcolor[RGB]{94,92,230}{10}} \newcommand{\numcofn}{\textcolor[RGB]{191,90,242}{10}} \DeclareMathOperator*{\argmin}{arg\,min}\end{aligned}\end{align} \]

去噪扩散概率模型(DDPM)#

学习目标与记号#

  1. 推导前向后验 \(q(\bx_{t-1}\mid\bx_t,\bx_0)\) 的均值和方差;

  2. 理解证据下界(ELBO,也称变分下界)与噪声预测损失之间的联系,并比较 \(\tilde{\bepsilon}_t\)⁠、\(\bx_0\)\(\bv\) 参数化;

  3. 实现基本 DDPM 训练与采样步骤,识别时间下标、方差和裁剪错误;

  本节沿用统一记号:普通小写字母表示标量,粗体小写字母表示向量,粗体大写字母表示矩阵或高阶张量;时间步写作下标 \(t\)⁠,网络层编号写作上标 \([l]\)⁠,转置写作 \(\trans\)⁠。用 \(m\) 表示一个批量中的图片数,用 \(d_C^{[l]}\)⁠、\(d_H^{[l]}\)\(d_W^{[l]}\) 分别表示第 \(l\) 层特征图的通道数、高和宽。\(\bx_0\) 表示干净数据,\(\bx_t\) 表示第 \(t\) 个噪声等级的随机变量;\(\bepsilon_t\) 表示从 \(\bx_{t-1}\)\(\bx_t\) 时单独加入的第 \(t\) 步噪声,\(\tilde{\bepsilon}_t\) 表示把前 \(t\) 步噪声按各自系数合并并标准化后得到、可用于从 \(\bx_0\) 一步构造 \(\bx_t\) 的累计等价噪声。正文与练习中的程序都应同时检查数值结果和数组维度。本节练习的参考答案见 DDPM 答案⁠。

  本节直接使用 扩散随机变量与逐步加噪参数解析前向采样⁠;理解训练与生成的方向差异时,可回看 训练路径与生成路径⁠。

  去噪扩散概率模型(denoising diffusion probabilistic model,DDPM) 1Ho、Jain 与 Abbeel(2020)⁠,Denoising Diffusion Probabilistic Models⁠。 是离散扩散模型的经典形式。它保留固定的前向高斯马尔可夫链,并用参数化高斯分布近似真实反向条件分布。

反向模型与可计算后验#

DDPM 的前向后验、噪声预测网络和反向采样步骤

图 41 DDPM 训练网络预测反向所需信息,并把标准高斯噪声逐步转换为数据样本#

  反向模型写为

\[p_{\btheta}(\bx_{0:T}) =p(\bx_T)\prod_{t=1}^{T}p_{\btheta}(\bx_{t-1}\mid\bx_t), \qquad p(\bx_T)=\mathcal{N}(\bzero,\bI),\]
\[p_{\btheta}(\bx_{t-1}\mid\bx_t) =\mathcal{N}\!\left( \bmu_{\btheta}(\bx_t,t), \bSigma_{\btheta}(\bx_t,t) \right).\]

  DDPM 常用一个带时间条件的神经网络预测解析前向采样中使用的累计等价噪声,并把该网络的输出记为 \(\bepsilon_{\btheta}(\bx_t,t)\)⁠。这里,\(\btheta\) 表示网络中需要通过训练确定的全部参数;第一个输入 \(\bx_t\) 是第 \(t\) 个噪声等级下的样本,第二个输入 \(t\) 告诉网络当前的噪声等级;输出 \(\bepsilon_{\btheta}(\bx_t,t)\)\(\bx_t\)\(\tilde{\bepsilon}_t\) 具有相同的维度。训练希望它接近构造当前 \(\bx_t\) 时实际使用的累计等价噪声,即

\[\bepsilon_{\btheta}(\bx_t,t) \approx \tilde{\bepsilon}_t,\]

其中 \(\tilde{\bepsilon}_t\) 来自前向采样式 \(\bx_t=\sqrt{\bar\alpha_t}\bx_0+\sqrt{1-\bar\alpha_t}\tilde{\bepsilon}_t\)⁠,并满足 \(\tilde{\bepsilon}_t\sim\mathcal{N}(\bzero,\bI)\)⁠,它是把 \(\bepsilon_1,\ldots,\bepsilon_t\) 按各步传播系数加权,并除以累计噪声的标准差后得到的等价标准高斯变量,不是把这些单步噪声直接相加,也不是第 \(t\) 步转移单独加入的 \(\bepsilon_t\)⁠。训练时,程序可以直接随机生成 \(\tilde{\bepsilon}_t\)⁠,因此它可以作为监督目标;生成新样本时并不知道该真实累计噪声,只能使用网络的预测结果。上面的近似关系表示训练目标,并不是对每个样本都严格成立的恒等式。

  直接计算 \(q(\bx_{t-1}\mid\bx_t)\) 需要未知的数据分布,但在额外给定 \(\bx_0\) 后,高斯乘积可得

\[q(\bx_{t-1}\mid\bx_t,\bx_0) =\mathcal{N}(\tilde{\bmu}_t,\tilde\beta_t\bI),\]

其中

\[\tilde\beta_t =\frac{1-\bar\alpha_{t-1}}{1-\bar\alpha_t}\beta_t,\]
\[\tilde{\bmu}_t =\frac{\sqrt{\bar\alpha_{t-1}}\beta_t}{1-\bar\alpha_t}\bx_0 +\frac{\sqrt{\alpha_t}(1-\bar\alpha_{t-1})}{1-\bar\alpha_t}\bx_t.\]

  这个后验为反向均值提供了监督目标。利用上面定义的噪声预测网络,由

\[\hat{\bx}_0 =\frac{\bx_t-\sqrt{1-\bar\alpha_t} \bepsilon_{\btheta}(\bx_t,t)} {\sqrt{\bar\alpha_t}},\]

可把预测的累计等价噪声换算成干净样本估计 \(\hat{\bx}_0\)⁠;把后验均值 \(\tilde{\bmu}_t\) 中的 \(\bx_0\) 替换为这个估计并化简,即得噪声预测形式的常用写法 \(\bmu_{\btheta}(\bx_t,t)\)⁠:

\[\bmu_{\btheta}(\bx_t,t) =\frac{1}{\sqrt{\alpha_t}} \left( \bx_t- \frac{\beta_t}{\sqrt{1-\bar\alpha_t}} \bepsilon_{\btheta}(\bx_t,t) \right).\]

  这里需要澄清两个记号的关系:\(\tilde{\bmu}_t\)\(\bmu_{\btheta}(\bx_t,t)\) 分别来自两个不同的分布。\(\tilde{\bmu}_t\)可计算后验 的均值:该后验是前向过程在额外给定 \(\bx_0\) 时的条件分布 \(q(\bx_{t-1}\mid\bx_t,\bx_0)=\mathcal{N}(\tilde{\bmu}_t,\tilde\beta_t\bI)\)⁠,它由加噪参数和 \(\bx_0,\bx_t\) 直接算出,是训练阶段可以对照的"理想目标",但因为依赖 \(\bx_0\)⁠,采样时无法使用。\(\bmu_{\btheta}(\bx_t,t)\) 则是 反向模型 的均值:该模型是 \(p_{\btheta}(\bx_{t-1}\mid\bx_t)=\mathcal{N}(\bmu_{\btheta}(\bx_t,t),\bSigma_{\btheta}(\bx_t,t))\)⁠,即网络在给定 \(\bx_t,t\) 时对 \(\bx_{t-1}\) 的预测,只依赖 \(\bx_t\)\(t\)⁠,采样时可以直接使用。DDPM 的训练目标正是让 \(\bmu_{\btheta}(\bx_t,t)\) 逼近 \(\tilde{\bmu}_t\)⁠:把 \(\tilde{\bmu}_t\) 中依赖 \(\bx_0\) 的项,用网络的累计噪声预测 \(\bepsilon_{\btheta}(\bx_t,t)\) 替换,即先通过 \(\hat{\bx}_0\) 换算成 \(\bx_0\) 的估计,就得到上面 \(\bmu_{\btheta}(\bx_t,t)\) 的常用形式。当噪声预测精确等于真实累计等价噪声,即 \(\bepsilon_{\btheta}(\bx_t,t)=\tilde{\bepsilon}_t\) 时,两者完全相等;二者之差完全来自噪声预测误差,这一结论的验证留作本节 综合练习 的习题。

证据下界(evidence lower bound,ELBO)与简化损失#

为什么需要一个下界?#

  概率生成模型不仅要生成外观合理的样本,还要给真实数据较高的概率。对一条已经观察到的干净数据 \(\bx_0\)⁠,这个概率写作 \(p_{\btheta}(\bx_0)\)⁠;它随模型参数 \(\btheta\) 改变。最大似然训练就是调整 \(\btheta\)⁠,使训练数据的平均对数似然尽可能大。如果对“似然”还不熟悉,可以先回看 线性回归与逻辑回归 中最大似然与负对数似然的介绍。

  DDPM 的生成过程不是一步得到 \(\bx_0\)⁠,而是从 \(\bx_T\) 出发,依次经过 \(\bx_{T-1},\ldots,\bx_1\)⁠,最后得到 \(\bx_0\)⁠。这些中间状态在训练数据中没有被观察到,因此属于隐藏变量。为简化表达,\(\bx_{1:T}\) 表示整个中间状态序列 \((\bx_1,\ldots,\bx_T)\)⁠,\(\bx_{0:T}\) 则还包括干净数据 \(\bx_0\)⁠。要求 \(\bx_0\) 的边缘概率,原则上需要把所有可能的中间路径都考虑进去:

\[p_{\btheta}(\bx_0) =\int p_{\btheta}(\bx_{0:T}) \,\mathrm{d}\bx_1\cdots\mathrm{d}\bx_T.\]

  对高维图像和很长的扩散链,这个积分包含数量极其庞大的可能路径,通常无法直接计算。困难不在于 \(\bx_0\) 没有观测到,而在于我们不知道生成 \(\bx_0\) 时可能经过哪些中间状态。ELBO 的作用就是把这个难以直接计算的对数似然,换成一个可以通过前向加噪样本来估计的下界。

从辅助分布到 ELBO#

  DDPM 已经给出了容易采样的前向过程 \(q(\bx_{1:T}\mid\bx_0)\)⁠。只要它在生成模型可能经过的路径上不为 0,就可以在上面的积分中乘以再除以同一个 \(q(\bx_{1:T}\mid\bx_0)\)⁠:

\[\begin{split}\begin{aligned} p_{\btheta}(\bx_0) &=\int q(\bx_{1:T}\mid\bx_0) \frac{p_{\btheta}(\bx_{0:T})} {q(\bx_{1:T}\mid\bx_0)} \,\mathrm{d}\bx_1\cdots\mathrm{d}\bx_T\\ &=\mathbb{E}_{q(\bx_{1:T}\mid\bx_0)} \left[ \frac{p_{\btheta}(\bx_{0:T})} {q(\bx_{1:T}\mid\bx_0)} \right]. \end{aligned}\end{split}\]

  这一步没有近似,只是把积分改写成关于已知前向过程的条件期望;这里的条件期望可以理解为:给定一条干净数据 \(\bx_0\)⁠,按照前向过程随机生成许多条加噪路径,计算括号中的概率比值,再对结果取平均。令随机变量

\[R(\bx_{1:T}) =\frac{p_{\btheta}(\bx_{0:T})} {q(\bx_{1:T}\mid\bx_0)},\]

\(p_{\btheta}(\bx_0)=\mathbb{E}_q[R(\bx_{1:T})]\)⁠。对两边取对数后,得到 \(\log p_{\btheta}(\bx_0)=\log\mathbb{E}_q[R(\bx_{1:T})]\)⁠。由于对数函数的凸性,Jensen 不等式说明“先求平均再取对数”不小于“先取对数再求平均”,即

\[\begin{split}\begin{aligned} \log p_{\btheta}(\bx_0) &=\log\mathbb{E}_{q(\bx_{1:T}\mid\bx_0)}[R(\bx_{1:T})]\\ &\geq \mathbb{E}_{q(\bx_{1:T}\mid\bx_0)}[\log R(\bx_{1:T})]\\ &= \mathbb{E}_{q(\bx_{1:T}\mid\bx_0)} \left[ \log p_{\btheta}(\bx_{0:T}) -\log q(\bx_{1:T}\mid\bx_0) \right]\\ &\equiv \mathcal{L}_{\mathrm{ELBO}}(\bx_0), \end{aligned}\end{split}\]

其中,ELBO 是 evidence lower bound 的缩写,中文通常称为“证据下界”或“变分下界”。“证据”指已经观察到的数据 \(\bx_0\)⁠,“下界”表示 \(\mathcal{L}_{\mathrm{ELBO}}(\bx_0)\) 不会超过数据的对数似然 \(\log p_{\btheta}(\bx_0)\)⁠。在变分推断的表述中,\(q(\bx_{1:T}\mid\bx_0)\) 称为变分分布;在 DDPM 中,它就是预先设定的前向加噪过程,不需要与反向网络一起学习。

Jensen 不等式在这里做了什么?

  假设一个正随机变量以相同概率取 \(1\)\(9\)⁠。“先平均再取对数”得到 \(\log[(1+9)/2]=\log 5\)⁠;“先取对数再平均”得到 \((\log 1+\log 9)/2=\log 3\)⁠。因为 \(\log 5>\log 3\)⁠,后者是前者的下界。ELBO 使用的是同样的关系,只是随机变量换成了路径概率之比 \(R(\bx_{1:T})\)⁠。

  为了看清下界何时接近真实对数似然,可以把两者之差继续整理。根据条件概率公式,\(p_{\btheta}(\bx_{1:T}\mid\bx_0)=p_{\btheta}(\bx_{0:T})/p_{\btheta}(\bx_0)\)⁠,因此

\[\begin{split}\begin{aligned} &\log p_{\btheta}(\bx_0) -\mathcal{L}_{\mathrm{ELBO}}(\bx_0)\\ &\quad= \mathbb{E}_{q(\bx_{1:T}\mid\bx_0)} \left[ \log \frac{q(\bx_{1:T}\mid\bx_0)} {p_{\btheta}(\bx_{0:T})} +\log p_{\btheta}(\bx_0) \right]\\ &\quad= \mathbb{E}_{q(\bx_{1:T}\mid\bx_0)} \left[ \log \frac{q(\bx_{1:T}\mid\bx_0)} {p_{\btheta}(\bx_{1:T}\mid\bx_0)} \right]\\ &\quad= D_{\mathrm{KL}}\!\left( q(\bx_{1:T}\mid\bx_0) \,\middle\|\, p_{\btheta}(\bx_{1:T}\mid\bx_0) \right) \geq 0. \end{aligned}\end{split}\]

这个 KL 散度可以理解为:当路径按照 \(q(\bx_{1:T}\mid\bx_0)\) 产生时,若改用 \(p_{\btheta}(\bx_{1:T}\mid\bx_0)\) 描述这些路径,会产生多大的平均对数概率差异。它总是不小于 0,并且只有两个分布相同时才等于 0。因此,ELBO 必然是下界;当 \(q(\bx_{1:T}\mid\bx_0)\) 恰好等于当前生成模型所确定的精确路径后验 \(p_{\btheta}(\bx_{1:T}\mid\bx_0)\) 时,ELBO 就等于当前模型给出的精确对数似然 \(\log p_{\btheta}(\bx_0)\)⁠。这里的“精确”是指该后验由当前的模型参数 \(\btheta\) 按照贝叶斯公式确定,并不表示它一定等于现实世界中的真实分布。两个路径分布的差异越大,下界通常就越松。需要注意的是,KL 散度一般不满足对称性,所以它不是通常几何意义下的距离。

例 8.1:只有两个隐藏状态时的 ELBO

  假设观察结果 \(x\) 可能通过两个没有被直接观察到的状态 \(z=A\)\(z=B\) 产生,并且模型给出的联合概率分别为 \(p_{\btheta}(x,A)=0.18\)\(p_{\btheta}(x,B)=0.02\)⁠。把两个可能状态相加,容易得到 \(p_{\btheta}(x)=0.20\)⁠,所以真实对数似然为 \(\log 0.20\approx-1.609\)⁠。

  先选择一个简单的辅助分布 \(q(A\mid x)=q(B\mid x)=0.5\)⁠。相应的 ELBO 为

\[ \begin{align}\begin{aligned}\begin{split} \begin{aligned} \mathcal{L}_{\mathrm{ELBO}}(x) &=0.5\log\frac{0.18}{0.5} +0.5\log\frac{0.02}{0.5}\\ &=\log 0.12 \approx-2.120. \end{aligned}\end{split}\\这个数确实小于真实对数似然 :math:`-1.609`。如果改用真实后验 :math:`q(A\mid x)=0.18/0.20=0.9`、:math:`q(B\mid x)=0.02/0.20=0.1`,则\end{aligned}\end{align} \]
\[\mathcal{L}_{\mathrm{ELBO}}(x) =0.9\log\frac{0.18}{0.9} +0.1\log\frac{0.02}{0.1} =\log 0.20.\]

此时 ELBO 与真实对数似然完全相等。DDPM 的中间路径远比两个状态复杂,但“引入容易处理的 \(q\)⁠,得到可计算下界”这一基本思路相同。

DDPM 中怎样分解负 ELBO?#

  增大 ELBO 与减小它的相反数完全等价,所以程序通常把负 ELBO 当作损失函数。对单个干净样本,先由 ELBO 的定义写出

\[-\mathcal{L}_{\mathrm{ELBO}}(\bx_0) =\mathbb{E}_{q(\bx_{1:T}\mid\bx_0)} \left[ \log q(\bx_{1:T}\mid\bx_0) -\log p_{\btheta}(\bx_{0:T}) \right].\]

  前向过程与反向生成过程都是马尔可夫链,因此它们的联合分布分别可以分解为

\[q(\bx_{1:T}\mid\bx_0) =\prod_{t=1}^{T}q(\bx_t\mid\bx_{t-1}),\]
\[p_{\btheta}(\bx_{0:T}) =p(\bx_T) \prod_{t=1}^{T} p_{\btheta}(\bx_{t-1}\mid\bx_t).\]

  若直接把这两个乘积代入,仍然不容易看出每个反向步骤应当逼近什么。对 \(t\geq2\)⁠,利用 Bayes 公式和前向过程的马尔可夫性质,可以把单步前向转移改写为

\[q(\bx_t\mid\bx_{t-1}) =q(\bx_{t-1}\mid\bx_t,\bx_0) \frac{q(\bx_t\mid\bx_0)} {q(\bx_{t-1}\mid\bx_0)}.\]

上式需要同时使用 Bayes 公式与前向过程的马尔可夫性质。把 \(t=2,\ldots,T\) 的这些等式相乘时,分母中的 \(q(\bx_{t-1}\mid\bx_0)\) 会与相邻时间步分子中的同类项依次相消。于是,前向路径概率与反向联合概率的比值可以整理为

\[\frac{q(\bx_{1:T}\mid\bx_0)} {p_{\btheta}(\bx_{0:T})} =\frac{q(\bx_T\mid\bx_0)}{p(\bx_T)} \frac{1}{p_{\btheta}(\bx_0\mid\bx_1)} \prod_{t=2}^{T} \frac{q(\bx_{t-1}\mid\bx_t,\bx_0)} {p_{\btheta}(\bx_{t-1}\mid\bx_t)}.\]

上面两个结果的验证合并在本节 综合练习 的同一道题中。

  最后,对上式取对数并关于前向过程求期望,乘积就变成求和,单个样本的负 ELBO 因而分解为

\[\begin{split}\begin{aligned} -\mathcal{L}_{\mathrm{ELBO}}(\bx_0) ={}&D_{\mathrm{KL}}\!\left( q(\bx_T\mid\bx_0)\,\middle\|\,p(\bx_T) \right)\\ &+\sum_{t=2}^{T} \mathbb{E}_{q(\bx_t\mid\bx_0)} \left[ D_{\mathrm{KL}}\!\left( q(\bx_{t-1}\mid\bx_t,\bx_0) \,\middle\|\, p_{\btheta}(\bx_{t-1}\mid\bx_t) \right) \right]\\ &-\mathbb{E}_{q(\bx_1\mid\bx_0)} \left[ \log p_{\btheta}(\bx_0\mid\bx_1) \right]. \end{aligned}\end{split}\]

第一项比较加噪终点与预先选定的高斯先验;当加噪参数和先验固定时,它不含需要学习的参数 \(\btheta\)⁠。中间的求和项逐步比较可计算的真实后验 \(q(\bx_{t-1}\mid\bx_t,\bx_0)\) 与神经网络给出的反向分布 \(p_{\btheta}(\bx_{t-1}\mid\bx_t)\)⁠,是学习反向去噪过程的核心。最后一项衡量从 \(\bx_1\) 恢复 \(\bx_0\) 的能力,通常称为重构项。

  实际训练会先从训练集取出一批 \(\bx_0\)⁠,再对这些样本的负 ELBO 求平均。这样做的含义是:既让最终加噪状态接近简单先验,又让每一个反向步骤接近相应的可计算后验,同时保证最后一步能够恢复干净数据。

怎样理解 ELBO?

  ELBO 不是一种新的网络结构或采样方法,而是用来训练概率模型的目标。真实对数似然难以直接计算时,我们先优化它的一个可计算下界。在 DDPM 中,固定的前向加噪过程提供了构造这个下界所需的中间状态分布,而神经网络负责学习下界中的反向条件分布。

为什么最后变成预测噪声?#

  负 ELBO 中 \(t=2,\ldots,T\) 的每个中间项都在比较两个高斯分布。真实后验的均值和方差分别记为 \(\tilde{\bmu}_t\)\(\tilde\beta_t\bI\)⁠,反向模型的均值和方差分别记为 \(\bmu_{\btheta}(\bx_t,t)\)\(\sigma_t^2\bI\)⁠。若 \(\sigma_t^2\) 预先给定、不随 \(\btheta\) 改变,则高斯 KL 散度中与 \(\btheta\) 有关的部分为

\[D_{\mathrm{KL}}\!\left( q(\bx_{t-1}\mid\bx_t,\bx_0) \,\middle\|\, p_{\btheta}(\bx_{t-1}\mid\bx_t) \right) =\frac{1}{2\sigma_t^2} \left\| \tilde{\bmu}_t -\bmu_{\btheta}(\bx_t,t) \right\|_2^2 +C_t,\]

其中 \(C_t\) 收集两个高斯方差产生的项,并且与需要训练的参数 \(\btheta\) 无关。因此,优化这个 KL 项的关键是让两个高斯分布的均值尽量接近。上述高斯 KL 散度的化简过程留作本节 综合练习 的习题。

  利用 \(\bx_t=\sqrt{\bar\alpha_t}\bx_0+\sqrt{1-\bar\alpha_t}\tilde{\bepsilon}_t\) 消去真实后验均值中的 \(\bx_0\)⁠,并把反向模型的均值写成噪声预测形式,可以分别得到

\[\begin{split}\begin{aligned} \tilde{\bmu}_t &=\frac{1}{\sqrt{\alpha_t}} \left( \bx_t- \frac{\beta_t}{\sqrt{1-\bar\alpha_t}} \tilde{\bepsilon}_t \right),\\ \bmu_{\btheta}(\bx_t,t) &=\frac{1}{\sqrt{\alpha_t}} \left( \bx_t- \frac{\beta_t}{\sqrt{1-\bar\alpha_t}} \bepsilon_{\btheta}(\bx_t,t) \right). \end{aligned}\end{split}\]

两式相减,便有

\[\bmu_{\btheta}(\bx_t,t) -\tilde{\bmu}_t =\frac{\beta_t} {\sqrt{\alpha_t}\sqrt{1-\bar\alpha_t}} \left[ \tilde{\bepsilon}_t -\bepsilon_{\btheta}(\bx_t,t) \right].\]

  这说明“让反向均值接近真实后验均值”与“让网络预测的累计等价噪声接近实际累计等价噪声”本质上是同一个要求,只是二者相差一个由时间步决定的比例系数。把上式代回高斯 KL 散度,并对训练样本、时间步和前向噪声取期望,便得到

\[\mathbb{E} \left[ w_t \left\| \tilde{\bepsilon}_t -\bepsilon_{\btheta}(\bx_t,t) \right\|_2^2 \right] +C,\]

其中

\[w_t =\frac{\beta_t^2} {2\sigma_t^2\alpha_t(1-\bar\alpha_t)} >0,\]

\(C\)\(\btheta\) 无关。权重 \(w_t\) 说明完整负 ELBO 会让不同时间步以不同强度影响训练。经典简化目标省略这个随时间变化的权重,并常写为

\[\mathcal{L}_{\mathrm{simple}} =\mathbb{E}_{\bx_0,t,\tilde{\bepsilon}_t} \left[ \left\| \tilde{\bepsilon}_t -\bepsilon_{\btheta}(\bx_t,t) \right\|_2^2 \right].\]

  简化损失与完整负 ELBO 关系紧密,但不同时间步的权重已经改变,因而不能简单地说两者在数值上完全相等。实践中还可预测 \(\bx_0\)⁠。另一种常见的 \(\bv\) 参数化定义为

\[\bv_t =\sqrt{\bar\alpha_t}\tilde{\bepsilon}_t -\sqrt{1-\bar\alpha_t}\bx_0.\]

不同参数化可以相互转换,但具有不同的优化尺度和误差分布。

时间条件网络#

  前面定义的 DDPM 噪声预测网络 记为 \(\bepsilon_{\btheta}(\bx_t,t)\)⁠:它接收含噪样本 \(\bx_t\) 和时间步 \(t\)⁠,输出与 \(\bx_t\) 同维度的累计等价噪声预测。从整体上看,这一模型通常是带时间条件的 U-Net:左侧(编码器)通过残差块与下采样逐步降低特征图分辨率并扩大感受野,同时把时间特征注入各个残差块;在最低或较低分辨率处处理图像的整体结构,并可加入自注意力;右侧(解码器)通过上采样逐步恢复分辨率,并用跳跃连接取回左侧保存的同分辨率细节。下面依次介绍该模型的各个模块:时间步条件、时间嵌入与时间特征、多尺度 U-Net 结构,以及自注意力。

为什么需要告诉网络当前时间步?#

  噪声预测网络的主要输入是当前的含噪样本 \(\bx_t\)⁠,但只给出 \(\bx_t\) 还不够。前向扩散在不同时间步使用不同的累计噪声强度:当 \(t\) 较小时,\(\bx_t\) 仍保留大量原始图像信息,网络主要需要识别并去除较弱的噪声;当 \(t\) 较大时,图像结构已经十分模糊,网络必须更多地依靠从训练数据中学到的形状、纹理和物体结构来推测应当保留的内容。因此,同一个网络在不同时间步面对的任务难度和处理方式并不相同。

  时间步 \(t\) 可以理解为当前样本所处的“噪声等级编号”。它不是神经网络已经训练了多少轮,也不是现实世界中的时间。把 \(t\) 或与它对应的连续噪声尺度明确提供给网络,可以避免网络仅凭 \(\bx_t\) 猜测当前噪声有多强。前面定义的 DDPM 噪声预测网络 因而写成 \(\bepsilon_{\btheta}(\bx_t,t)\)⁠,其输出用于估计与当前 \(\bx_t\) 对应的累计等价噪声 \(\tilde{\bepsilon}_t\)⁠。

例 8.2:不同噪声等级需要不同的处理方式

  假设输入图像中原本有一只熊猫。当 \(t\) 较小时,耳朵、眼睛和身体轮廓仍然清楚,网络只需去除细小的随机斑点;当 \(t\) 较大时,图像可能只剩下模糊的明暗区域,网络需要先判断整体结构,再逐步恢复局部细节。若不告诉网络当前的 \(t\)⁠,它就需要同时猜测“噪声有多强”和“噪声是什么”,学习任务会明显变得困难。

怎样把一个时间步变成网络能够使用的信息?#

  时间步 \(t\) 只是一个标量,而卷积网络的中间特征通常包含很多通道。为了让不同网络层都能方便地使用时间信息,常先把 \(t\) 转换为一个多维的 时间嵌入向量 (time embedding)。一种常见方法是先令 \(\tau_t=t/T\)⁠,再使用不同频率的正弦函数和余弦函数构造

\[\be_t =\left( \sin(\omega_1\tau_t),\cos(\omega_1\tau_t), \ldots, \sin(\omega_{d_t/2}\tau_t),\cos(\omega_{d_t/2}\tau_t) \right)^{\trans},\]

其中,\(d_t\) 是时间嵌入的维度,并假设它为偶数;\(\omega_1,\ldots,\omega_{d_t/2}\) 是从较低频率到较高频率的一组数。低频分量反映时间步的大范围变化,高频分量帮助区分相近的时间步。这种表示会使相近时间步得到相近但不完全相同的向量,也能让网络表示较大的时间范围。它与 Transformer 中的位置编码形式相似,但在这里编码的是噪声等级,而不是词语在句子中的位置。

  得到 \(\be_t\) 后,通常再把它送入 多层感知机⁠。例如,两层变换可写为

\[\bh_t =\phi\!\left( \bW_2\cdot \phi\!\left(\bW_1\cdot\be_t+\bb_1\right) +\bb_2 \right),\]

其中,\(\phi\) 表示激活函数,\(\bh_t\) 是经过学习后得到的时间特征。正弦和余弦编码本身没有需要训练的参数;MLP 中的权重和偏置则会在训练过程中得到更新,使时间特征更适合当前的噪声预测任务。

  U-Net 框架 中通常有多个残差块。残差块先对输入特征进行卷积和非线性变换,再把变换结果与原输入相加,因此可以理解为在原有特征上学习一个“修正量”。对于第 \(l\) 个残差块,可先用该块的权重 \(\bW^{[l]}_t\) 和偏置 \(\bb^{[l]}_t\) 对时间特征 \(\bh_t\) 做线性变换,得到时间向量 \(\br^{[l]}_t\)⁠:

\[\br^{[l]}_t =\bW^{[l]}_t\cdot\bh_t+\bb^{[l]}_t,\]

它的维度与该块通道数 \(d_C^{[l]}\) 一致,再整理为 \(d_C^{[l]}\times1\times1\) 的形式,并加到该块各个空间位置的特征上。这样,同一个卷积块在接收不同 \(t\) 时会产生不同的处理结果。有些实现还会让时间向量同时控制特征的缩放和偏移;基本目的仍然相同,即让各层都知道当前需要处理多强的噪声。

多尺度 U-Net 怎样处理图像?#

  图像 DDPM 常使用 U-Net 框架 作为噪声预测网络。不熟悉其基本结构的读者,可先通过该链接了解编码器、瓶颈层、解码器和同尺度跳跃连接。这里的“U”来自网络结构图的大致形状:左侧逐步降低特征图的空间分辨率,右侧再逐步恢复分辨率,中间连接最低分辨率的特征。假设输入 \(\bx_t\) 的大小为 \(d_C^{[l]}\times d_H^{[l]}\times d_W^{[l]}\)⁠,网络通常先用卷积把它转换为包含更多通道的特征图,然后依次完成以下处理。

  1. 下采样。 特征图的高和宽逐步减小,例如从 \(d_H^{[l]}\times d_W^{[l]}\) 变为 \(d_H^{[l]}/2\times d_W^{[l]}/2\)⁠,同时通道数 \(d_C^{[l]}\) 通常增加。分辨率降低后,一个特征位置能够综合更大图像区域的信息,因此网络更容易判断物体的大致位置、整体轮廓和不同区域之间的关系。这里所说的“提取全局上下文”,就是让网络不仅观察一个很小的局部区域,还能利用图像中较远位置的信息。

  2. 中间层。 在最低或较低分辨率处,网络继续使用残差块处理已经汇总的特征。由于空间位置较少,此处适合学习图像的整体结构,也常加入自注意力模块。

  3. 上采样。 网络逐步提高特征图的高和宽,把低分辨率的整体信息重新转换为与输入相同的空间分辨率。在这一过程中,网络逐渐恢复边缘、纹理和其他局部细节。

  4. 跳跃连接。 每次下采样前保存相应分辨率的特征;上采样回到同一分辨率时,再把保存的特征与当前特征拼接或相加。这样可以把早期层保留的精确位置和局部细节直接传给右侧网络,避免这些信息在连续下采样中完全丢失。U-Net 的跳跃连接连接左右两侧的同分辨率特征;它与残差块内部把输入加到输出上的残差连接作用不同,不应混为一谈。

  U-Net 的最终输出与 \(\bx_t\) 具有相同的维度。对参数化的噪声预测而言,每一个输出位置都对应输入中的一个位置,表示网络对累计等价噪声 \(\tilde{\bepsilon}_t\) 相应分量的估计。U-Net 并不是在一次前向计算中直接产生完整的最终图像;采样程序会反复调用它,依次完成从 \(\bx_T\)\(\bx_{T-1}\)⁠、再到 \(\bx_{T-2}\) 的更新,直至得到 \(\bx_0\)⁠。

为什么还会加入自注意力?#

  普通卷积首先汇总邻近像素的信息。随着网络加深,一个位置能够间接利用更远区域的特征,但这种信息传递需要经过多层。自注意力(self-attention)则允许一个空间位置直接根据相关程度汇总其他位置的信息。例如,在生成左右对称的眼睛、重复纹理或彼此相隔较远但属于同一物体的区域时,自注意力有助于保持这些区域的一致性。

  若特征图包含 \(d_H^{[l]}d_W^{[l]}\) 个空间位置,直接比较所有位置两两之间的关系会需要大约 \((d_H^{[l]}d_W^{[l]})^2\) 次配对计算,分辨率越高,计算和内存开销越大。因此,实际模型通常只在较低或中等分辨率处加入自注意力,而主要依靠卷积处理高分辨率的局部细节。是否使用自注意力以及在哪些分辨率使用,是网络结构的设计选择,并不是 DDPM 概率公式本身的必要条件。

从输入到输出:一次噪声预测包含哪些步骤?

  1. 输入当前含噪样本 \(\bx_t\)⁠,并把时间步 \(t\) 转换为时间嵌入。

  2. 通过左侧残差块和下采样逐渐汇总更大范围的图像信息,同时把时间特征送入各个残差块。

  3. 在较低分辨率处处理整体结构,并按需要使用自注意力联系相距较远的区域。

  4. 通过上采样逐步恢复空间分辨率,并利用跳跃连接取回左侧保存的位置和细节信息。

  5. 输出与 \(\bx_t\) 同维度的 \(\bepsilon_{\btheta}(\bx_t,t)\)⁠,将其作为累计等价噪声 \(\tilde{\bepsilon}_t\) 的估计,供反向采样公式使用。

DDPM 噪声预测网络从输入到输出的完整流程图

图 42 \(\bx_t\) 与时间步 \(t\) 出发,依次经过时间嵌入、U-Net 编码器(残差块与下采样)、瓶颈(自注意力)与解码器(上采样、跳跃连接),最终输出噪声预测 \(\bepsilon_{\btheta}(\bx_t,t)\)⁠。图中 灰色 部分表示与标准 U-Net 一致的结构(U 形编码器—解码器、同尺度跳跃连接等);彩色 部分为扩散模型相对标准 U-Net 新增的内容:蓝色表示时间条件(\(t\to\be_t\to\bh_t\)⁠)与各阶段的残差块,紫色表示瓶颈层的自注意力,橙色表示噪声预测输出 \(\bepsilon_{\btheta}(\bx_t,t)\approx\tilde{\bepsilon}_t\)#

基本采样步骤#

  反向采样从接近标准高斯的 \(\bx_T\sim\mathcal{N}(\bzero,\bI)\) 出发,按数学时间 \(t=T,T-1,\ldots,1\) 逐次更新。除最后一步外,每一步都从反向模型给出的 预测高斯 中随机抽取 \(\bx_{t-1}\)⁠。这里的"预测高斯"指条件分布 \(p_{\btheta}(\bx_{t-1}\mid\bx_t)=\mathcal{N}(\bmu_{\btheta}(\bx_t,t),\tilde\beta_t\bI)\)⁠,其中均值 \(\bmu_{\btheta}(\bx_t,t)\) 由网络预测的累计等价噪声按前面推导的常用形式换算而来,方差取已知的后验方差 \(\tilde\beta_t\)⁠;"采样"则指先取该高斯的均值,再加一个尺度为 \(\sqrt{\tilde\beta_t}\) 的标准高斯随机扰动,即

\[\bx_{t-1} =\bmu_{\btheta}(\bx_t,t) +\sqrt{\tilde\beta_t}\,\bz, \qquad \bz\sim\mathcal{N}(\bzero,\bI).\]

最后一步,即数学时间 \(t=1\)⁠、下面代码中的数组位置 0,通常不再加入随机噪声,直接以均值作为干净的最终输出 \(\bx_0\)⁠,避免在生成结果中残留一层噪点。下面的 ddpm_step 实现一次从 \(\bx_t\)\(\bx_{t-1}\) 的更新。它接收当前带噪样本 xt数组中的 0 基时间位置 txt 同形的模型噪声预测 eps_pred记录各时间步系数的 betaalphaalpha_bar 数组,以及 NumPy 随机数生成器,返回一个与 xt 同形的上一步样本。这里的 eps_pred 对应 \(\bepsilon_{\btheta}(\bx_t,t)\)⁠,即网络对累计等价噪声 \(\tilde{\bepsilon}_t\) 的预测,而不是新加入的第 \(t\) 步单步噪声 \(\bepsilon_t\)⁠:

import numpy as np

def ddpm_step(xt, t, eps_pred, beta, alpha, alpha_bar, rng):
    if not (0 <= t < len(beta)):
        raise IndexError("t 超出参数数组范围")
    if eps_pred.shape != xt.shape:
        raise ValueError("噪声预测必须与 xt 同形")
    if not (0.0 < alpha[t] < 1.0 and 0.0 < alpha_bar[t] < 1.0):
        raise ValueError("本函数约定 alpha 与 alpha_bar 位于 (0, 1)")
    mean = (xt - beta[t] / np.sqrt(1.0 - alpha_bar[t]) * eps_pred)
    mean = mean / np.sqrt(alpha[t])
    if t == 0:
        return mean
    alpha_bar_prev = alpha_bar[t - 1]
    posterior_var = beta[t] * (1.0 - alpha_bar_prev) / (1.0 - alpha_bar[t])
    if posterior_var < -1e-12:
        raise ValueError("后验方差为负,请检查时间下标和加噪参数")
    noise = rng.normal(size=xt.shape)
    return mean + np.sqrt(max(posterior_var, 0.0)) * noise

  程序先用网络预测的累计等价噪声计算反向高斯的均值 meant == 0这已是数组约定下的最后去噪步,因此直接返回均值;否则计算后验方差,再加入相应尺度的高斯噪声。max 只用来消除浮点误差导致的极小负数;明显为负时会报错。

  这个函数只是单步采样器:它不运行噪声预测网络,也不包含从纯噪声到样本的完整时间循环。这里的数组使用 Python 的 0 基下标:数组位置 \(t=0\) 对应数学中的第一步。项目中若另外保存 \(\bar\alpha_0=1\)⁠,公式下标会发生偏移,必须统一约定而不能照抄代码。

  下面的动画从接近高斯噪声的样本开始,展示 DDPM 如何在每个时间步利用网络预测的累计等价噪声逐步完成反向更新,直至得到较为清晰的样本。

为什么还要加噪声,而不直接以预测均值作为 \(\bx_{t-1}\)⁠?

  读者可能会问:既然 \(\bmu_{\btheta}(\bx_t,t)\) 是反向高斯 \(p_{\btheta}(\bx_{t-1}\mid\bx_t)\) 的均值,即给定 \(\bx_t\) 时对 \(\bx_{t-1}\) 的最可能预测,直接把它当作输出不是更省事吗?关键在于区分 估计采样⁠。训练阶段用高斯 KL 散度对齐的是整个分布 \(q(\bx_{t-1}\mid\bx_t,\bx_0)\)\(p_{\btheta}(\bx_{t-1}\mid\bx_t)\)⁠,而不是让某个单点相等;因此采样阶段也应当从 \(p_{\btheta}(\bx_{t-1}\mid\bx_t)\) 中随机抽取,而不是每次都取分布中心的均值。若每一步都只输出均值,得到的 \(\bx_{t-1}\) 永远落在分布中心,相当于把后验“坍塌”成一个点,无法与目标后验匹配。

  更具体地说,给定同一个含噪输入 \(\bx_t\)⁠,真正干净的 \(\bx_{t-1}\) 往往有多个同样合理的取值,例如同一团模糊像素既可能对应数字 3,也可能对应数字 8。条件均值 \(\bmu_{\btheta}(\bx_t,t)\) 是这些候选的加权平均,直接取均值会把不同模式“平均”成一张模糊的混合图像;只有在均值上加上尺度为 \(\sqrt{\tilde\beta_t}\) 的随机扰动,才能从多个可行解中随机选出一个具体、清晰的样本。这也解释了生成多样性的来源:即使固定初始噪声 \(\bx_T\) 与网络参数,每一步加入的小噪声仍会让最终图像在不同合理答案之间变化,而不是永远返回同一张“平均脸”。

  加噪声还有一个与训练目标一致的好处:噪声预测存在误差时,\(\bmu_{\btheta}(\bx_t,t)\) 会系统性地偏向误差方向,而随机项使样本围绕后验分布波动,单步偏差不会被整条链累积成固定的系统性偏移。噪声尺度 \(\sqrt{\tilde\beta_t}\) 在接近干净端点时明显缩小,最后一步 \(\tilde\beta_1=0\)⁠,因此越接近 \(\bx_0\) 扰动越小,最后一步直接返回均值也不会在最终图像中留下可见噪点。需要说明的是,确定性更新并非完全不可用:后文 DDIM 等方法会刻意去掉随机项并采用配套的确定性更新式,但那是另一种采样方案,不能理解为“DDPM 每步省略噪声”。

每步都在加噪声,为什么最后得到的还是清晰图像?

  反向过程确实每一步都加入随机噪声,但样本仍然越来越清晰,原因是每步"去除的噪声"远多于"加入的噪声"。前向过程定义了一条噪声阶梯:解析式 \(\bx_t=\sqrt{\bar\alpha_t}\bx_0+\sqrt{1-\bar\alpha_t}\tilde{\bepsilon}_t\) 中,累计信号保留系数 \(\bar\alpha_t\)\(t\) 增大而单调下降,样本的信噪比 \(\operatorname{SNR}(t)=\bar\alpha_t/(1-\bar\alpha_t)\) 也随之单调下降;\(t=T\)\(\bar\alpha_T\approx0\)⁠,样本几乎只剩噪声,\(t=0\)\(\bar\alpha_0=1\)⁠,样本就是原图。反向链要做的正是沿这条阶梯逐级爬回:每步的均值 \(\bmu_{\btheta}(\bx_t,t)\) 相当于网络在判断"这团噪声下面藏着什么结构",并把样本从 \(t\) 级噪声推向 \(t-1\) 级更干净的估计;而这一步加入的噪声尺度只有 \(\sqrt{\tilde\beta_t}\)⁠,它是这一级后验的剩余不确定性,数值上通常远小于样本中已有的噪声幅度 \(\sqrt{1-\bar\alpha_t}\)⁠,且噪声等级越高差距越大、最后一步则完全为 0。因此每完成一步,样本中的噪声占比就下降一档,信噪比严格上升。

  也可以这样理解:加入的小噪声并不会破坏网络刚恢复的结构,它只是让 \(\bx_{t-1}\) 的分布保持正确;而这部分新噪声又会成为下一步均值要去除的对象。于是整条链表现为"去噪、加一小点噪声、再去噪、再加一小点噪声……"的循环,其中每一次"去噪"都比"加噪"多,经过 \(T\) 步后,残留在 \(\bx_1\) 中的噪声已经微乎其微,最后一步 \(\tilde\beta_1=0\)⁠,不再加噪,直接给出清晰的 \(\bx_0\)⁠。整个过程很像逐层"对焦":网络每步先根据当前噪声等级判断整体轮廓,再逐步补充边缘和局部细节,越到后面每一步处理得越精细,最终图像也就越清晰。

训练与采样核验#

  DDPM 的完整流程包含两条相互联系的路径。 训练路径 用解析式采样一步构造带噪样本,让网络学习预测样本中隐藏的累计等价噪声; 生成路径 (也称采样路径)从纯噪声出发,反复调用网络、逐步去噪,最后得到接近真实数据的样本。两条路径共用同一组加噪系数 \(\beta_t,\alpha_t,\bar\alpha_t\) 和同一个网络 \(\bepsilon_{\btheta}\)⁠,因此最容易在 形状时间下标系数配对 三处出错。下面分别给出最简训练循环和最简采样循环,并逐项说明应当核对什么;对程序还不熟悉的读者,先照着一行行实现、再回头对照公式理解,是掌握这条流程最直接的方式。

训练循环的核对#

  训练阶段只做一件事:让网络学会“看到带噪样本和时间步之后,猜出样本里藏着多少噪声”。一个最小训练循环通常按以下顺序执行:

  1. 从训练集抽取一批干净样本 \(\bx_0\)⁠。若每个样本是图像,则这批数据的形状为 \(m\times d_C\times d_H\times d_W\)⁠,其中 \(m\) 是批量大小,后三个数依次是通道数、高和宽。

  2. 为每个样本独立随机抽取一个时间步 \(t\in\{1,\ldots,T\}\)⁠。不同样本可以使用不同的 \(t\)⁠,它们对应不同强度的噪声。

  3. 用解析式 \(\bx_t=\sqrt{\bar\alpha_t}\bx_0+\sqrt{1-\bar\alpha_t}\tilde{\bepsilon}_t\) 一步构造带噪样本,并保留本次实际使用的累计等价噪声 \(\tilde{\bepsilon}_t\) 作为监督目标。这里并不需要真的从 \(\bx_0\) 连续加噪 \(t\) 次。

  4. 网络接收 \(\bx_t\) 和时间步 \(t\)⁠,输出预测 \(\bepsilon_{\btheta}(\bx_t,t)\)⁠;用简化损失 \(\mathcal{L}_{\mathrm{simple}}=\mathbb{E}_{\bx_0,t,\tilde{\bepsilon}_t}\left[\|\tilde{\bepsilon}_t-\bepsilon_{\btheta}(\bx_t,t)\|_2^2\right]\) 度量预测与真实噪声在整批上取平均后的逐元素误差。

  5. 对损失反向传播,更新网络参数 \(\btheta\)⁠,然后回到第 1 步重复。

  训练循环中最容易出错的是形状、下标和数值范围,下面分别说明。

数据规模:输出必须与噪声同形

  网络输出 \(\bepsilon_{\btheta}(\bx_t,t)\) 必须与 \(\bx_t\)⁠、\(\tilde{\bepsilon}_t\) 完全同形,逐元素相减、平方和取平均才有意义。程序可以在前向之后加一条形状断言,例如要求 pred.shape == noise.shape第一时间发现维度错位。

常见误区:时间下标从 0 还是从 1 开始

  数学公式从第 1 步开始编号,程序数组常从位置 0 开始存放,二者相差 1。若取系数 \(\bar\alpha_t\) 时下标错位,构造出的 \(\bx_t\) 会用错噪声强度,而训练损失可能仍然很小,难以察觉。稳妥的做法是在程序开头把下标约定写成注释,并分别测试第一步(数组位置 0)和最后一步(数组位置 \(T-1\)⁠)取出的系数是否与公式一致。

  每一步之后都应检查损失和中间量为有限值;若出现 NaN 或无穷大,先检查输入范围、平方根下的数是否非负、除法分母是否可能为 0,而不是加大学习率或用裁剪掩盖问题。

采样循环的核对#

  生成样本时,从接近标准高斯的 \(\bx_T\sim\mathcal{N}(\bzero,\bI)\) 出发,按 \(t=T,T-1,\ldots,1\) 的反向顺序逐步去噪。每个反向步骤包含四件事:

  1. 用网络预测累计等价噪声 \(\bepsilon_{\btheta}(\bx_t,t)\)⁠。

  2. 由预测反推干净样本估计 \(\hat{\bx}_0=\frac{\bx_t-\sqrt{1-\bar\alpha_t}\bepsilon_{\btheta}(\bx_t,t)}{\sqrt{\bar\alpha_t}}\)⁠。它本身不直接作为输出,但可以用于检查预测是否合理,例如看数值范围是否落在图像取值范围内。

  3. 计算反向高斯均值 \(\bmu_{\btheta}(\bx_t,t)=\frac{1}{\sqrt{\alpha_t}}\left(\bx_t-\frac{\beta_t}{\sqrt{1-\bar\alpha_t}}\bepsilon_{\btheta}(\bx_t,t)\right)\)⁠,它表示本步最可能的 \(\bx_{t-1}\)⁠。

  4. \(t>1\)⁠,在均值上加入尺度为 \(\sqrt{\tilde\beta_t}\) 的随机噪声,得到 \(\bx_{t-1}\)⁠;若 \(t=1\)⁠,直接返回均值作为最终输出 \(\bx_0\)⁠。

常见误区:每一步的系数必须成套使用

  每一步都必须使用属于同一个时间步的 \(\beta_t\)⁠、\(\alpha_t\)\(\bar\alpha_t\)⁠,其中 \(\alpha_t=1-\beta_t\)⁠、\(\bar\alpha_t=\prod_{s=1}^{t}\alpha_s\)⁠。若从不同位置取系数,例如用 \(\bar\alpha_t\)\(\beta_{t-1}\)⁠,均值和方差的尺度会错位,生成结果可能过暗、过亮甚至发散。实现时应把所有系数按时间步放在同一组数组里,用同一个下标取出。

数据规模:后验方差与单步方差不同

  反向一步加入的随机噪声,其方差是后验方差 \(\tilde\beta_t=\frac{1-\bar\alpha_{t-1}}{1-\bar\alpha_t}\beta_t\)⁠,而不是第 \(t\) 步前向加噪的方差 \(\beta_t\)⁠。后验方差考虑了从 \(\bx_0\) 一步构造 \(\bx_t\) 时已经积累的全部噪声,数值通常小于 \(\beta_t\)⁠。若把 \(\beta_t\) 直接当作反向噪声方差,生成结果会带有过量噪声,除非明确说明所采用的模型变体。最后一步(\(t=1\)⁠)应直接返回均值、不再加噪,否则最终样本会残留一层噪声。

动手检查:用“已知答案”验证每一步

  有几种低成本核验方法:① 把网络预测替换成构造 \(\bx_t\) 时使用的真实累计等价噪声 \(\tilde{\bepsilon}_t\)⁠,此时反向均值应等于理论表达式,可逐项手算对比;② 在 \(t=1\)\(t=T\) 两端各测一次,核对数组下标与数学编号的对应关系;③ 检查最终输出与干净样本的数值范围一致,例如都落在 \([-1,1]\)⁠;④ 固定随机种子重复运行,确认结果可以复现。

  对 \(\hat{\bx}_0\) 的裁剪或动态阈值化可以防止极端值,但会改变采样动力学,应作为明确配置记录并在报告中说明,不能把加裁剪与不加裁剪的结果直接混为一谈。

Shiny 交互演示:DDPM 前向采样与反向更新

  交互页面可以核对前向解析采样中的信号系数、噪声系数和信噪比,并在同一条教学轨迹中查看任意一次 \(t\to t-1\) 的 DDPM 更新。页面会显示后验方差,并明确最后一步不再加入随机噪声;其中噪声预测器读取真实样本,只用于演示公式。

点击打开“扩散模型逐步实验室”

核心推导与实现核验#

核心关系

\[\mathcal{L}_{\mathrm{simple}}=\mathbb{E}_{\bx_0,t,\tilde{\bepsilon}_t}\left[\|\tilde{\bepsilon}_t-\bepsilon_{\btheta}(\bx_t,t)\|_2^2\right].\]

  推导路径。 先用高斯乘积求 \(q(\bx_{t-1}\mid\bx_t,\bx_0)\)⁠,再让 \(p_{\btheta}\) 的均值逼近该后验均值。把 \(\bx_t\) 的解析重参数化公式代入后,可将均值误差写成噪声预测误差;删除时间权重得到常用简化目标。

关键条件

  高斯后验公式依赖线性高斯前向链。简化 MSE 与完整的 ELBO 的逐项最优解在常见固定方差设定下相关,但权重被改变,因此评价似然时仍需使用正确的变分项。

数据规模

  若输入是 \(m\times d_C^{[l]}\times d_H^{[l]}\times d_W^{[l]}\) 的一批图像,噪声网络 \(\bepsilon_{\btheta}\) 也必须输出同样大小的结果。每张图像的时间嵌入先表示为 \(m\times d_t\)⁠,当前时间步对应的加噪系数则整理为 \(m\times1\times1\times1\)⁠,从而分别作用于对应图像。

常见误区

  最常见错误是把 \(\alpha_t\)\(\bar\alpha_t\) 混淆、把单步噪声 \(\bepsilon_t\) 误当成网络预测的累计等价噪声 \(\tilde{\bepsilon}_t\)⁠、从 0 与从 1 开始的时间下标错位、在 \(t=0\) 时仍加噪、使用错误的后验方差公式,以及训练预处理与采样后处理不一致。

动手检查

  在一维或二维高斯混合数据上训练小网络,手算后验均值方差并与代码比较;将网络预测替换为构造当前 \(\bx_t\) 时使用的真实累计等价噪声 \(\tilde{\bepsilon}_t\)⁠,检查反向一步的均值是否等于理论表达。

数值稳定性与规模

  除以 \(\sqrt{1-\bar\alpha_t}\)\(\sqrt{\bar\alpha_t}\) 时要防止端点数值问题。方差先在高精度中预计算并裁剪理论上的微小负值;混合精度训练应监控极端时间步损失。

本节小结#

  1. DDPM 用参数化高斯反向链逼近可计算的条件后验,并从标准高斯逐步采样。

  2. 经典噪声预测 MSE 来自变分目标的重参数化与重新加权,而不是与完整的 ELBO 数值完全相同。

  3. \(\tilde{\bepsilon}_t\)⁠、\(\bx_0\)\(\bv\) 参数化可以转换,但时间权重、数值尺度与采样公式必须成套使用。

综合练习#

  程序题应固定随机种子、写出维度断言并报告运行环境;比较题还应固定数据划分、随机种子集合和训练预算。全部参考答案见 DDPM 答案⁠。

  1. 前向边缘分布。\(q(\bx_t\mid\bx_{t-1})=\mathcal N(\sqrt{\alpha_t}\bx_{t-1},\beta_t\bI)\)⁠,其中 \(\alpha_t=1-\beta_t\)⁠、\(\bar\alpha_t=\prod_{s=1}^t\alpha_s\)⁠。用数学归纳法证明 \(q(\bx_t\mid\bx_0)=\mathcal N(\sqrt{\bar\alpha_t}\bx_0,(1-\bar\alpha_t)\bI)\)⁠。

  2. Bayes 公式与连乘整理。\(t\geq2\)⁠,从条件概率的定义出发,并利用前向过程的马尔可夫性质 \(q(\bx_t\mid\bx_{t-1},\bx_0)=q(\bx_t\mid\bx_{t-1})\)⁠,验证

    \[q(\bx_t\mid\bx_{t-1}) =q(\bx_{t-1}\mid\bx_t,\bx_0) \frac{q(\bx_t\mid\bx_0)} {q(\bx_{t-1}\mid\bx_0)}.\]

    然后,将 \(t=2,\ldots,T\) 对应的等式相乘,并结合前向过程和反向生成过程的联合分布分解,验证

    \[\frac{q(\bx_{1:T}\mid\bx_0)} {p_{\btheta}(\bx_{0:T})} =\frac{q(\bx_T\mid\bx_0)}{p(\bx_T)} \frac{1}{p_{\btheta}(\bx_0\mid\bx_1)} \prod_{t=2}^{T} \frac{q(\bx_{t-1}\mid\bx_t,\bx_0)} {p_{\btheta}(\bx_{t-1}\mid\bx_t)}.\]
  3. 条件后验推导。 利用 \(q(\bx_{t-1}\mid\bx_t,\bx_0)\propto q(\bx_t\mid\bx_{t-1})q(\bx_{t-1}\mid\bx_0)\)⁠,通过合并关于 \(\bx_{t-1}\) 的二次项,推导后验方差 \(\tilde\beta_t\) 和均值 \(\tilde{\bmu}_t\)⁠。

  4. 各向同性高斯分布的 KL 散度。\(q(\bx_{t-1}\mid\bx_t,\bx_0)=\mathcal{N}(\tilde{\bmu}_t,\tilde\beta_t\bI)\)⁠,\(p_{\btheta}(\bx_{t-1}\mid\bx_t)=\mathcal{N}(\bmu_{\btheta}(\bx_t,t),\sigma_t^2\bI)\)⁠,并假设 \(\tilde\beta_t\)\(\sigma_t^2\) 均不依赖 \(\btheta\)⁠。从两个多元高斯分布的 KL 散度公式出发,验证

    \[D_{\mathrm{KL}}\!\left( q(\bx_{t-1}\mid\bx_t,\bx_0) \,\middle\|\, p_{\btheta}(\bx_{t-1}\mid\bx_t) \right) =\frac{1}{2\sigma_t^2} \left\| \tilde{\bmu}_t -\bmu_{\btheta}(\bx_t,t) \right\|_2^2 +C_t,\]

    并写出 \(C_t\) 的具体表达式,说明它为什么与 \(\btheta\) 无关。

  5. 反向均值的一致性。 由解析式 \(\bx_t=\sqrt{\bar\alpha_t}\bx_0+\sqrt{1-\bar\alpha_t}\tilde{\bepsilon}_t\)\(\bx_0\) 表示成 \(\bx_t\) 与累计等价噪声 \(\tilde{\bepsilon}_t\) 的表达式,代入正文中的条件后验均值 \(\tilde{\bmu}_t\) 并化简;将结果与正文中 \(\bmu_{\btheta}(\bx_t,t)\) 的常用形式比较,验证当 \(\bepsilon_{\btheta}(\bx_t,t)=\tilde{\bepsilon}_t\) 时两者完全相等,并推导二者之差的表达式

    \[\tilde{\bmu}_t-\bmu_{\btheta}(\bx_t,t) =\frac{\beta_t}{\sqrt{\alpha_t(1-\bar\alpha_t)}} \left( \bepsilon_{\btheta}(\bx_t,t)-\tilde{\bepsilon}_t \right),\]

    说明二者之差为什么完全来自噪声预测误差。

  6. 前向加噪计算。 取两步加噪参数 \(\beta_1=0.1\)⁠、\(\beta_2=0.2\)⁠,标量数据 \(x_0=2\)⁠,用于从 \(x_0\) 一步构造 \(x_2\) 的累计等价噪声 \(\tilde{\epsilon}_2=0.5\)⁠。计算 \(\alpha_1,\alpha_2,\bar\alpha_2\)\(x_2=\sqrt{\bar\alpha_2}x_0+\sqrt{1-\bar\alpha_2}\tilde{\epsilon}_2\)⁠;若网络精确预测 \(\tilde{\epsilon}_2\)⁠,验证恢复公式得到原来的 \(x_0\)⁠。

  7. 平方损失的最优预测。 对固定的 \((\bx_t,t)\)⁠,证明使 \(\mathbb E[\lVert\tilde{\bepsilon}_t-f(\bx_t,t)\rVert_2^2\mid\bx_t,t]\) 最小的函数值是 \(\mathbb E(\tilde{\bepsilon}_t\mid\bx_t,t)\)⁠,并说明噪声预测网络一般不能恢复某次未知累计等价噪声的精确取值。

  8. DDPM 核心实现。Python 预先计算各时间步的加噪参数,并实现随机时间步采样、解析形式的前向加噪、简化噪声预测损失和一次反向采样;明确数组 0 基下标与数学时间下标的对应关系。

  9. 分布与自动微分核验。 在一维高斯数据上比较逐步前向采样和解析形式采样的经验均值、方差;对一个小型噪声网络比较简化损失的 自动微分梯度中心差分⁠,并验证网络预测等于真实累计等价噪声 \(\tilde{\epsilon}_t\)\(\hat x_0=x_0\)⁠。

  10. 测试设计。 至少测试:\(\bar\alpha_t\) 单调下降且位于 \((0,1)\)⁠、\(t=0\) 不额外加随机噪声、网络输出与输入同形、后验方差非负、固定种子可复现、时间下标越界报错,以及不同批次样本能使用不同时间步。

  11. 加噪参数设置比较。 在同一 DDPM 和数据集上,比较“让 \(\beta_t\) 随时间步线性变化”和“按照余弦函数设置 \(\bar\alpha_t\)⁠”两种方案。固定训练/验证划分、随机种子集合、网络、优化器、训练预算和采样步数;报告生成质量与覆盖、训练时间、固定样本数的生成时间、网络评估次数 (number of function evaluations,NFE)、参数量及训练/生成峰值内存。

  12. 预测参数化比较。 比较 \(\tilde{\bepsilon}_t\)⁠、\(\bx_0\)\(\bv\) 三种预测参数化,并使用与各参数化匹配的损失权重和采样转换。固定数据划分、随机种子集合、网络容量、训练预算和采样器;报告生成质量、训练稳定性、训练与生成时间、NFE、参数量及峰值内存。

  13. 采样步数比较。 对同一个已训练 DDPM 使用完整步数以及两个较短时间子序列采样。固定初始噪声集合、随机种子集合和后处理;报告生成质量、多样性、固定样本数生成时间、NFE、参数量及生成峰值内存,并说明减少步数带来的速度—质量权衡。