条件引导

前文的无条件模型学习整个数据分布 pdata(z)p_{data}(z),因此只能生成“像训练集”的样本,却不知道用户此刻想要其中哪一部分。文本、类别或其他条件 yYy\in\mathcal Y 的作用,是把这个大分布收缩为与条件相符的子分布 pdata(y)p_{data}(\cdot\mid y)。为此,训练和推理时都将 yy 输入向量场网络:

uθ:Rd×Y×[0,1]Rd,(x,y,t)utθ(xy).u^\theta:\mathbb R^d\times\mathcal Y\times[0,1]\rightarrow\mathbb R^d, \qquad (x,y,t)\mapsto u_t^\theta(x\mid y).

同一个带噪状态 xx 在不同条件下可能需要向完全不同的终点移动,网络必须同时看到 xx、噪声时刻 tt 和条件 yy,才能判断当前应该保留哪些结构、朝哪一类样本修正。给定 yy 后,先从简单先验采样 X0pinitX_0\sim p_{init},再从 t=0t=0 积分到 t=1t=1

dXt=utθ(Xty)dt.dX_t=u_t^\theta(X_t\mid y)\,dt.

若使用分数参数化的条件扩散模型,还要按所选 SDE 将分数转换为漂移。例如,SDE Extension 写成

dXt=[utθ(Xty)+σt22stθ(Xty)]dt+σtdWt.dX_t= \left[ u_t^\theta(X_t\mid y) +\frac{\sigma_t^2}{2}s_t^\theta(X_t\mid y) \right]dt +\sigma_t\,dW_t.

其中 stθs_t^\theta 近似条件边缘分数。条件只改变每个中间状态应采用的动力学,并不改变“从简单先验逐步走到数据”的生成框架;训练完成后,固定不同的 yy,同一份初始噪声便会被引向不同的条件分布。

对于流匹配,数据集需要返回配对样本 (z,y)pdata(z,y)(z,y)\sim p_{data}(z,y)。在固定的条件概率路径 pt(z)p_t(\cdot\mid z) 下,训练目标为

LG-CFM(θ)=E(z,y)pdata(z,y),tUnif[0,1]xpt(z)utθ(xy)uttarget(xz)2.\mathcal L_{\mathrm{G\text{-}CFM}}(\theta) = \mathbb E_{\substack{ (z,y)\sim p_{data}(z,y),\,t\sim\mathrm{Unif}[0,1]\\ x\sim p_t(\cdot\mid z) }} \left\| u_t^\theta(x\mid y)-u_t^{target}(x\mid z) \right\|^2.

这里已知完整样本 zz,所以可以用它构造中间状态 xx 和精确的条件路径速度 uttarget(xz)u_t^{target}(x\mid z);网络推理时看不到 zz,只能根据较粗的描述 yy 推断应当朝哪些可能的终点移动。对所有具有同一条件的训练样本取平均后,平方损失的最优解正是条件边缘向量场,因此不需要在训练中真的从 pdata(y)p_{data}(\cdot\mid y) 启动一条生成轨迹。这个设计使训练仍然只需随机抽取一个时刻便可监督,但也留下了条件信息不足的问题:一段描述通常不能唯一确定图像,图文配对误差和模型欠拟合还会进一步稀释条件影响,直接条件化得到的样本便可能语义正确却不够“听指令”。

CFG

为了进一步增强条件影响,需要先弄清条件究竟把无条件动力学改向了哪里。对高斯概率路径

pt(z)=N(αtz,βt2Id).p_t(\cdot\mid z)=\mathcal N(\alpha_tz,\beta_t^2I_d).

条件边缘向量场与条件分数之间有仿射关系:

uttarget(xy)=atxlogpt(xy)+btx,u_t^{target}(x\mid y) = a_t\nabla_x\log p_t(x\mid y)+b_tx,

其中

at=βt2α˙tαtβ˙tβt,bt=α˙tαt.a_t=\beta_t^2\frac{\dot\alpha_t}{\alpha_t}-\dot\beta_t\beta_t, \qquad b_t=\frac{\dot\alpha_t}{\alpha_t}.

其中 btxb_tx 是由当前状态和调度共同决定的公共运动,atxlogpt(xy)a_t\nabla_x\log p_t(x\mid y) 才携带“哪些方向更像条件数据”的信息。贝叶斯公式把条件分数拆为

xlogpt(xy)=xlogpt(x)+xlogpt(yx).\nabla_x\log p_t(x\mid y) = \nabla_x\log p_t(x)+\nabla_x\log p_t(y\mid x).

右侧第一项把样本推向任意高概率数据,第二项则提高条件 yy 在当前带噪状态下的似然。代回向量场关系后,公共的 btxb_tx 项相消,理想条件向量场可写成

uttarget(xy)=uttarget(x)+atxlogpt(yx).u_t^{target}(x\mid y) = u_t^{target}(x)+a_t\nabla_x\log p_t(y\mid x).

这说明“让样本更符合条件”可以具体化为沿 xlogpt(yx)\nabla_x\log p_t(y\mid x) 移动:它是对 xx 做微小改变时,条件似然增长最快的方向。Classifier guidance 将这项修正乘以引导尺度 ww

u~t(xy)=uttarget(x)+watxlogpt(yx).\widetilde u_t(x\mid y) = u_t^{target}(x)+wa_t\nabla_x\log p_t(y\mid x).

分类器必须在每个噪声时刻都能识别 yy,因为干净图像分类器面对几乎纯噪声的 xx 时不会给出可靠梯度;开放文本还意味着条件集合不是有限类别,另训这样一个分类器既昂贵又难覆盖。CFG 利用同一个贝叶斯关系,把似然梯度写成条件分数与无条件分数之差:

xlogpt(yx)=xlogpt(xy)xlogpt(x),\nabla_x\log p_t(y\mid x) = \nabla_x\log p_t(x\mid y)-\nabla_x\log p_t(x),

条件预测和无条件预测都包含把噪声拉回数据流形的共同部分,相减后这部分大致抵消,留下的正是“为了满足 yy 还需额外改变什么”。由于高斯路径下分数与向量场只相差同一组仿射系数,这个差值可以直接在向量场上计算:

u~ttarget(xy)=(1w)uttarget(x)+wuttarget(xy).\widetilde u_t^{target}(x\mid y) = (1-w)u_t^{target}(x\mid\emptyset) +wu_t^{target}(x\mid y).

实际采样时使用同一网络的两个预测:

u~tθ(xy)=utθ(x)+w[utθ(xy)utθ(x)].\widetilde u_t^\theta(x\mid y) = u_t^\theta(x\mid\emptyset) +w\left[ u_t^\theta(x\mid y)-u_t^\theta(x\mid\emptyset) \right].

这里第一项负责一般的去噪与成像,方括号内的差值负责条件修正,ww 决定每一步愿意沿这条修正方向走多远。为了让两次预测处在同一参数化和同一误差尺度下,通常让一个网络同时学习二者;训练时以概率 η\eta 丢弃条件。令

yˉ={y,概率 1η,,概率 η,\bar y= \begin{cases} y, & \text{概率 }1-\eta,\\ \emptyset, & \text{概率 }\eta, \end{cases}

则 CFG 的训练目标为

LCFG-CFM(θ)=E(z,y)pdata(z,y),tUnif[0,1]xpt(z),yˉDropη(y)utθ(xyˉ)uttarget(xz)2.\mathcal L_{\mathrm{CFG\text{-}CFM}}(\theta) = \mathbb E_{\substack{ (z,y)\sim p_{data}(z,y),\,t\sim\mathrm{Unif}[0,1]\\ x\sim p_t(\cdot\mid z),\,\bar y\sim\mathrm{Drop}_\eta(y) }} \left\| u_t^\theta(x\mid\bar y)-u_t^{target}(x\mid z) \right\|^2.

若训练时从未见过空条件,推理时传入 \emptyset 只会得到没有语义保证的外推结果,条件差值也就失去贝叶斯解释。随机丢弃使带条件样本教会网络使用 yy,空条件样本则逼近对所有 yy 边缘化后的数据动力学。一次训练迭代可按以下顺序进行:

  1. 从数据集中采样 (z,y)(z,y),再采样 tUnif[0,1]t\sim\mathrm{Unif}[0,1] 和噪声 ϵN(0,Id)\epsilon\sim\mathcal N(0,I_d)
  2. 对高斯路径构造 x=αtz+βtϵx=\alpha_tz+\beta_t\epsilon
  3. 以概率 η\etayy 替换为 \emptyset,得到 yˉ\bar y
  4. 计算 utθ(xyˉ)u_t^\theta(x\mid\bar y)uttarget(xz)u_t^{target}(x\mid z) 的平方误差。
  5. 反向传播并更新网络参数。

引导尺度

训练时只随机保留或丢弃条件,并不使用大于 11 的引导尺度;ww 是推理阶段控制条件强度的旋钮。先采样 X0pinitX_0\sim p_{init},在每个数值积分步分别计算 utθ(Xty)u_t^\theta(X_t\mid y)utθ(Xt)u_t^\theta(X_t\mid\emptyset),再按 CFG 公式合成为 u~tθ(Xty)\widetilde u_t^\theta(X_t\mid y)。条件流模型直接求解

dXt=u~tθ(Xty)dt.dX_t=\widetilde u_t^\theta(X_t\mid y)\,dt.

若使用 SDE Extension,还需同样组合条件与无条件分数,得到 s~tθ\widetilde s_t^\theta,再求解

dXt=[u~tθ(Xty)+σt22s~tθ(Xty)]dt+σtdWt.dX_t= \left[ \widetilde u_t^\theta(X_t\mid y) +\frac{\sigma_t^2}{2}\widetilde s_t^\theta(X_t\mid y) \right]dt +\sigma_t\,dW_t.

若模型参数化的是噪声、去噪结果或某个特定 SDE 的完整漂移,则应先确认该参数化和采样方程之间的线性关系,再在一致的量上组合条件与无条件预测。否则,同一个数值 ww 可能对应不同强度,甚至把修正方向的符号弄反。

w=0w=0 时只使用无条件模型,条件不起作用;w=1w=1 恢复训练得到的普通条件模型;w>1w>1 则越过条件预测,沿“条件相对无条件多做的那一步”继续外推。适度外推能放大文字与图像之间容易被平均掉的对应关系,但它已经离开训练时见过的向量场范围,终点分布也不再保证严格等于 pdata(y)p_{data}(\cdot\mid y)。当 ww 很大时,少数最能提高条件似然的模式会被反复强化,因而常以多样性下降、颜色过饱和或结构失真为代价。条件差值在不同时间和不同参数化下尺度并不恒定,合适的 ww 需要随模型、调度与采样器共同验证。