Skip to content
huc
Go back

扩散模型的条件引导

目录 1 / 5

《生成扩散模型漫谈(九):条件控制生成结果》 —— 苏剑林

条件生成改变的是反向分布

无条件生成学习 p(xt1xt)p(x_{t-1}\mid x_t),加入条件 yy 后,目标变成 p(xt1xt,y)p(x_{t-1}\mid x_t,y)。两类常见方法的差异在于条件信息何时进入模型:

  • Classifier Guidance 保留已有的无条件扩散模型,在采样时用额外判别模型修改反向一步;
  • Classifier-Free Guidance 从训练开始就让去噪网络接收 yy,采样时组合条件与无条件预测。

前者可以复用无条件模型,但每一步都要额外求条件模型对输入的梯度;后者推理结构更直接,却需要带条件数据重新训练扩散模型。

贝叶斯分解怎样引出分类器梯度

先在每一步应用条件贝叶斯:

p(xt1xt,y)=p(xt1xt)p(yxt1,xt)p(yxt).(1)p(x_{t-1}\mid x_t,y) =p(x_{t-1}\mid x_t) \frac{p(y\mid x_{t-1},x_t)}{p(y\mid x_t)}. \tag{1}

正向马尔可夫链满足 yxt1xty\to x_{t-1}\to x_t:一旦给定较干净的 xt1x_{t-1},再观察由它加噪得到的 xtx_t 不应提供额外的 yy 信息。因此

p(yxt1,xt)=p(yxt1).p(y\mid x_{t-1},x_t)=p(y\mid x_{t-1}).

把公式 (1) 写成指数形式:

p(xt1xt,y)p(xt1xt)exp[logp(yxt1)logp(yxt)].p(x_{t-1}\mid x_t,y) \propto p(x_{t-1}\mid x_t) \exp\left[ \log p(y\mid x_{t-1})-\log p(y\mid x_t) \right].

当反向一步足够小时,xt1x_{t-1} 接近 xtx_t,所以在 xtx_t 处做一阶展开:

logp(yxt1)logp(yxt)(xt1xt)xtlogp(yxt).(2)\log p(y\mid x_{t-1})-\log p(y\mid x_t) \approx (x_{t-1}-x_t)^\top \nabla_{x_t}\log p(y\mid x_t). \tag{2}

这里关于时间变化的项与待采样变量 xt1x_{t-1} 无关,会被归一化常数吸收。若无条件反向核是

p(xt1xt)=N(xt1;μt(xt),Σt),p(x_{t-1}\mid x_t) =\mathcal N(x_{t-1};\mu_t(x_t),\Sigma_t),

把公式 (2) 代入并对高斯指数配方,就得到

p(xt1xt,y)N(xt1;μt(xt)+Σtxtlogp(yxt),Σt).(3)p(x_{t-1}\mid x_t,y) \approx \mathcal N\left( x_{t-1}; \mu_t(x_t)+\Sigma_t\nabla_{x_t}\log p(y\mid x_t), \Sigma_t \right). \tag{3}

各向同性协方差 Σt=σt2I\Sigma_t=\sigma_t^2I 时,Classifier Guidance 的采样式为

xt1=μt(xt)+σt2xtlogp(yxt)+σtz,zN(0,I).(4)x_{t-1} =\mu_t(x_t) +\sigma_t^2\nabla_{x_t}\log p(y\mid x_t) +\sigma_t z, \qquad z\sim\mathcal N(0,I). \tag{4}

梯度乘的是方差而不是标准差,因为它来自高斯自然参数的平移。若 Σt\Sigma_t 不是单位阵的倍数,正确修正是 Σtlogp(yxt)\Sigma_t\nabla\log p(y\mid x_t),不同方向会按反向核的不确定性缩放。

分类器必须理解带噪输入 xtx_t。直接把只见过干净样本的分类器用于高噪声状态,会产生分布外梯度;一种近似补救是先用 x^0(xt)\hat x_0(x_t) 去噪再分类,但此时反向传播还会穿过 x^0\hat x_0,得到的已不是严格的 xtlogp(yxt)\nabla_{x_t}\log p(y\mid x_t)

Guidance Scale 控制条件能量

实际采样会加入强度 γ\gamma

xt1=μt(xt)+γΣtxtlogp(yxt)+Σt1/2z.(5)x_{t-1} =\mu_t(x_t) +\gamma\Sigma_t\nabla_{x_t}\log p(y\mid x_t) +\Sigma_t^{1/2}z. \tag{5}

γ>1\gamma>1 会更强地把样本推向条件模型的高分区域,通常提高条件一致性,同时压缩多样性。不能简单把它严格解释为从归一化分布

p~(yxt)=p(yxt)γZ(xt)\tilde p(y\mid x_t) =\frac{p(y\mid x_t)^\gamma}{Z(x_t)}

采样,因为

xtlogp~(yxt)=γxtlogp(yxt)xtlogZ(xt),\nabla_{x_t}\log\tilde p(y\mid x_t) =\gamma\nabla_{x_t}\log p(y\mid x_t) -\nabla_{x_t}\log Z(x_t),

γ1\gamma\ne1Z(xt)=yp(yxt)γZ(x_t)=\sum_y p(y\mid x_t)^\gamma 一般依赖 xtx_t。忽略第二项后,缩放梯度更准确的理解是人为调节条件能量,而不是某个已归一化分类分布的精确 score。

这个角度也允许把类别概率推广为任意可微相似度 s(xt,y)s(x_t,y)。直接定义能量重加权

pγ(xt1xt,y)p(xt1xt)exp[γs(xt1,y)],p_\gamma(x_{t-1}\mid x_t,y) \propto p(x_{t-1}\mid x_t) \exp\left[\gamma s(x_{t-1},y)\right],

再对 s(xt1,y)s(x_{t-1},y) 局部展开,可得

pγ(xt1xt,y)N(xt1;μt(xt)+γΣtxts(xt,y),Σt).(6)p_\gamma(x_{t-1}\mid x_t,y) \approx \mathcal N\left( x_{t-1}; \mu_t(x_t)+\gamma\Sigma_t\nabla_{x_t}s(x_t,y), \Sigma_t \right). \tag{6}

yy 因而可以是类别、文本或参考图像;ss 可以来自共享 embedding 的余弦相似度,也可以是适合具体任务的可微度量。关键条件仍然是度量网络要能处理当前噪声水平。

条件 score 统一随机与确定性采样

离散配方公式 (3) 的修正含 Σt\Sigma_t,于是当确定性 DDIM 取 Σt=0\Sigma_t=0 时,看起来 Guidance 会消失。原因是这个公式从一个带方差的局部高斯核推导,不是确定性概率流的统一表达。

在 score 框架中,条件贝叶斯直接给出

xlogpt(xy)=xlogpt(x)+xlogpt(yx).(7)\nabla_x\log p_t(x\mid y) =\nabla_x\log p_t(x) +\nabla_x\log p_t(y\mid x). \tag{7}

无论反向过程选择随机 SDE 还是 probability flow ODE,只需把无条件 score sθ(x,t)s_\theta(x,t) 换成公式 (7)。若采用

sθ(xt,t)=εθ(xt,t)βˉt,s_\theta(x_t,t) =-\frac{\varepsilon_\theta(x_t,t)}{\bar\beta_t},

那么等价的噪声预测替换是

εθguided(xt,t,y)=εθ(xt,t)γβˉtxtlogpt(yxt).(8)\varepsilon_\theta^{\mathrm{guided}}(x_t,t,y) =\varepsilon_\theta(x_t,t) -\gamma\bar\beta_t \nabla_{x_t}\log p_t(y\mid x_t). \tag{8}

公式 (8) 不依赖采样方差是否为零,所以也适用于确定性 DDIM/ODE。离散均值平移和条件 score 替换不是两套无关技巧:在小步长、高斯反向核下,前者正是后者的一阶离散结果。

Classifier-Free Guidance 如何得到两次预测

Classifier-Free 模型直接训练条件噪声预测器:

Lcond=Ex0,y,t,ε[εεθ(αˉtx0+βˉtε,y,t)2].(9)L_{\mathrm{cond}} =\mathbb E_{x_0,y,t,\varepsilon} \left[ \left\| \varepsilon- \varepsilon_\theta( \bar\alpha_t x_0+\bar\beta_t\varepsilon, y,t ) \right\|^2 \right]. \tag{9}

训练时以一定概率把 yy 替换为空条件 \varnothing。同一网络于是同时学习

εθ(xt,y,t)εθ(xt,,t).\varepsilon_\theta(x_t,y,t) \quad\text{和}\quad \varepsilon_\theta(x_t,\varnothing,t).

条件 score 与无条件 score 的差满足

xlogpt(xy)xlogpt(x)=xlogpt(yx)=εθ(xt,y,t)εθ(xt,,t)βˉt.(10)\begin{aligned} \nabla_x\log p_t(x\mid y)-\nabla_x\log p_t(x) &=\nabla_x\log p_t(y\mid x)\\ &=-\frac{ \varepsilon_\theta(x_t,y,t) -\varepsilon_\theta(x_t,\varnothing,t) }{\bar\beta_t}. \end{aligned} \tag{10}

因此不需要显式分类器,也可以用两次网络预测构造引导:

ε~θ(xt,y,t)=(1+w)εθ(xt,y,t)wεθ(xt,,t)(11)\boxed{ \tilde\varepsilon_\theta(x_t,y,t) =(1+w)\varepsilon_\theta(x_t,y,t) -w\varepsilon_\theta(x_t,\varnothing,t) } \tag{11}

其中 w=0w=0 是普通条件预测,w>0w>0 把预测沿“无条件 \to 条件”的方向外推。若把 γ=1+w\gamma=1+w,也可写成 εuncond+γ(εcondεuncond)\varepsilon_{\mathrm{uncond}}+\gamma (\varepsilon_{\mathrm{cond}}-\varepsilon_{\mathrm{uncond}})

这种线性外推与 Classifier Guidance 一样会以多样性换条件一致性。它并不保证对应某个归一化概率分布的精确 score,尤其当 scale 很大、网络预测又有误差时,外推结果可能离开训练时见过的 score 范围。空条件训练比例与 guidance scale 因此共同决定无条件基线是否可靠以及外推能走多远。