DDPM:参考答案

目录

\[ \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:参考答案#

返回正文练习 · 返回答案索引

说明#

以下答案与正文12道题逐题对应。推导题给出主要步骤;编程和实验题给出参考程序和检查方法,并说明结果适用于哪些条件。

  1. \(t=1\) 时结论由定义直接成立。假设 \(\bx_{t-1}=\sqrt{\bar\alpha_{t-1}}\bx_0+\sqrt{1-\bar\alpha_{t-1}}\tilde{\boldsymbol{\epsilon}}_{t-1}\)⁠,其中 \(\tilde{\boldsymbol{\epsilon}}_{t-1}\) 是构造 \(\bx_{t-1}\) 时前 \(t-1\) 步噪声合并后的累计噪声。第 \(t\) 步再独立加入 \(\boldsymbol{\epsilon}_t\)⁠,得 \(\bx_t=\sqrt{\alpha_t\bar\alpha_{t-1}}\bx_0+\sqrt{\alpha_t(1-\bar\alpha_{t-1})}\tilde{\boldsymbol{\epsilon}}_{t-1}+\sqrt{\beta_t}\boldsymbol{\epsilon}_t\)⁠。后两项之和仍为高斯,其协方差系数和为 \(\alpha_t(1-\bar\alpha_{t-1})+\beta_t=1-\bar\alpha_t\)⁠。将标准化后的总噪声记作 \(\tilde{\boldsymbol{\epsilon}}_t\)⁠,便得到 \(\bx_t=\sqrt{\bar\alpha_t}\bx_0+\sqrt{1-\bar\alpha_t}\tilde{\boldsymbol{\epsilon}}_t\)⁠,故归纳成立。

  2. 根据条件概率的定义,

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

    前向过程具有马尔可夫性质:给定 \(\bx_{t-1}\) 后,\(\bx_t\) 的分布不再依赖更早的 \(\bx_0\)⁠,所以

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

    将这个等式代入前式,再把 \(q(\bx_t\mid\bx_0)\) 移到等号左侧,并除以 \(q(\bx_{t-1}\mid\bx_0)\)⁠,便得到

    \[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)}.\]

    以上改写默认所除的概率密度在当前讨论的位置大于 \(0\)⁠。

    接下来,把前向路径分布和反向生成模型的联合分布展开:

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

    单独取出 \(t=1\) 对应的因子,并把刚刚得到的单步改写代入 \(t=2,\ldots,T\) 的各个前向转移,得到

    \[\begin{split}\begin{aligned} \frac{q(\bx_{1:T}\mid\bx_0)} {p_{\btheta}(\bx_{0:T})} ={}&\frac{q(\bx_1\mid\bx_0)} {p_{\btheta}(\bx_0\mid\bx_1)} \frac{1}{p(\bx_T)} \\ &\quad \prod_{t=2}^{T} \left[ \frac{q(\bx_{t-1}\mid\bx_t,\bx_0)} {p_{\btheta}(\bx_{t-1}\mid\bx_t)} \frac{q(\bx_t\mid\bx_0)} {q(\bx_{t-1}\mid\bx_0)} \right]. \end{aligned}\end{split}\]

    其中,与边缘分布有关的因子依次相消:

    \[q(\bx_1\mid\bx_0) \prod_{t=2}^{T} \frac{q(\bx_t\mid\bx_0)} {q(\bx_{t-1}\mid\bx_0)} =q(\bx_T\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)}.\]
  3. 两个高斯因子的精度相加:后验精度为 \(\alpha_t/\beta_t+1/(1-\bar\alpha_{t-1})=(1-\bar\alpha_t)/[\beta_t(1-\bar\alpha_{t-1})]\)⁠,因此 \(\tilde\beta_t=\beta_t(1-\bar\alpha_{t-1})/(1-\bar\alpha_t)\)⁠。把线性项乘以后验方差可得 \(\tilde{\boldsymbol\mu}_t=[\sqrt{\bar\alpha_{t-1}}\beta_t/(1-\bar\alpha_t)]\bx_0+[\sqrt{\alpha_t}(1-\bar\alpha_{t-1})/(1-\bar\alpha_t)]\bx_t\)⁠。推导依赖各向同性线性高斯前向链。

  4. \(\bx_{t-1}\) 的维度为 \(d\)⁠。对于 \(q=\mathcal{N}(\boldsymbol m_q,\boldsymbol\Sigma_q)\)\(p=\mathcal{N}(\boldsymbol m_p,\boldsymbol\Sigma_p)\)⁠,多元高斯分布的 KL 散度为

    \[D_{\mathrm{KL}}(q\|p) =\frac{1}{2} \left[ \operatorname{tr}(\boldsymbol\Sigma_p^{-1}\cdot\boldsymbol\Sigma_q) +(\boldsymbol m_p-\boldsymbol m_q)\trans \cdot\boldsymbol\Sigma_p^{-1} \cdot(\boldsymbol m_p-\boldsymbol m_q) -d +\log\frac{\det(\boldsymbol\Sigma_p)} {\det(\boldsymbol\Sigma_q)} \right].\]

    在本题中,\(\boldsymbol\Sigma_q=\tilde\beta_t\bI\)⁠、\(\boldsymbol\Sigma_p=\sigma_t^2\bI\)⁠,所以

    \[\begin{split}\begin{aligned} \operatorname{tr}(\boldsymbol\Sigma_p^{-1}\cdot\boldsymbol\Sigma_q) &=d\frac{\tilde\beta_t}{\sigma_t^2},\\ \log\frac{\det(\boldsymbol\Sigma_p)} {\det(\boldsymbol\Sigma_q)} &=d\log\frac{\sigma_t^2}{\tilde\beta_t},\\ (\boldsymbol m_p-\boldsymbol m_q)\trans \cdot\boldsymbol\Sigma_p^{-1} \cdot(\boldsymbol m_p-\boldsymbol m_q) &=\frac{1}{\sigma_t^2} \left\| \bmu_{\btheta}(\bx_t,t)-\tilde{\boldsymbol\mu}_t \right\|_2^2. \end{aligned}\end{split}\]

    代回一般公式便得到

    \[D_{\mathrm{KL}}(q\|p) =\frac{1}{2\sigma_t^2} \left\| \tilde{\boldsymbol\mu}_t-\bmu_{\btheta}(\bx_t,t) \right\|_2^2 +C_t,\]

    其中

    \[C_t =\frac{d}{2} \left( \frac{\tilde\beta_t}{\sigma_t^2} -1 +\log\frac{\sigma_t^2}{\tilde\beta_t} \right).\]

    因为 \(\tilde\beta_t\)⁠、\(\sigma_t^2\)\(d\) 都不依赖需要训练的参数 \(\btheta\)⁠,所以 \(C_t\) 对优化 \(\btheta\) 没有影响。与 \(\btheta\) 有关的部分只剩两个均值之差的平方。

  5. 由解析式解得 \(\bx_0=[\bx_t-\sqrt{1-\bar\alpha_t}\tilde{\bepsilon}_t]/\sqrt{\bar\alpha_t}\)⁠,代入条件后验均值:

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

    利用 \(\bar\alpha_t=\alpha_t\bar\alpha_{t-1}`(即 :math:\)sqrt{baralpha_{t-1}}/sqrt{baralpha_t}=1/sqrt{alpha_t}`)和 \(\beta_t=1-\alpha_t\)⁠,\(\bx_t\) 的系数化为 \([\beta_t+\alpha_t(1-\bar\alpha_{t-1})]/[\sqrt{\alpha_t}(1-\bar\alpha_t)]=(1-\bar\alpha_t)/[\sqrt{\alpha_t}(1-\bar\alpha_t)]=1/\sqrt{\alpha_t}\)⁠,\(\tilde{\bepsilon}_t\) 的系数为 \(-\beta_t/[\sqrt{\alpha_t}\sqrt{1-\bar\alpha_t}]\)⁠,故

    \[\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)\) 相减,得

    \[\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).\]

    因此当 \(\bepsilon_{\btheta}(\bx_t,t)=\tilde{\bepsilon}_t\) 时两者完全相等;一般情形下二者之差只由网络的噪声预测误差 \(\bepsilon_{\btheta}(\bx_t,t)-\tilde{\bepsilon}_t\) 决定,系数 \(\beta_t/\sqrt{\alpha_t(1-\bar\alpha_t)}\) 只含已知的加噪参数,与 \(\bx_t\) 无关。这也解释了训练目标为何是最小化该预测误差:反向均值与理想后验均值的偏差随它线性增长。

  6. \(\alpha_1=0.9\)⁠、\(\alpha_2=0.8\)⁠、\(\bar\alpha_2=0.72\)⁠。因此 \(x_2=\sqrt{0.72}(2)+\sqrt{0.28}(0.5)\approx1.9617\)⁠。代入恢复式 \(\widehat x_0=[x_2-\sqrt{0.28}(0.5)]/\sqrt{0.72}=2\)⁠,数值误差只来自舍入。这里的 \(\tilde{\epsilon}_2\) 是与 \(x_2\) 对应的解析重参数化累计噪声,不是从 \(x_1\)\(x_2\) 时单独加入的 \(\epsilon_2\)⁠。

  7. \(\boldsymbol m=\mathbb E(\tilde{\boldsymbol{\epsilon}}_t\mid\bx_t,t)\)⁠,把 \(\tilde{\boldsymbol{\epsilon}}_t-f=(\tilde{\boldsymbol{\epsilon}}_t-\boldsymbol m)+(\boldsymbol m-f)\)⁠。条件平方期望中的交叉项为 0,故风险等于与 \(f\) 无关的条件方差加 \(\lVert\boldsymbol m-f\rVert_2^2\)⁠,在 \(f=\boldsymbol m\) 时最小。相同 \(\bx_t\) 可能由不同 \(\bx_0\) 与噪声产生,所以网络学习的是与 \(\bx_t\) 对应的累计噪声 \(\tilde{\boldsymbol{\epsilon}}_t\) 的条件均值,不是为每个含噪输入唯一反推出某次噪声。

  8. 下面的实现把数组下标与数学时间统一为 \(t=0,1,\ldots,T\)⁠,其中 \(\bar\alpha_0=1\) 表示干净数据:

    def make_schedule(beta):
        beta = np.asarray(beta, dtype=np.float64)
        if beta.ndim != 1 or np.any((beta <= 0) | (beta >= 1)):
            raise ValueError("beta 必须是一维且位于 (0,1)")
        alpha = 1.0 - beta
        return {
            "beta": np.concatenate(([0.0], beta)),
            "alpha": np.concatenate(([1.0], alpha)),
            "alpha_bar": np.concatenate(([1.0], np.cumprod(alpha))),
        }
    
    def sample_timesteps(batch_size, schedule, rng):
        # 训练噪声预测时只抽取 1,...,T;t=0 留给边界测试。
        return rng.integers(1, len(schedule["beta"]), size=batch_size)
    
    def q_sample(x0, t, schedule, rng):
        x0 = np.asarray(x0, dtype=np.float64)
        t = np.asarray(t, dtype=np.int64)
        if t.shape != (x0.shape[0],):
            raise ValueError("每个样本需要一个时间下标")
        if np.any((t < 0) | (t >= len(schedule["alpha_bar"]))):
            raise IndexError("时间下标越界")
        # eps_tilde 对应解析式中的累计等价噪声。
        eps_tilde = np.zeros_like(x0)
        active = t > 0
        if np.any(active):
            eps_tilde[active] = rng.normal(size=x0[active].shape)
        shape = (-1,) + (1,) * (x0.ndim - 1)
        a = schedule["alpha_bar"][t].reshape(shape)
        return np.sqrt(a) * x0 + np.sqrt(1.0-a) * eps_tilde, eps_tilde
    
    def simple_loss(model, x0, schedule, rng):
        t = sample_timesteps(x0.shape[0], schedule, rng)
        xt, eps_tilde = q_sample(x0, t, schedule, rng)
        pred = model(xt, t)
        if pred.shape != x0.shape: raise ValueError("网络输出维度错误")
        return ((pred - eps_tilde) ** 2).mean()
    
    def p_sample(model, xt, t, schedule, rng):
        # t 是单个反向时间;一次调用处理整个批次。
        if not 1 <= t < len(schedule["beta"]):
            raise IndexError("反向时间下标越界")
        time = np.full(xt.shape[0], t, dtype=np.int64)
        eps_pred = model(xt, time)
        if eps_pred.shape != xt.shape: raise ValueError("网络输出维度错误")
        beta_t = schedule["beta"][t]
        alpha_t = schedule["alpha"][t]
        abar_t = schedule["alpha_bar"][t]
        mean = (xt - beta_t * eps_pred / np.sqrt(1.0-abar_t)) / np.sqrt(alpha_t)
        if t == 1:
            return mean, 1
        abar_prev = schedule["alpha_bar"][t-1]
        variance = beta_t * (1.0-abar_prev) / (1.0-abar_t)
        return mean + np.sqrt(variance) * rng.normal(size=xt.shape), 1
    

    p_sample 的第二个返回值是本步 NFE。完整采样器依次调用 \(t=T,T-1,\ldots,1\) 并累加该计数;在 \(t=1\) 时直接返回均值,不再抽取随机噪声。

  9. 逐步采样与解析采样不要求逐元素相同,但在大量样本下都应接近均值 \(\sqrt{\bar\alpha_t}\mathbb E(x_0)\) 和由公式给出的方差;使用置信区间判断差异而非要求精确相等。自动微分检查 应在 float64 小网络上抽查参数,并用 中心差分 尝试多个步长。把模型输出替换为直接构造 \(x_t\) 时保存的真实累计噪声 \(\tilde{\epsilon}_t\)⁠,恢复的 \(\widehat x_0\) 应与 \(x_0\) 在严格容差内一致。

  10. 检查 np.diff(alpha_bar) < 0 及有限性;最后反向一步不调用随机数生成器;批量输入、预测和目标同形;所有 \(\tilde\beta_t\geq0\)⁠;同种子重跑完全一致;负时间或超出参数数组长度时报错;维度为 \((m,)\) 的时间数组通过重整为 \(m\times1\times1\times1\) 作用于各自样本。另应拒绝非有限输入和不合法 \(\beta_t\)⁠。

  11. 两种参数设置只改变 \(\beta_t\) 序列,复用数据、种子、网络初始化、更新次数与采样器。质量可用 FID/KID 或玩具数据 MMD,并同时检查覆盖;报告多种子训练秒数、相同样本数的生成秒数、每个样本 NFE、参数量和训练/生成峰值内存。采样步数相同时 NFE 通常相同,质量差异来自噪声分配;不能让某种参数设置额外训练更久。

  12. 三种参数化必须正确转换到同一反向均值,且损失尺度按正文公式匹配;否则比较的是实现错误。固定数据、种子、U-Net 宽度、更新次数与采样时间表,报告质量与覆盖、极端时间步损失、训练及生成时间、NFE、参数量和峰值内存。网络结构相同时参数量和名义 NFE 相同,但数值稳定性与最终质量可能不同。

  13. 三个采样设置使用同一检查点和同一批初始噪声,训练时间与参数量相同,应注明为共享成本。每次反向步通常调用噪声网络一次,因此 NFE 近似等于实际步数。减少步数会近似线性降低生成时间,但时间子序列更稀会放大离散误差并可能降低质量或多样性。报告 FID/KID 等分布指标、固定样本生成秒数、NFE 和峰值内存,不能只展示筛选后的最好样本。