条件扩散与引导:参考答案

目录

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

条件扩散与引导:参考答案#

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

说明#

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

  1. 取对数得 \(\log p_t(\bx_t\mid c)=\log p_t(\bx_t)+\log p_t(c\mid\bx_t)-\log p_t(c)\)⁠。最后一项与 \(\bx_t\) 无关,求梯度后得到 \(\nabla_{\bx_t}\log p_t(\bx_t\mid c)=\nabla_{\bx_t}\log p_t(\bx_t)+\nabla_{\bx_t}\log p_{\boldsymbol\phi}(c\mid\bx_t,t)\)⁠。分类器必须在各噪声等级上估计含噪数据的类别概率;若只在干净图像上训练,其梯度不对应采样中的 \(p_t(c\mid\bx_t)\)⁠。

  2. \(w=0\) 时得到 \(\boldsymbol\epsilon_{\mathrm{uncond}}\)⁠;\(w=1\) 时得到 \(\boldsymbol\epsilon_{\mathrm{cond}}\)⁠;\(w>1\) 时可写为 \((1-w)\boldsymbol\epsilon_{\mathrm{uncond}}+w\boldsymbol\epsilon_{\mathrm{cond}}\)⁠,无条件项系数为负,因此不是凸组合。它沿 \(\boldsymbol\epsilon_{\mathrm{cond}}-\boldsymbol\epsilon_{\mathrm{uncond}}\) 方向越过条件预测,常增强条件一致性,也可能降低多样性或产生伪影。

  3. 差向量为 \((2,3)\)⁠。三种结果依次为 \((1,-1)+0.5(2,3)=(2,0.5)\)⁠、\((3,2)\)\((5,5)\)⁠。相对无条件预测的改变量范数为 \(w\sqrt{2^2+3^2}=w\sqrt{13}\)⁠,因此依次为 \(\sqrt{13}/2\)⁠、\(\sqrt{13}\)⁠、\(2\sqrt{13}\)⁠。

  4. 普通条件扩散的噪声网络 NFE 为 50,CFG 为 \(2(50)=100\) 次样本等价评估,分类器引导的噪声网络 NFE 为 50,另有 50 次分类器前向传播与后向梯度计算。拼接 CFG 可把两次 Python 调用合成一次、提高设备利用率,但批量大小翻倍,网络仍处理 100 个单批等价输入;因此应同时报告调用次数、样本等价 NFE 和实际时间。

  5. 实现应把数值组合与网络调用分开:

    def cfg(eps_u, eps_c, scale):
        if eps_u.shape != eps_c.shape: raise ValueError("预测必须同形")
        if not np.isfinite(scale) or scale < 0: raise ValueError("引导强度错误")
        return eps_u + scale * (eps_c - eps_u)
    
    def cfg_predict(model, x, t, cond, null_cond, scale):
        # 一次程序调用处理两倍批量;样本等价 NFE 仍为 2。
        t = np.broadcast_to(np.asarray(t), (x.shape[0],))
        x_pair = np.concatenate([x, x], axis=0)
        t_pair = np.concatenate([t, t], axis=0)
        c_pair = np.concatenate([null_cond, cond], axis=0)
        prediction = model(x_pair, t_pair, c_pair)
        eps_u, eps_c = np.split(prediction, 2, axis=0)
        return cfg(eps_u, eps_c, scale), {
            "network_calls": 1,
            "sample_equivalent_nfe": 2,
        }
    
    def drop_condition(cond, probability, rng, null_value=0):
        keep = rng.random(cond.shape[0]) >= probability
        out = cond.copy(); out[~keep] = null_value
        return out, keep
    
    def composite(known, generated, mask):
        if np.any((mask < 0) | (mask > 1)): raise ValueError("掩码越界")
        return mask * known + (1-mask) * generated
    

    完整采样器在每一步累加 network_callssample_equivalent_nfe若改为分开执行条件与无条件分支,则两项分别增加 2 和 2。这样可以区分程序调用次数、样本等价计算量与实际墙钟时间。

  6. 若分类器的类别对数概率是 \(\log p(c\mid\bx)=\ba\trans\cdot\bx+b\)⁠,手算梯度就是 \(\ba\)⁠;更一般的线性 Softmax 分类器可按交叉熵导数手算。自动微分结果 应与之逐元素一致。CFG 拼接对照必须关闭 Dropout、冻结 BatchNorm 统计量,并按批次轴把无条件与条件输入拼接后再拆开;若结果不同,通常说明模型仍在训练模式或条件索引错位。

  7. 直接断言 cfg(u,c,0)==ucfg(u,c,1)==ccfg(u,u,w)==u全 0 掩码返回生成区域,全 1 掩码返回已知区域;大量条件样本的丢弃比例应在二项波动容差内;同种子得到相同丢弃掩码和采样;用网络包装器记录调用,普通 CFG 每步两次,拼接版每步一次程序调用但样本等价计数加 2。非法维度、概率、引导强度和掩码应报错。

  8. 所有强度使用同一检查点、条件、初始噪声和时间表,训练时间与参数量为共享常数。通常增大 \(w\) 会提高条件一致性但可能降低多样性,最优值依数据而定。报告 FID/KID、条件分类或对齐指标、覆盖/重复率、固定样本生成秒数、样本等价 NFE 和峰值内存;展示样本不得只挑最好结果。

  9. 分类器引导需要额外训练一个能够处理含噪输入的分类器,CFG 则在扩散训练中随机丢弃条件,并在采样时分别进行有条件和无条件预测。比较这两种方法时,应使用相同的扩散数据、随机种子、更新次数和采样步数,并把分类器的训练时间、参数量和内存计入总成本。报告生成质量、条件准确性、多样性、总训练与生成时间、噪声网络及分类器的函数运行次数、总参数量和峰值内存。

  10. 三种网络尽量匹配主干宽度和总参数量,复用划分、种子、更新次数、条件丢弃率、CFG 强度和采样时间表。拼接适合空间对齐条件,FiLM 用条件控制通道尺度与平移,交叉注意力适合可变长度序列但注意力内存随词元数增长。报告质量与条件一致性、训练秒数、固定样本生成秒数、NFE、参数量和峰值内存,不能只按参数量预测实际速度。