《生成扩散模型漫谈(十九):作为扩散ODE的GAN》 —— 苏剑林
《生成扩散模型漫谈(二十):从ReFlow到WGAN-GP》 —— 苏剑林
渐变发生在生成器训练时间里
普通扩散模型显式维护 pt 并让样本沿许多时刻逐渐运动;GAN 的生成器 gθ(z) 看起来一次就把噪声映射成样本。联系两者的关键是把优化进度 τ 当作另一种时间:
xτ=gθτ(z),z∼N(0,I).
随着交替训练更新 θ0,θ1,…,生成分布 qτ 也形成一条渐变路径。于是 GAN 可以理解为:先在样本空间确定 qτ 应该怎样向真实分布 p 移动,再让生成器参数更新去拟合这一步样本运动。
这里必须区分两个时间变量:扩散或 Flow 的路径时间 t 描述一次样本运输,τ 描述生成器训练进度。后面“在 t=0 前进一步”会变成一次 θτ→θτ+1,二者不能混用。
MonoFlow:判别器估计密度比方向
令密度比
rτ(x)=qτ(x)p(x).
最小化 qτ 到 p 的 KL 散度所对应的一条 Wasserstein 梯度流满足
∂τqτ(x)=−∇⋅(qτ(x)∇logrτ(x)).
由连续性方程,样本速度是
dτdx=∇xlogrτ(x).(1)
真实密度与生成密度都未知,但能分别采样。Vanilla GAN 的判别目标
DmaxEx∼plogσ(D(x))+Ex∼qτlog(1−σ(D(x)))
在无限容量和精确最优化下有
D∗(x)=logqτ(x)p(x).(2)
所以当前判别器给出了公式 (1) 的速度势。它只估计当前 qτ 下的密度比,不能一次得到未来所有时刻的 ODE;但可以先做一个小 Euler 步:
x+=x+ϵ∇xD(x),x=gθτ(z).(3)
把样本移动蒸馏回生成器参数
下一组生成器参数应让同一个潜变量 z 产生公式 (3) 的目标:
Ltransport(θ)=Ez[∥gθ(z)−gθτ(z)−ϵ∇xD(gθτ(z))∥2].(4)
在 θ=θτ 处求梯度,前两项抵消,通过链式法则得到
∇θLtransport∝−∇θD(gθ(z))at θ=θτ.
因此若判别器固定且生成器只更新一小步,公式 (4) 与常用生成器损失
LG(θ)=Ez[−D(gθ(z))](5)
给出相同方向,差一个正比例系数。等价只在当前参数处的一阶梯度意义上成立,不表示两个 loss 在有限距离内具有相同函数值或相同最优解。这也解释了为什么固定判别器后不宜让生成器连续走很多步:离开 θτ 后,公式 (5) 不再忠实拟合最初的局部运输目标;若要多步拟合,应保留带旧生成器锚点的公式 (4)。
MonoFlow 还允许把 logr 换成它的单调递增变换 h(logr)。这会改变速度大小,且通常也会改变参数化后的轨迹,但上升方向仍由密度比排序决定。很多 f-GAN 判别器最优解都是密度比的单调函数,因此可以放进同一梯度流视角。
这条推导与 GAN 的实际交替训练相吻合:当前生成器采样,判别器估计当前密度比,样本前进一步,再把这一步拟合回生成器。传统“对任意生成分布先把判别器完全解到最优,再把结果代入某个散度”的论证,更接近嵌套优化,并不直接描述有限步交替训练。
从 Rectified Flow 重建同一个交替过程
第二条路线不从 KL 的 Wasserstein 梯度流出发。令当前生成器样本为
x0=gθτ(z),
真实样本为 x1∼p,用直线连接:
xt=(1−t)x0+tx1,dtdxt=x1−x0.(6)
Rectified Flow 训练速度场
Lv(φ)=Ex0,x1,t[21∥vφ(xt,t)−(x1−x0)∥2].(7)
训练完成后,vφ∗(x,0) 给出当前生成分布向真实分布运动的初始边缘速度。于是同样可以构造生成器局部运输目标:
Ez[∥gθ(z)−gθτ(z)−ϵvφ∗(gθτ(z),0)∥2].(8)
更新生成器后重新训练速度场,便得到与 GAN 类似的交替过程。不同于 MonoFlow,这里速度通过真假样本之间的条件位移回归得到,不需要先证明密度比梯度流。
梯度速度场把 ReFlow 变成 WGAN-GP
展开公式 (7) 并删去与 φ 无关的 ∥x1−x0∥2/2:
21∥vφ(xt,t)∥2−⟨vφ(xt,t),x1−x0⟩.
接下来需要两个额外假设:忽略显式时间输入,并把向量场限制为某个标量 critic 的梯度,vφ(x)=∇xDφ(x)。由公式 (6) 和链式法则,
⟨∇Dφ(xt),x1−x0⟩=dtdDφ(xt).
如果对 t∼U[0,1] 求期望,使用微积分基本定理可以精确写成端点差:
∫01dtdDφ(xt)dt=Dφ(x1)−Dφ(x0).
于是速度回归目标化为
LD=E[Dφ(x0)−Dφ(x1)+21∥∇Dφ(xt)∥2](9)
这就是带零中心梯度惩罚的 WGAN 形式,而且惩罚位置 xt 正好是真假样本线性插值。原始 WGAN-GP 常惩罚 (∥∇D∥−1)2,公式 (9) 则惩罚 ∥∇D∥2;二者相关但不是同一个正则项。
当 v=∇D 时,公式 (8) 在当前参数处的一步梯度又等价于公式 (5)。因此 WGAN 的 critic 损失与生成器损失都能从 ReFlow 的局部运输得到。
两条推导的共同闭环
MonoFlow 和 ReFlow 选择的速度来源不同:前者用 ∇log(p/qτ),后者回归人为端点耦合下的条件位移;只有加入特定参数化与理想化假设后,它们才落到熟悉的 GAN loss。共同结构则很稳定:
- 当前生成器定义可采样分布 qτ;
- 判别器或速度网络估计 qτ 向 p 的局部运动方向;
- 生成器用一步参数更新吸收这段样本运动;
- 分布改变后重新估计局部方向。
因此“GAN 是扩散 ODE”不是说推理时也要多步去噪,而是说它把原本发生在单个样本上的连续运输,压进了生成器训练轨迹。训练完成后仍由 gθ(z) 一步生成;渐进性已经储存在参数历史中。