《生成扩散模型漫谈(九):条件控制生成结果》 —— 苏剑林
条件生成改变的是反向分布
无条件生成学习 p(xt−1∣xt),加入条件 y 后,目标变成
p(xt−1∣xt,y)。两类常见方法的差异在于条件信息何时进入模型:
- Classifier Guidance 保留已有的无条件扩散模型,在采样时用额外判别模型修改反向一步;
- Classifier-Free Guidance 从训练开始就让去噪网络接收 y,采样时组合条件与无条件预测。
前者可以复用无条件模型,但每一步都要额外求条件模型对输入的梯度;后者推理结构更直接,却需要带条件数据重新训练扩散模型。
贝叶斯分解怎样引出分类器梯度
先在每一步应用条件贝叶斯:
p(xt−1∣xt,y)=p(xt−1∣xt)p(y∣xt)p(y∣xt−1,xt).(1)
正向马尔可夫链满足 y→xt−1→xt:一旦给定较干净的
xt−1,再观察由它加噪得到的 xt 不应提供额外的 y 信息。因此
p(y∣xt−1,xt)=p(y∣xt−1).
把公式 (1) 写成指数形式:
p(xt−1∣xt,y)∝p(xt−1∣xt)exp[logp(y∣xt−1)−logp(y∣xt)].
当反向一步足够小时,xt−1 接近 xt,所以在 xt 处做一阶展开:
logp(y∣xt−1)−logp(y∣xt)≈(xt−1−xt)⊤∇xtlogp(y∣xt).(2)
这里关于时间变化的项与待采样变量 xt−1 无关,会被归一化常数吸收。若无条件反向核是
p(xt−1∣xt)=N(xt−1;μt(xt),Σt),
把公式 (2) 代入并对高斯指数配方,就得到
p(xt−1∣xt,y)≈N(xt−1;μt(xt)+Σt∇xtlogp(y∣xt),Σt).(3)
各向同性协方差 Σt=σt2I 时,Classifier Guidance 的采样式为
xt−1=μt(xt)+σt2∇xtlogp(y∣xt)+σtz,z∼N(0,I).(4)
梯度乘的是方差而不是标准差,因为它来自高斯自然参数的平移。若
Σt 不是单位阵的倍数,正确修正是
Σt∇logp(y∣xt),不同方向会按反向核的不确定性缩放。
分类器必须理解带噪输入 xt。直接把只见过干净样本的分类器用于高噪声状态,会产生分布外梯度;一种近似补救是先用
x^0(xt) 去噪再分类,但此时反向传播还会穿过 x^0,得到的已不是严格的
∇xtlogp(y∣xt)。
Guidance Scale 控制条件能量
实际采样会加入强度 γ:
xt−1=μt(xt)+γΣt∇xtlogp(y∣xt)+Σt1/2z.(5)
γ>1 会更强地把样本推向条件模型的高分区域,通常提高条件一致性,同时压缩多样性。不能简单把它严格解释为从归一化分布
p~(y∣xt)=Z(xt)p(y∣xt)γ
采样,因为
∇xtlogp~(y∣xt)=γ∇xtlogp(y∣xt)−∇xtlogZ(xt),
而 γ=1 时 Z(xt)=∑yp(y∣xt)γ 一般依赖
xt。忽略第二项后,缩放梯度更准确的理解是人为调节条件能量,而不是某个已归一化分类分布的精确 score。
这个角度也允许把类别概率推广为任意可微相似度 s(xt,y)。直接定义能量重加权
pγ(xt−1∣xt,y)∝p(xt−1∣xt)exp[γs(xt−1,y)],
再对 s(xt−1,y) 局部展开,可得
pγ(xt−1∣xt,y)≈N(xt−1;μt(xt)+γΣt∇xts(xt,y),Σt).(6)
y 因而可以是类别、文本或参考图像;s 可以来自共享 embedding 的余弦相似度,也可以是适合具体任务的可微度量。关键条件仍然是度量网络要能处理当前噪声水平。
条件 score 统一随机与确定性采样
离散配方公式 (3) 的修正含
Σt,于是当确定性 DDIM 取 Σt=0 时,看起来 Guidance 会消失。原因是这个公式从一个带方差的局部高斯核推导,不是确定性概率流的统一表达。
在 score 框架中,条件贝叶斯直接给出
∇xlogpt(x∣y)=∇xlogpt(x)+∇xlogpt(y∣x).(7)
无论反向过程选择随机 SDE 还是 probability flow ODE,只需把无条件 score
sθ(x,t) 换成公式 (7)。若采用
sθ(xt,t)=−βˉtεθ(xt,t),
那么等价的噪声预测替换是
εθguided(xt,t,y)=εθ(xt,t)−γβˉt∇xtlogpt(y∣xt).(8)
公式 (8) 不依赖采样方差是否为零,所以也适用于确定性 DDIM/ODE。离散均值平移和条件 score 替换不是两套无关技巧:在小步长、高斯反向核下,前者正是后者的一阶离散结果。
Classifier-Free Guidance 如何得到两次预测
Classifier-Free 模型直接训练条件噪声预测器:
Lcond=Ex0,y,t,ε[ε−εθ(αˉtx0+βˉtε,y,t)2].(9)
训练时以一定概率把 y 替换为空条件 ∅。同一网络于是同时学习
εθ(xt,y,t)和εθ(xt,∅,t).
条件 score 与无条件 score 的差满足
∇xlogpt(x∣y)−∇xlogpt(x)=∇xlogpt(y∣x)=−βˉtεθ(xt,y,t)−εθ(xt,∅,t).(10)
因此不需要显式分类器,也可以用两次网络预测构造引导:
ε~θ(xt,y,t)=(1+w)εθ(xt,y,t)−wεθ(xt,∅,t)(11)
其中 w=0 是普通条件预测,w>0 把预测沿“无条件 → 条件”的方向外推。若把
γ=1+w,也可写成
εuncond+γ(εcond−εuncond)。
这种线性外推与 Classifier Guidance 一样会以多样性换条件一致性。它并不保证对应某个归一化概率分布的精确 score,尤其当 scale 很大、网络预测又有误差时,外推结果可能离开训练时见过的 score 范围。空条件训练比例与 guidance scale 因此共同决定无条件基线是否可靠以及外推能走多远。