Skip to content
huc
Go back

把 GAN 看成参数空间中的扩散 ODE

目录 1 / 6

《生成扩散模型漫谈(十九):作为扩散ODE的GAN》 —— 苏剑林

《生成扩散模型漫谈(二十):从ReFlow到WGAN-GP》 —— 苏剑林

渐变发生在生成器训练时间里

普通扩散模型显式维护 ptp_t 并让样本沿许多时刻逐渐运动;GAN 的生成器 gθ(z)g_\theta(z) 看起来一次就把噪声映射成样本。联系两者的关键是把优化进度 τ\tau 当作另一种时间:

xτ=gθτ(z),zN(0,I).x_\tau=g_{\theta_\tau}(z), \qquad z\sim\mathcal N(0,I).

随着交替训练更新 θ0,θ1,\theta_0,\theta_1,\ldots,生成分布 qτq_\tau 也形成一条渐变路径。于是 GAN 可以理解为:先在样本空间确定 qτq_\tau 应该怎样向真实分布 pp 移动,再让生成器参数更新去拟合这一步样本运动。

这里必须区分两个时间变量:扩散或 Flow 的路径时间 tt 描述一次样本运输,τ\tau 描述生成器训练进度。后面“在 t=0t=0 前进一步”会变成一次 θτθτ+1\theta_\tau\to\theta_{\tau+1},二者不能混用。

MonoFlow:判别器估计密度比方向

令密度比

rτ(x)=p(x)qτ(x).r_\tau(x)=\frac{p(x)}{q_\tau(x)}.

最小化 qτq_\taupp 的 KL 散度所对应的一条 Wasserstein 梯度流满足

τqτ(x)=(qτ(x)logrτ(x)).\partial_\tau q_\tau(x) =-\nabla\cdot \left(q_\tau(x)\nabla\log r_\tau(x)\right).

由连续性方程,样本速度是

dxdτ=xlogrτ(x).(1)\frac{dx}{d\tau}=\nabla_x\log r_\tau(x). \tag{1}

真实密度与生成密度都未知,但能分别采样。Vanilla GAN 的判别目标

maxDExplogσ(D(x))+Exqτlog(1σ(D(x)))\max_D \mathbb E_{x\sim p}\log\sigma(D(x)) +\mathbb E_{x\sim q_\tau}\log(1-\sigma(D(x)))

在无限容量和精确最优化下有

D(x)=logp(x)qτ(x).(2)D^*(x)=\log\frac{p(x)}{q_\tau(x)}. \tag{2}

所以当前判别器给出了公式 (1) 的速度势。它只估计当前 qτq_\tau 下的密度比,不能一次得到未来所有时刻的 ODE;但可以先做一个小 Euler 步:

x+=x+ϵxD(x),x=gθτ(z).(3)x^+=x+\epsilon\nabla_xD(x), \qquad x=g_{\theta_\tau}(z). \tag{3}

把样本移动蒸馏回生成器参数

下一组生成器参数应让同一个潜变量 zz 产生公式 (3) 的目标:

Ltransport(θ)=Ez[gθ(z)gθτ(z)ϵxD(gθτ(z))2].(4)\mathcal L_{\text{transport}}(\theta) =\mathbb E_z \left[ \|g_\theta(z)-g_{\theta_\tau}(z) -\epsilon\nabla_xD(g_{\theta_\tau}(z))\|^2 \right]. \tag{4}

θ=θτ\theta=\theta_\tau 处求梯度,前两项抵消,通过链式法则得到

θLtransportθD(gθ(z))at θ=θτ.\nabla_\theta\mathcal L_{\text{transport}} \propto -\nabla_\theta D(g_\theta(z)) \quad\text{at }\theta=\theta_\tau.

因此若判别器固定且生成器只更新一小步,公式 (4) 与常用生成器损失

LG(θ)=Ez[D(gθ(z))](5)\mathcal L_G(\theta) =\mathbb E_z[-D(g_\theta(z))] \tag{5}

给出相同方向,差一个正比例系数。等价只在当前参数处的一阶梯度意义上成立,不表示两个 loss 在有限距离内具有相同函数值或相同最优解。这也解释了为什么固定判别器后不宜让生成器连续走很多步:离开 θτ\theta_\tau 后,公式 (5) 不再忠实拟合最初的局部运输目标;若要多步拟合,应保留带旧生成器锚点的公式 (4)

MonoFlow 还允许把 logr\log r 换成它的单调递增变换 h(logr)h(\log r)。这会改变速度大小,且通常也会改变参数化后的轨迹,但上升方向仍由密度比排序决定。很多 f-GAN 判别器最优解都是密度比的单调函数,因此可以放进同一梯度流视角。

这条推导与 GAN 的实际交替训练相吻合:当前生成器采样,判别器估计当前密度比,样本前进一步,再把这一步拟合回生成器。传统“对任意生成分布先把判别器完全解到最优,再把结果代入某个散度”的论证,更接近嵌套优化,并不直接描述有限步交替训练。

从 Rectified Flow 重建同一个交替过程

第二条路线不从 KL 的 Wasserstein 梯度流出发。令当前生成器样本为

x0=gθτ(z),x_0=g_{\theta_\tau}(z),

真实样本为 x1px_1\sim p,用直线连接:

xt=(1t)x0+tx1,dxtdt=x1x0.(6)x_t=(1-t)x_0+tx_1, \qquad \frac{dx_t}{dt}=x_1-x_0. \tag{6}

Rectified Flow 训练速度场

Lv(φ)=Ex0,x1,t[12vφ(xt,t)(x1x0)2].(7)\mathcal L_v(\varphi) =\mathbb E_{x_0,x_1,t} \left[ \frac12 \|v_\varphi(x_t,t)-(x_1-x_0)\|^2 \right]. \tag{7}

训练完成后,vφ(x,0)v_{\varphi^*}(x,0) 给出当前生成分布向真实分布运动的初始边缘速度。于是同样可以构造生成器局部运输目标:

Ez[gθ(z)gθτ(z)ϵvφ(gθτ(z),0)2].(8)\mathbb E_z \left[ \|g_\theta(z)-g_{\theta_\tau}(z) -\epsilon v_{\varphi^*}(g_{\theta_\tau}(z),0)\|^2 \right]. \tag{8}

更新生成器后重新训练速度场,便得到与 GAN 类似的交替过程。不同于 MonoFlow,这里速度通过真假样本之间的条件位移回归得到,不需要先证明密度比梯度流。

梯度速度场把 ReFlow 变成 WGAN-GP

展开公式 (7) 并删去与 φ\varphi 无关的 x1x02/2\|x_1-x_0\|^2/2

12vφ(xt,t)2vφ(xt,t),x1x0.\frac12\|v_\varphi(x_t,t)\|^2 -\langle v_\varphi(x_t,t),x_1-x_0\rangle.

接下来需要两个额外假设:忽略显式时间输入,并把向量场限制为某个标量 critic 的梯度,vφ(x)=xDφ(x)v_\varphi(x)=\nabla_xD_\varphi(x)。由公式 (6) 和链式法则,

Dφ(xt),x1x0=ddtDφ(xt).\left\langle \nabla D_\varphi(x_t),x_1-x_0 \right\rangle =\frac{d}{dt}D_\varphi(x_t).

如果对 tU[0,1]t\sim U[0,1] 求期望,使用微积分基本定理可以精确写成端点差:

01ddtDφ(xt)dt=Dφ(x1)Dφ(x0).\int_0^1\frac{d}{dt}D_\varphi(x_t)dt =D_\varphi(x_1)-D_\varphi(x_0).

于是速度回归目标化为

LD=E[Dφ(x0)Dφ(x1)+12Dφ(xt)2](9)\boxed{ \mathcal L_D =\mathbb E \left[ D_\varphi(x_0)-D_\varphi(x_1) +\frac12\|\nabla D_\varphi(x_t)\|^2 \right] } \tag{9}

这就是带零中心梯度惩罚的 WGAN 形式,而且惩罚位置 xtx_t 正好是真假样本线性插值。原始 WGAN-GP 常惩罚 (D1)2(\|\nabla D\|-1)^2,公式 (9) 则惩罚 D2\|\nabla D\|^2;二者相关但不是同一个正则项。

v=Dv=\nabla D 时,公式 (8) 在当前参数处的一步梯度又等价于公式 (5)。因此 WGAN 的 critic 损失与生成器损失都能从 ReFlow 的局部运输得到。

两条推导的共同闭环

MonoFlow 和 ReFlow 选择的速度来源不同:前者用 log(p/qτ)\nabla\log(p/q_\tau),后者回归人为端点耦合下的条件位移;只有加入特定参数化与理想化假设后,它们才落到熟悉的 GAN loss。共同结构则很稳定:

  1. 当前生成器定义可采样分布 qτq_\tau
  2. 判别器或速度网络估计 qτq_\taupp 的局部运动方向;
  3. 生成器用一步参数更新吸收这段样本运动;
  4. 分布改变后重新估计局部方向。

因此“GAN 是扩散 ODE”不是说推理时也要多步去噪,而是说它把原本发生在单个样本上的连续运输,压进了生成器训练轨迹。训练完成后仍由 gθ(z)g_\theta(z) 一步生成;渐进性已经储存在参数历史中。