条件扩散与引导:参考答案#
说明#
以下答案与正文10道题逐题对应。推导题给出主要步骤;编程和实验题给出参考程序和检查方法,并说明结果适用于哪些条件。
取对数得 \(\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)\)。
\(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}}\) 方向越过条件预测,常增强条件一致性,也可能降低多样性或产生伪影。
差向量为 \((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}\)。
普通条件扩散的噪声网络 NFE 为 50,CFG 为 \(2(50)=100\) 次样本等价评估,分类器引导的噪声网络 NFE 为 50,另有 50 次分类器前向传播与后向梯度计算。拼接 CFG 可把两次 Python 调用合成一次、提高设备利用率,但批量大小翻倍,网络仍处理 100 个单批等价输入;因此应同时报告调用次数、样本等价 NFE 和实际时间。
实现应把数值组合与网络调用分开:
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_calls与sample_equivalent_nfe;若改为分开执行条件与无条件分支,则两项分别增加 2 和 2。这样可以区分程序调用次数、样本等价计算量与实际墙钟时间。若分类器的类别对数概率是 \(\log p(c\mid\bx)=\ba\trans\cdot\bx+b\),手算梯度就是 \(\ba\);更一般的线性 Softmax 分类器可按交叉熵导数手算。自动微分结果 应与之逐元素一致。CFG 拼接对照必须关闭 Dropout、冻结 BatchNorm 统计量,并按批次轴把无条件与条件输入拼接后再拆开;若结果不同,通常说明模型仍在训练模式或条件索引错位。
直接断言
cfg(u,c,0)==u、cfg(u,c,1)==c和cfg(u,u,w)==u;全 0 掩码返回生成区域,全 1 掩码返回已知区域;大量条件样本的丢弃比例应在二项波动容差内;同种子得到相同丢弃掩码和采样;用网络包装器记录调用,普通 CFG 每步两次,拼接版每步一次程序调用但样本等价计数加 2。非法维度、概率、引导强度和掩码应报错。所有强度使用同一检查点、条件、初始噪声和时间表,训练时间与参数量为共享常数。通常增大 \(w\) 会提高条件一致性但可能降低多样性,最优值依数据而定。报告 FID/KID、条件分类或对齐指标、覆盖/重复率、固定样本生成秒数、样本等价 NFE 和峰值内存;展示样本不得只挑最好结果。
分类器引导需要额外训练一个能够处理含噪输入的分类器,CFG 则在扩散训练中随机丢弃条件,并在采样时分别进行有条件和无条件预测。比较这两种方法时,应使用相同的扩散数据、随机种子、更新次数和采样步数,并把分类器的训练时间、参数量和内存计入总成本。报告生成质量、条件准确性、多样性、总训练与生成时间、噪声网络及分类器的函数运行次数、总参数量和峰值内存。
三种网络尽量匹配主干宽度和总参数量,复用划分、种子、更新次数、条件丢弃率、CFG 强度和采样时间表。拼接适合空间对齐条件,FiLM 用条件控制通道尺度与平移,交叉注意力适合可变长度序列但注意力内存随词元数增长。报告质量与条件一致性、训练秒数、固定样本生成秒数、NFE、参数量和峰值内存,不能只按参数量预测实际速度。