数学推导
前向过程与条件分布
前向过程是
$$x_t\mid x_{t-1} \sim \mathcal{N}(\sqrt{1-\beta_t}x_{t-1}, \beta_t I)$$
即
$$x_t = \sqrt{1-\beta_t} x_{t-1} + \sqrt{\beta_t}\varepsilon_t$$
从而可以得到在给定 \(x_0\) 的情况下,任意 \(x_t\) 的条件分布
$$\begin{align*} x_1\mid x_0 &\sim \mathcal{N}(\sqrt{1-\beta_1}x_0, \beta_1 I)\\ x_2\mid x_0 &\sim \sqrt{1-\beta_2}\mathcal{N}(\sqrt{1-\beta_1}x_0, \beta_1 I) + \mathcal{N}(0, \beta_2 I)\\ & = \mathcal{N}(\sqrt{(1-\beta_1)(1-\beta_2)}x_0, 1-(1-\beta_1)(1-\beta_2)I)\\ &\dots \end{align*}$$
令 \(\alpha_t = \prod_{i=1}^t(1-\beta_i)\),则
$$x_t\mid x_0 \sim \mathcal{N}(\sqrt{\alpha_t}x_0, (1-\alpha_t)I)$$
逆向过程的贝叶斯推导
我们可以知道,当 \(t\to \infty\) 时,\(\alpha_t\to 0\),上述分布渐进于 \(\mathcal{N}(0, I)\)。由于上面的分布都是条件分布,因此,在逆向过程中,我们希望得到 \(p(x_{t-1}\mid x_t)\),却只能得到
$$\begin{align*} p(x_{t-1}\mid x_t, x_0) &= \frac{p(x_{t-1}, x_t, x_0)}{p(x_t, x_0)} = \frac{p(x_{t-1}, x_t\mid x_0)}{p(x_t\mid x_0)}\\ &= \frac{p(x_t\mid x_{t-1}, x_0)p(x_{t-1}\mid x_0)}{p(x_t\mid x_0)}\\ &= \frac{p(x_t\mid x_{t-1})p(x_{t-1}\mid x_0)}{p(x_t\mid x_0)} \quad(\text{由马氏性}) \end{align*}$$
从而
$$\begin{align*} p(x_{t-1}\mid x_t, x_0) &\propto \exp\big(-\frac{1}{2\beta_t}\|x_t-\sqrt{1-\beta_t}x_{t-1}\|^2-\frac{1}{2(1-\alpha_{t-1})}\|x_{t-1}-\sqrt{\alpha_{t-1}}x_0\|^2\big)\\ &\propto \exp\bigg(-\frac{(1-\beta_t)\|x_{t-1}\|^2-2\sqrt{1-\beta_t} x_{t-1}'x_t}{2\beta_t}-\frac{\|x_{t-1}\|^2-2\sqrt{\alpha_{t-1}}x_{t-1}'x_0}{2(1-\alpha_{t-1})}\bigg)\\ &\propto\exp\bigg(-\frac{1}{2}\big(\frac{1-\beta_t}{\beta_t}+\frac{1}{1-\alpha_{t-1}}\big)\|x_{t-1}\|^2+x_{t-1}'\big(\frac{\sqrt{1-\beta_t}}{\beta_t}x_t+\frac{\sqrt{\alpha_{t-1}}}{1-\alpha_{t-1}}x_0\big)\bigg) \end{align*}$$
其中
$$\frac{1-\beta_t}{\beta_t}+\frac{1}{1-\alpha_{t-1}} = \frac{1-\beta_t - \alpha_t + \beta_t}{\beta_t(1-\alpha_{t-1})} = \frac{1-\alpha_t}{\beta_t(1-\alpha_{t-1})}$$
$$\frac{\beta_t(1-\alpha_{t-1})}{1-\alpha_t}\big(\frac{\sqrt{1-\beta_t}}{\beta_t}x_t + \frac{\sqrt{\alpha_{t-1}}}{1-\alpha_{t-1}}x_0\big) = \frac{(1-\alpha_{t-1})\sqrt{1-\beta_t}}{1-\alpha_t}x_t + \frac{\beta_t\sqrt{\alpha_{t-1}}}{1-\alpha_t}x_0$$
所以
$$x_{t-1}\mid x_t, x_0\sim \mathcal{N}\bigg(\frac{(1-\alpha_{t-1})\sqrt{1-\beta_t}}{1-\alpha_t}x_t + \frac{\beta_t\sqrt{\alpha_{t-1}}}{1-\alpha_t}x_0, \frac{\beta_t(1-\alpha_{t-1})}{1-\alpha_t}I\bigg)\tag{1}$$
现在问题来了,如何从 \(p(x_{t-1}\mid x_t, x_0)\) 得到 \(p(x_{t-1}\mid x_t)\)。因为条件中都有一个 \(x_t\),所以可以拆解成
$$p(x_{t-1}, x_0\mid x_t) = p(x_{t-1}\mid x_t, x_0)p(x_0\mid x_t)$$
从而
$$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{$\ast$}$$
但是注意
$$p(x_0\mid x_t) = \frac{p(x_t\mid x_0)p(x_0)}{p(x_t)}\tag{2}$$
它依赖 \(x_0\) 本身的分布,而这正是我们要求的。
举个例子,如果 \(t\) 很大,那么 \(p(x_t)\) 和 \(p(x_t\mid x_0)\) 就近似于标准正态分布,那么上式的结果就 \(\approx p(x_0)\)
近似与采样
因此我们知道了,积分 \((\ast)\) 是不可算的。而且我们也不需要得到一个分布,我们需要的是根据已有的 \(x_t\) 得到 \(x_{t-1}\) 的一个采样。浅薄地理解,
$$p(x_{t-1}\mid x_t) = \mathbb{E}_{x_0\sim p(x_0\mid x_t)}p(x_{t-1}\mid x_t, x_0)$$
如果我们能在给定 \(x_t\) 的条件下估计出 \(x_0\)(这里指的是得到 \(p(x_0\mid x_t)\),然后从中采样得到 \(x_0\)),我们就能得到 \(p(x_{t-1}\mid x_t)\) 的点估计,然后从中采样出 \(x_{t-1}\)。具体见下。
注意 \(x_t = \sqrt{\alpha_t}x_0 + \sqrt{1-\alpha_t}\bar\varepsilon_t\),\(\bar\varepsilon_t\) 表示累积噪声,因此
$$x_0 = \frac{1}{\sqrt{\alpha_t}}(x_t - \sqrt{1-\alpha_t}\bar\varepsilon_t)$$
因此估计 \(x_0\) 就是在估计 \(\tilde\varepsilon_t = \bar\varepsilon_t\mid x_t\)。一旦有了 \(\tilde\varepsilon_t\),就有
$$p(x_{t-1}\mid x_t) \approx p(x_{t-1}\mid x_t, \tilde x_0)$$
其中
$$x_{t-1}\mid x_t, \tilde x_0 = \mathcal{N}\bigg(\frac{(1-\alpha_{t-1})\sqrt{1-\beta_t}}{1-\alpha_t}x_t + \frac{\beta_t\sqrt{\alpha_{t-1}}}{1-\alpha_t}\cdot\frac{1}{\sqrt{\alpha_t}}(x_t-\sqrt{1-\alpha_t}\tilde\varepsilon_t), \frac{\beta_t(1-\alpha_{t-1})}{1-\alpha_t}I\bigg)$$
$$\begin{align*} \tilde\mu_t &= \bigg(\frac{(1-\alpha_{t-1})\sqrt{1-\beta_t}}{1-\alpha_t} + \frac{\beta_t}{(1-\alpha_t)\sqrt{1-\beta_t}}\bigg)x_t - \frac{\beta_t\sqrt{\alpha_{t-1}}}{\sqrt{1-\alpha_t}\sqrt{\alpha_t}}\tilde\varepsilon_t\\ &=\frac{(1-\alpha_{t-1})(1-\beta_t)+\beta_t}{(1-\alpha_t)\sqrt{1-\beta_t}}x_t - \frac{\beta_t}{\sqrt{(1-\alpha_t)(1-\beta_t)}}\tilde\varepsilon_t\\ &=\frac{1-\beta_t-\alpha_t+\beta_t}{(1-\alpha_t)\sqrt{1-\beta_t}}x_t - \frac{\beta_t}{\sqrt{(1-\alpha_t)(1-\beta_t)}}\tilde\varepsilon_t\\ &= \frac{1}{\sqrt{1-\beta_t}}\big(x_t-\frac{\beta_t}{\sqrt{1-\alpha_t}}\tilde\varepsilon_t\big)\\ \\ \tilde\sigma_t^2 &= \frac{\beta_t(1-\alpha_{t-1})}{1-\alpha_t} \end{align*}$$
因此只要采样一个 \(z\sim\mathcal{N}(0,I)\),就可以得到 \(x_{t-1}\) 的估计
$$x_{t-1} = \tilde\mu_t + \tilde\sigma_t z$$
DDPM
上面的推导是基于概率密度的,而 DDPM 原论文则不然。它预测噪声的条件期望,即直接生成
$$\mathbb{E}[\bar\varepsilon_t\mid x_t]$$
这样一来,根据公式 \((1)\),
$$\mathbb{E}[x_{t-1}\mid x_t] = \mathbb{E}\bigg[\mathbb{E}[x_{t-1}\mid x_t, x_0]\bigm| x_t\bigg] = \mathbb{E}[Ax_t+Bx_0\mid x_t] = Ax_t + \mathbb{E}[Bx_0\mid x_t]$$
而要预测 \(\mathbb{E}[x_0\mid x_t]\),只要预测 \(\mathbb{E}[\bar\varepsilon_t\mid x_t]\) 就够了。
原则上,\(x_{t-1}\mid x_t\) 并非简单的高斯分布,但原论文选择使用高斯分布,并设置其方差为 \(\beta_t\) 或 \(\frac{1-\alpha_{t-1}}{1-\alpha_t}\beta_t\),便可采样出 \(x_{t-1}\)。
(由 Claude 生成)
问题一:\(\beta_t\) 和 \(\tilde\beta_t = \frac{1-\alpha_{t-1}}{1-\alpha_t}\beta_t\) 从何而来
关键在于区分两个不同的对象:
- \(\mathrm{Var}(x_{t-1}\mid x_t, x_0)\):已知 \(x_0\) 时的条件方差,这就是你公式 (1) 里算出来的 \(\tilde\sigma_t^2 = \tilde\beta_t\),它是个常数,不依赖 \(x_0\) 具体取值。
- \(\mathrm{Var}(x_{t-1}\mid x_t)\):真正想要的逆向过程方差,这是我们无法精确求出、只能选一个固定值去近似的对象。
用全概方差公式把后者展开,对 \(x_0\mid x_t\) 取期望: $$\mathrm{Var}(x_{t-1}\mid x_t) = \underbrace{\mathbb E_{x_0\mid x_t}\big[\mathrm{Var}(x_{t-1}\mid x_t,x_0)\big]}_{=\ \tilde\beta_t(常数,直接提出来)} + \underbrace{\mathrm{Var}_{x_0\mid x_t}\big(\mathbb E[x_{t-1}\mid x_t,x_0]\big)}_{\text{均值的波动项}}$$
第二项里,\(\mathbb E[x_{t-1}\mid x_t,x_0] = Ax_t + Bx_0\)(你公式 (1) 里的线性形式,\(B = \frac{\beta_t\sqrt{\alpha_{t-1}}}{1-\alpha_t}\)),所以
$$\mathrm{Var}(x_{t-1}\mid x_t) = \tilde\beta_t + B^2\cdot \mathrm{Var}(x_0\mid x_t)$$
这个式子就是全部秘密:真实方差 = \(\tilde\beta_t\) + 一个正的修正项,修正项的大小取决于给定 \(x_t\) 后,\(x_0\) 本身还有多"不确定"。
如果 \(x_0\) 是确定的一个点(\(\mathrm{Var}(x_0\mid x_t)=0\),比如数据分布是个 delta 函数),修正项消失,\(\mathrm{Var}(x_{t-1}\mid x_t)\) 恰好等于 \(\tilde\beta_t\)。这就是"固定 \(x_0\)"对应 \(\tilde\beta_t\) 的来源——不是近似,是精确成立。
如果 \(x_0\sim\mathcal N(0,I)\)(数据本身就是标准正态,也就是熵最大、最"扩散"的情形),可以直接算出 \(x_0\mid x_t\) 的后验(\(x_0,x_t\) 联合高斯,标准公式): $$\mathrm{Var}(x_0\mid x_t) = 1-\alpha_t$$ 代入上式: $$\tilde\beta_t + B^2(1-\alpha_t) = \frac{(1-\alpha_{t-1})\beta_t}{1-\alpha_t} + \frac{\beta_t^2\alpha_{t-1}}{(1-\alpha_t)^2}(1-\alpha_t) = \frac{\beta_t\big[(1-\alpha_{t-1}) + \beta_t\alpha_{t-1}\big]}{1-\alpha_t}$$ 注意 \(\alpha_{t-1}(1-\beta_t) = \alpha_t\),所以中括号里 \(= 1-\alpha_{t-1}(1-\beta_t) = 1-\alpha_t\),于是整体约掉,恰好得到 $$\mathrm{Var}(x_{t-1}\mid x_t) = \beta_t$$
完全精确地等于 \(\beta_t\),不是凑出来的近似。所以论文那句话的真正含义是:\(\tilde\beta_t\) 和 \(\beta_t\) 是真实方差 \(\mathrm{Var}(x_{t-1}\mid x_t)\) 的两个精确端点——分别对应"数据是单点(零熵)"和"数据本身就是各向同性高斯(最大熵、与噪声先验重合)"这两种极端数据分布。真实数据介于两者之间,所以两个选择都"差不多合理",这也是原论文两个都试了、发现效果接近的原因。(后来 Improved DDPM 里 Nichol & Dhariwal 干脆学习一个插值系数 \(v\),用 \(\sigma_t^2=\exp(v\log\beta_t+(1-v)\log\tilde\beta_t)\) 在两个端点之间插值,也印证了这个"端点"的理解是对的。)
问题二:没用变分下界,为什么还是推出了同一个模型和 MSE 损失
这其实不是巧合,而是同一个数学事实的两种呈现方式。你走的路径本质上是最小均方误差(MMSE)估计的论证,比 ELBO 更直接:
对任意随机变量 \(Y\) 和观测 \(X\),在 \(L^2\) 损失下,\(\mathbb E[(Y-f(X))^2]\) 的最优预测器 \(f^*(X)\) 就是条件期望 \(\mathbb E[Y\mid X]\)。这是个纯粹的统计事实,跟扩散模型、ELBO 都没关系。
你笔记里要的 \(\tilde\varepsilon_t := \mathbb E[\bar\varepsilon_t\mid x_t]\) 正是这个意义下的最优点估计。所以只要你约定"用 MSE 训练网络去预测 \(\bar\varepsilon_t\)",网络学到的自动就是 \(\mathbb E[\bar\varepsilon_t\mid x_t,t]\)——这一步完全不需要 ELBO,直接从 MSE 损失的性质就能保证。这就是为什么你绕开变分推导,仍然"巧合地"落到了同一个训练目标。
但这里要注意,ELBO 那条路径给出的其实是一个更精细的结论,你的推导是它的一个特例:
ELBO 逐项展开后,第 \(t\) 项是 $$D_{KL}\big(q(x_{t-1}\mid x_t,x_0)\,\|\,p_\theta(x_{t-1}\mid x_t)\big)$$ 如果强行让 \(p_\theta\) 的方差也固定为 \(\tilde\sigma_t^2\)(跟 \(q\) 一样),两个高斯的 KL 就退化成均值之差的平方: $$= \frac{1}{2\tilde\sigma_t^2}\big\|\tilde\mu_t(x_t,x_0) - \mu_\theta(x_t,t)\big\|^2 + \text{const}$$ 再把 \(\mu_\theta\) 用你上面那套噪声参数化代进去(形式跟 \(\tilde\mu_t\) 完全一样,只是用 \(\varepsilon_\theta\) 换掉 \(\tilde\varepsilon_t\)),系数会精确约化成 $$\propto \frac{\beta_t^2}{\tilde\sigma_t^2\,\alpha_t(1-\alpha_t)}\,\big\|\bar\varepsilon_t - \varepsilon_\theta(x_t,t)\big\|^2$$
也就是说,严格从 ELBO 推出来的是一个带时间步权重 \(\lambda(t)\) 的加权 MSE,而不是你我直觉上写的那个无权重版本。DDPM 论文里的 \(L_\text{simple}\),是人为把这个权重丢掉、设成 1,用的正是你那种"MMSE 点估计"式的朴素目标。论文发现这样训练效果反而更好(原因大致是无权重版本会给大 \(t\) 的步骤更大的相对权重,迫使模型更关注结构性的、感知上重要的部分,而不是被小 \(t\) 时那些像素级的琐碎误差主导)。
所以准确的关系是:你的推导(MMSE 论证)直接给出了 \(L_\text{simple}\);ELBO 推导给出的是 \(L_\text{simple}\) 的加权推广版 \(L_\text{vlb}\),而 \(L_\text{simple}\) 是舍弃权重后的特例。 两条路径殊途同归,是因为它们本质上都在问"给定 \(x_t\),对 \(\bar\varepsilon_t\) 的最优点估计是什么",只是 ELBO 的版本多告诉了你一件事——不同 \(t\) 步的误差理论上不该被同等对待,只是实践中发现同等对待反而更好。
模型与训练目标
如上所述,模型只需要预测 \(\tilde\varepsilon_t = \bar\varepsilon_t\mid x_t\)(或者预测 \(p(\bar\varepsilon_t\mid x_t)\) 然后从中采样出一个 \(\tilde\varepsilon_t\))即可,因此模型可以写为 \(f_\theta(x_t)\)。
显然即使 \(x_1 = x_t(t\neq 1)\),也不应预测相同的噪声(参见公式 \((2)\)),因此模型应该显式将时间步 \(t\) 也作为输入,即 \(f_\theta(x_t,t)\)。预测目标自然就是正向过程加的噪声,即
$$f_\theta(x_t,t) \approx \frac{1}{\sqrt{1-\alpha_t}}(x_t - \sqrt{\alpha_t}x_0)$$
损失使用 MSE 损失。