Diffusion Model(二):训练推导详解
paperreading
本文字数:2.3k 字 | 阅读时长 ≈ 10 min

Diffusion Model(二):训练推导详解

paperreading
本文字数:2.3k 字 | 阅读时长 ≈ 10 min

本文承接上一篇,从变分下界开始推导 DDPM 的训练目标,并在最后补充重参数化技巧。

1. Diffusion 的训练推导

1.1 最小化负对数似然与变分下界

在弄懂diffusion model前向和反向过程之后,最后我们需要了解其训练推导过程,即用什么loss以及为什么。在diffusion的反向过程中,根据\((3)\)式我们需要预测\(\mu_{\theta}(x_{t}, t), \Sigma_{\theta}(x_{t}, t)\),如何得到一个合理的均值和方差?类似于VAE,在对真实数据分布情况下,最大化模型预测分布的对数似然,即优化\(x_{0}\sim q(x_{0})\)下的\(p_{\theta}(x_{0})\)的交叉熵
\[ \mathcal{L} = \mathbb{E}_{q(x_{0})}[-\log p_{\theta}(x_{0})] \tag{12} \]

注意,我们的真实目标是最大化\(p_{\theta}(x)\),\((3)\)对\(p_{\theta}(x)\)取了负\(\log\),所以这里要最小化\(\mathcal{L}\)。与VAE类似,我们通过变分下限VLB来优化\((12)\)的负对数似然,如下
\[ \begin{aligned} -\log p_{\theta}(x_{0}) &\leq -\log p_{\theta}(x_{0}) + D_{KL}\!\left(q(x_{1:T}\mid x_{0}) \parallel p_{\theta}(x_{1:T}\mid x_{0})\right) \\ &= -\log p_{\theta}(x_{0}) + \mathbb{E}_{q(x_{1:T}\mid x_{0})}\left[\log \frac{q(x_{1:T}\mid x_{0})}{p_{\theta}(x_{1:T}\mid x_{0})}\right] \\ &= -\log p_{\theta}(x_{0}) + \mathbb{E}_{q(x_{1:T}\mid x_{0})}\left[\log \frac{q(x_{1:T}\mid x_{0})}{p_{\theta}(x_{0:T})/p_{\theta}(x_{0})}\right] \\ &= \mathbb{E}_{q(x_{1:T}\mid x_{0})}\left[\log \frac{q(x_{1:T}\mid x_{0})}{p_{\theta}(x_{0:T})}\right] \end{aligned} \tag{13} \]

现在我们得到了最大似然的变分下界,即式\((13)\),但是根据式\((12)\),我们要优化\(\mathbb{E}_{q(x_{0})}[-\log p_{\theta}(x_{0})]\),比\((13)\)多了个期望,所以这里我们加上期望,同时根据重积分中的Fubini定理,得到
\[ \begin{aligned} \mathbb{E}_{q(x_{0})}[-\log p_{\theta}(x_{0})] & \leq \mathcal{L}_{VLB} \\ &= \mathbb{E}_{q(x_{0})}\left( \mathbb{E}_{q(x_{1:T}|x_{0})} \log \frac{q(x_{1:T}|x_{0})}{p_{\theta}(x_{0:T})} \right) \\ &= \mathbb{E}_{q(x_{0:T})} \log \frac{q(x_{1:T}|x_{0})}{p_{\theta}(x_{0:T})} \end{aligned} \tag{14} \]

1.2 变分下界的详细推导与优化

再次回到式 \((12)\),我们的目标是最小化交叉熵。现在已经得到它的变分上界,因此可以转为最小化 \(\mathcal{L}_{VLB}\)。利用前向过程和反向过程的马尔可夫分解,可以得到:
\[ \begin{aligned} \mathcal{L}_{VLB} &= \mathbb{E}_{q(x_{0:T})}\left[\log \frac{q(x_{1:T}\mid x_{0})}{p_{\theta}(x_{0:T})}\right] \\ &= \mathbb{E}_{q(x_{0:T})}\left[-\log p_{\theta}(x_T) \mathrel{+} \sum_{t=1}^{T}\log \frac{q(x_t\mid x_{t-1})}{p_{\theta}(x_{t-1}\mid x_t)}\right] \\ &= \mathbb{E}_{q(x_{0:T})}\left[ \log \frac{q(x_T\mid x_0)}{p_{\theta}(x_T)} \mathrel{+} \sum_{t=2}^{T}\log \frac{q(x_{t-1}\mid x_t,x_0)}{p_{\theta}(x_{t-1}\mid x_t)} \mathrel{-} \log p_{\theta}(x_0\mid x_1)\right] \\ &= \mathbb{E}_{q(x_{0:T})}\left[ \underbrace{D_{KL}\!\left(q(x_T\mid x_0)\parallel p_{\theta}(x_T)\right)}_{L_T} \mathrel{+} \sum_{t=2}^{T}\underbrace{D_{KL}\!\left(q(x_{t-1}\mid x_t,x_0)\parallel p_{\theta}(x_{t-1}\mid x_t)\right)}_{L_t} \mathrel{-} \underbrace{\log p_{\theta}(x_0\mid x_1)}_{L_0}\right] \end{aligned} \]

最后一行把相应变量上的条件期望写成了 KL 散度;最外层的 \(\mathbb{E}_{q(x_{0:T})}\) 已经包含了从前向过程采样所需的期望。

公式补充说明

1. 使用贝叶斯公式改写前向转移

由\(x_{t-1}\)得\(x_{t}\)比较困难,但是当提供额外的\(x_{0}\)时,\(x_{t-1}\)和\(x_{t}\)的候选会减少,选择更加确定,因此这里我们加上\(x_{0}\),这样推导如下
\[ q(x_t\mid x_{t-1}) = \frac{q(x_{t-1}\mid x_t,x_0)q(x_t\mid x_0)}{q(x_{t-1}\mid x_0)} \]

2. 为什么单独处理 \(t=1\)?

当 \(t=1\) 时,\(x_0\) 是观测数据,后验 \(q(x_0\mid x_1,x_0)\) 会退化为确定分布。因此这一项通常不继续写成 KL 散度,而是单独记为重建项 \(L_0=-\log p_{\theta}(x_0\mid x_1)\)。

1.3 由变分下界得到优化 Loss

去掉繁琐的推导过程以及最后的期望符号,我们再回头看一下VLB的化简结果
\[ \mathcal{L}_{VLB} \propto \underbrace{D_{KL}\!\left(q(x_T\mid x_0)\parallel p_{\theta}(x_T)\right)}_{L_T} \mathrel{+} \sum_{t=2}^{T}\underbrace{D_{KL}\!\left(q(x_{t-1}\mid x_t,x_0)\parallel p_{\theta}(x_{t-1}\mid x_t)\right)}_{L_t} \mathrel{-} \underbrace{\log p_{\theta}(x_0\mid x_1)}_{L_0} \]

\[ \begin{aligned} L_t &= \mathbb{E}_{x_0,\overline{z}_t}\left[ \frac{\left\lVert\tilde{\mu}_t(x_t,x_0)-\mu_{\theta}(x_t,t)\right\rVert_2^2} {2\lVert\Sigma_{\theta}(x_t,t)\rVert_2^2} \right] \\ &= \mathbb{E}_{x_0,\overline{z}_t}\left[ \frac{1}{2\lVert\Sigma_{\theta}(x_t,t)\rVert_2^2} \left\lVert \frac{\beta_t}{\sqrt{\alpha_t}\sqrt{1-\overline{\alpha}_t}} \left(\overline{z}_t-z_{\theta}(x_t,t)\right) \right\rVert_2^2 \right] \\ &= \mathbb{E}_{x_0,\overline{z}_t}\left[ \frac{\beta_t^2} {2\alpha_t(1-\overline{\alpha}_t)\lVert\Sigma_{\theta}(x_t,t)\rVert_2^2} \left\lVert \overline{z}_t-z_{\theta}\!\left( \sqrt{\overline{\alpha}_t}x_0+\sqrt{1-\overline{\alpha}_t}\,\overline{z}_t,t \right) \right\rVert_2^2 \right] \end{aligned} \]

所有的\(L\)就是最终的变分下界,DDPM将其进一步化简如下
\[ L_t^{\mathrm{simple}} = \mathbb{E}_{x_0,\overline{z}_t,t}\left[ \left\lVert \overline{z}_t-z_{\theta}\!\left( \sqrt{\overline{\alpha}_t}x_0+\sqrt{1-\overline{\alpha}_t}\,\overline{z}_t,t \right) \right\rVert_2^2 \right] \]

论文没有将方差\(\Sigma_{\theta}\)考虑在训练和推断中,而是将untrained的\(\beta_{t}\)代替\(\tilde{\beta_{t}}\),因为\(\Sigma_{\theta}\)可能会导致训练不稳定

到这里我们知道了Diffusion model的训练其实也是去预测每一步的噪声,就像反向过程中对均值推导的那样,\(\tilde{\mu}_{t} = \frac{1}{\sqrt{\alpha_{t}}}(x_{t} - \frac{\beta_{t}}{\sqrt{1-\overline{\alpha}_{t}}}\overline{z}_{t})\),这里的均值不依赖\(x_{0}\),而他的求解本质上就是\(x_{t}\)减去随机噪声

1.4 训练与推断

最后我们附上论文中训练和推断的过程

DDPM 的训练与采样流程

在看论文的时候有一点困惑:解释一下为什么训练的时候reverse一步到位,sample的时候得一步一步的来。因为训练的时候我们对于均值u的预测,是建立在\(x_{0}\)已知的基础上的,所以能够通过公式直接预测噪声进行训练,测试的时候我们只有采样的高斯随机噪声\(x_{t}\),并不知道他的\(x_{0}\)是什么,所以需要一步一步的预测噪声

2. 知识点补充

2.1 重参数化技巧

重参数化技巧在VAE中被应用过,此技巧主要用来使采样可以进行反向传播,假设我们随机采样时从任意一个高斯分布\(\mathcal{N}(\mu, \sigma^{2})\)中采样,然后预测结果,最终结果是无法反向传播的(不可导),通常做法是使用标准高斯分布\(\mathcal{N}(0, I)\)作为引导

具体做法是首先从标准高斯分布中采样一个变量,然后根据高斯分布的均值\(\mu\)和方差\(\sigma^{2}\)来对采样变量进行线性变换,如下
\[ z = \mu + \sigma \odot \epsilon, \epsilon \sim \mathcal{N}(0, I) \]

重参数化之后得到的变量\(z\)具有随机性的,满足均值为\(\mu\),方差为\(\sigma^{2}\)的高斯分布,这样采样过程就可导了,随机性加到了\(\epsilon \sim \mathcal{N}(0, I)\)上,而不是\(\mathcal{N}(\mu, \sigma^{2})\)

再通俗解释一下:如果直接从原高斯分布采样,采样之后的计算可以反向传播,但梯度无法越过带有随机性的采样步骤。当我们将随机性转移到\(\epsilon \sim \mathcal{N}(0, I)\)上时,模型只需要对采样之后参与变换的\(\mu\)和\(\sigma\)求导。

Sep 06, 2026
Sep 09, 2024
Sep 06, 2024