条件引导
前文的无条件模型学习整个数据分布 p d a t a ( z ) p_{data}(z) p d a t a ( z ) ,因此只能生成“像训练集”的样本,却不知道用户此刻想要其中哪一部分。文本、类别或其他条件 y ∈ Y y\in\mathcal Y y ∈ Y 的作用,是把这个大分布收缩为与条件相符的子分布 p d a t a ( ⋅ ∣ y ) p_{data}(\cdot\mid y) p d a t a ( ⋅ ∣ y ) 。为此,训练和推理时都将 y y y 输入向量场网络:
u θ : R d × Y × [ 0 , 1 ] → R d , ( x , y , t ) ↦ u t θ ( x ∣ y ) . 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).
u θ : R d × Y × [ 0 , 1 ] → R d , ( x , y , t ) ↦ u t θ ( x ∣ y ) .
同一个带噪状态 x x x 在不同条件下可能需要向完全不同的终点移动,网络必须同时看到 x x x 、噪声时刻 t t t 和条件 y y y ,才能判断当前应该保留哪些结构、朝哪一类样本修正。给定 y y y 后,先从简单先验采样 X 0 ∼ p i n i t X_0\sim p_{init} X 0 ∼ p ini t ,再从 t = 0 t=0 t = 0 积分到 t = 1 t=1 t = 1 :
d X t = u t θ ( X t ∣ y ) d t . dX_t=u_t^\theta(X_t\mid y)\,dt.
d X t = u t θ ( X t ∣ y ) d t .
若使用分数参数化的条件扩散模型,还要按所选 SDE 将分数转换为漂移。例如,SDE Extension 写成
d X t = [ u t θ ( X t ∣ y ) + σ t 2 2 s t θ ( X t ∣ y ) ] d t + σ t d W t . 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.
d X t = [ u t θ ( X t ∣ y ) + 2 σ t 2 s t θ ( X t ∣ y ) ] d t + σ t d W t .
其中 s t θ s_t^\theta s t θ 近似条件边缘分数。条件只改变每个中间状态应采用的动力学,并不改变“从简单先验逐步走到数据”的生成框架;训练完成后,固定不同的 y y y ,同一份初始噪声便会被引向不同的条件分布。
对于流匹配,数据集需要返回配对样本 ( z , y ) ∼ p d a t a ( z , y ) (z,y)\sim p_{data}(z,y) ( z , y ) ∼ p d a t a ( z , y ) 。在固定的条件概率路径 p t ( ⋅ ∣ z ) p_t(\cdot\mid z) p t ( ⋅ ∣ z ) 下,训练目标为
L G - C F M ( θ ) = E ( z , y ) ∼ p d a t a ( z , y ) , t ∼ U n i f [ 0 , 1 ] x ∼ p t ( ⋅ ∣ z ) ∥ u t θ ( x ∣ y ) − u t t a r g e t ( x ∣ z ) ∥ 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.
L G - CFM ( θ ) = E ( z , y ) ∼ p d a t a ( z , y ) , t ∼ Unif [ 0 , 1 ] x ∼ p t ( ⋅ ∣ z ) u t θ ( x ∣ y ) − u t t a r g e t ( x ∣ z ) 2 .
这里已知完整样本 z z z ,所以可以用它构造中间状态 x x x 和精确的条件路径速度 u t t a r g e t ( x ∣ z ) u_t^{target}(x\mid z) u t t a r g e t ( x ∣ z ) ;网络推理时看不到 z z z ,只能根据较粗的描述 y y y 推断应当朝哪些可能的终点移动。对所有具有同一条件的训练样本取平均后,平方损失的最优解正是条件边缘向量场,因此不需要在训练中真的从 p d a t a ( ⋅ ∣ y ) p_{data}(\cdot\mid y) p d a t a ( ⋅ ∣ y ) 启动一条生成轨迹。这个设计使训练仍然只需随机抽取一个时刻便可监督,但也留下了条件信息不足的问题:一段描述通常不能唯一确定图像,图文配对误差和模型欠拟合还会进一步稀释条件影响,直接条件化得到的样本便可能语义正确却不够“听指令”。
CFG
为了进一步增强条件影响,需要先弄清条件究竟把无条件动力学改向了哪里。对高斯概率路径
p t ( ⋅ ∣ z ) = N ( α t z , β t 2 I d ) . p_t(\cdot\mid z)=\mathcal N(\alpha_tz,\beta_t^2I_d).
p t ( ⋅ ∣ z ) = N ( α t z , β t 2 I d ) .
条件边缘向量场与条件分数之间有仿射关系:
u t t a r g e t ( x ∣ y ) = a t ∇ x log p t ( x ∣ y ) + b t x , u_t^{target}(x\mid y)
=
a_t\nabla_x\log p_t(x\mid y)+b_tx,
u t t a r g e t ( x ∣ y ) = a t ∇ x log p t ( x ∣ y ) + b t x ,
其中
a t = β t 2 α ˙ t α t − β ˙ t β t , b t = α ˙ 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}.
a t = β t 2 α t α ˙ t − β ˙ t β t , b t = α t α ˙ t .
其中 b t x b_tx b t x 是由当前状态和调度共同决定的公共运动,a t ∇ x log p t ( x ∣ y ) a_t\nabla_x\log p_t(x\mid y) a t ∇ x log p t ( x ∣ y ) 才携带“哪些方向更像条件数据”的信息。贝叶斯公式把条件分数拆为
∇ x log p t ( x ∣ y ) = ∇ x log p t ( x ) + ∇ x log p t ( y ∣ x ) . \nabla_x\log p_t(x\mid y)
=
\nabla_x\log p_t(x)+\nabla_x\log p_t(y\mid x).
∇ x log p t ( x ∣ y ) = ∇ x log p t ( x ) + ∇ x log p t ( y ∣ x ) .
右侧第一项把样本推向任意高概率数据,第二项则提高条件 y y y 在当前带噪状态下的似然。代回向量场关系后,公共的 b t x b_tx b t x 项相消,理想条件向量场可写成
u t t a r g e t ( x ∣ y ) = u t t a r g e t ( x ) + a t ∇ x log p t ( y ∣ x ) . u_t^{target}(x\mid y)
=
u_t^{target}(x)+a_t\nabla_x\log p_t(y\mid x).
u t t a r g e t ( x ∣ y ) = u t t a r g e t ( x ) + a t ∇ x log p t ( y ∣ x ) .
这说明“让样本更符合条件”可以具体化为沿 ∇ x log p t ( y ∣ x ) \nabla_x\log p_t(y\mid x) ∇ x log p t ( y ∣ x ) 移动:它是对 x x x 做微小改变时,条件似然增长最快的方向。Classifier guidance 将这项修正乘以引导尺度 w w w :
u ~ t ( x ∣ y ) = u t t a r g e t ( x ) + w a t ∇ x log p t ( y ∣ x ) . \widetilde u_t(x\mid y)
=
u_t^{target}(x)+wa_t\nabla_x\log p_t(y\mid x).
u t ( x ∣ y ) = u t t a r g e t ( x ) + w a t ∇ x log p t ( y ∣ x ) .
分类器必须在每个噪声时刻都能识别 y y y ,因为干净图像分类器面对几乎纯噪声的 x x x 时不会给出可靠梯度;开放文本还意味着条件集合不是有限类别,另训这样一个分类器既昂贵又难覆盖。CFG 利用同一个贝叶斯关系,把似然梯度写成条件分数与无条件分数之差:
∇ x log p t ( y ∣ x ) = ∇ x log p t ( x ∣ y ) − ∇ x log p t ( x ) , \nabla_x\log p_t(y\mid x)
=
\nabla_x\log p_t(x\mid y)-\nabla_x\log p_t(x),
∇ x log p t ( y ∣ x ) = ∇ x log p t ( x ∣ y ) − ∇ x log p t ( x ) ,
条件预测和无条件预测都包含把噪声拉回数据流形的共同部分,相减后这部分大致抵消,留下的正是“为了满足 y y y 还需额外改变什么”。由于高斯路径下分数与向量场只相差同一组仿射系数,这个差值可以直接在向量场上计算:
u ~ t t a r g e t ( x ∣ y ) = ( 1 − w ) u t t a r g e t ( x ∣ ∅ ) + w u t t a r g e t ( x ∣ y ) . \widetilde u_t^{target}(x\mid y)
=
(1-w)u_t^{target}(x\mid\emptyset)
+wu_t^{target}(x\mid y).
u t t a r g e t ( x ∣ y ) = ( 1 − w ) u t t a r g e t ( x ∣ ∅ ) + w u t t a r g e t ( x ∣ y ) .
实际采样时使用同一网络的两个预测:
u ~ t θ ( x ∣ y ) = u t θ ( x ∣ ∅ ) + w [ u t θ ( x ∣ y ) − u t θ ( 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].
u t θ ( x ∣ y ) = u t θ ( x ∣ ∅ ) + w [ u t θ ( x ∣ y ) − u t θ ( x ∣ ∅ ) ] .
这里第一项负责一般的去噪与成像,方括号内的差值负责条件修正,w w w 决定每一步愿意沿这条修正方向走多远。为了让两次预测处在同一参数化和同一误差尺度下,通常让一个网络同时学习二者;训练时以概率 η \eta η 丢弃条件。令
y ˉ = { y , 概率 1 − η , ∅ , 概率 η , \bar y=
\begin{cases}
y, & \text{概率 }1-\eta,\\
\emptyset, & \text{概率 }\eta,
\end{cases}
y ˉ = { y , ∅ , 概率 1 − η , 概率 η ,
则 CFG 的训练目标为
L C F G - C F M ( θ ) = E ( z , y ) ∼ p d a t a ( z , y ) , t ∼ U n i f [ 0 , 1 ] x ∼ p t ( ⋅ ∣ z ) , y ˉ ∼ D r o p η ( y ) ∥ u t θ ( x ∣ y ˉ ) − u t t a r g e t ( x ∣ z ) ∥ 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.
L CFG - CFM ( θ ) = E ( z , y ) ∼ p d a t a ( z , y ) , t ∼ Unif [ 0 , 1 ] x ∼ p t ( ⋅ ∣ z ) , y ˉ ∼ Drop η ( y ) u t θ ( x ∣ y ˉ ) − u t t a r g e t ( x ∣ z ) 2 .
若训练时从未见过空条件,推理时传入 ∅ \emptyset ∅ 只会得到没有语义保证的外推结果,条件差值也就失去贝叶斯解释。随机丢弃使带条件样本教会网络使用 y y y ,空条件样本则逼近对所有 y y y 边缘化后的数据动力学。一次训练迭代可按以下顺序进行:
从数据集中采样 ( z , y ) (z,y) ( z , y ) ,再采样 t ∼ U n i f [ 0 , 1 ] t\sim\mathrm{Unif}[0,1] t ∼ Unif [ 0 , 1 ] 和噪声 ϵ ∼ N ( 0 , I d ) \epsilon\sim\mathcal N(0,I_d) ϵ ∼ N ( 0 , I d ) 。
对高斯路径构造 x = α t z + β t ϵ x=\alpha_tz+\beta_t\epsilon x = α t z + β t ϵ 。
以概率 η \eta η 将 y y y 替换为 ∅ \emptyset ∅ ,得到 y ˉ \bar y y ˉ 。
计算 u t θ ( x ∣ y ˉ ) u_t^\theta(x\mid\bar y) u t θ ( x ∣ y ˉ ) 与 u t t a r g e t ( x ∣ z ) u_t^{target}(x\mid z) u t t a r g e t ( x ∣ z ) 的平方误差。
反向传播并更新网络参数。
引导尺度
训练时只随机保留或丢弃条件,并不使用大于 1 1 1 的引导尺度;w w w 是推理阶段控制条件强度的旋钮。先采样 X 0 ∼ p i n i t X_0\sim p_{init} X 0 ∼ p ini t ,在每个数值积分步分别计算 u t θ ( X t ∣ y ) u_t^\theta(X_t\mid y) u t θ ( X t ∣ y ) 和 u t θ ( X t ∣ ∅ ) u_t^\theta(X_t\mid\emptyset) u t θ ( X t ∣ ∅ ) ,再按 CFG 公式合成为 u ~ t θ ( X t ∣ y ) \widetilde u_t^\theta(X_t\mid y) u t θ ( X t ∣ y ) 。条件流模型直接求解
d X t = u ~ t θ ( X t ∣ y ) d t . dX_t=\widetilde u_t^\theta(X_t\mid y)\,dt.
d X t = u t θ ( X t ∣ y ) d t .
若使用 SDE Extension,还需同样组合条件与无条件分数,得到 s ~ t θ \widetilde s_t^\theta s t θ ,再求解
d X t = [ u ~ t θ ( X t ∣ y ) + σ t 2 2 s ~ t θ ( X t ∣ y ) ] d t + σ t d W t . 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.
d X t = [ u t θ ( X t ∣ y ) + 2 σ t 2 s t θ ( X t ∣ y ) ] d t + σ t d W t .
若模型参数化的是噪声、去噪结果或某个特定 SDE 的完整漂移,则应先确认该参数化和采样方程之间的线性关系,再在一致的量上组合条件与无条件预测。否则,同一个数值 w w w 可能对应不同强度,甚至把修正方向的符号弄反。
w = 0 w=0 w = 0 时只使用无条件模型,条件不起作用;w = 1 w=1 w = 1 恢复训练得到的普通条件模型;w > 1 w>1 w > 1 则越过条件预测,沿“条件相对无条件多做的那一步”继续外推。适度外推能放大文字与图像之间容易被平均掉的对应关系,但它已经离开训练时见过的向量场范围,终点分布也不再保证严格等于 p d a t a ( ⋅ ∣ y ) p_{data}(\cdot\mid y) p d a t a ( ⋅ ∣ y ) 。当 w w w 很大时,少数最能提高条件似然的模式会被反复强化,因而常以多样性下降、颜色过饱和或结构失真为代价。条件差值在不同时间和不同参数化下尺度并不恒定,合适的 w w w 需要随模型、调度与采样器共同验证。