《生成扩散模型漫谈(二十一):中值定理加速ODE采样》 —— 苏剑林
从欧拉法的误差开始
扩散 ODE 把噪声端 xT 送到数据端 x0:
dtdxt=vθ(xt,t).(1)
在反向网格 0=t0<t1<⋯<tN=T 上,欧拉法只用右端点速度近似整段速度:
xtn≈xtn+1−(tn+1−tn)vθ(xtn+1,tn+1).
它每步只需一次网络调用,但局部误差是一阶的。Heun 法先用欧拉法得到预测点 x~tn,再平均区间两端的速度:
x~tnxtn=xtn+1−Δtnvθ(xtn+1,tn+1),≈xtn+1−2Δtn[vθ(xtn+1,tn+1)+vθ(x~tn,tn)],
其中 Δtn=tn+1−tn。高阶方法减少了截断误差,却增加每步 NFE;总 NFE 极小时,被迫使用的大步长仍会让通用高阶 Solver 失效。
精确目标其实是区间平均速度
对公式 (1) 在 [tn,tn+1] 上积分:
xtn+1−xtn=∫tntn+1vθ(xt,t)dt.(2)
因此一步真正需要的不是端点瞬时速度,而是
vˉn=Δtn1∫tntn+1vθ(xt,t)dt.(3)
标量连续函数满足积分中值定理,存在 sn∈(tn,tn+1) 使 vˉn=vθ(xsn,sn)。但速度是高维向量,通常不存在一个共同的 sn 让所有坐标同时满足中值等式。AMED 使用的是近似假设:短区间内轨迹足够接近直线,使平均速度可由某个中间点的速度近似。
把未知中间点变成可学习量
如果已知 sn,仍然不知道轨迹上的 xsn。AMED 先用右端点速度做一次欧拉预测:
x~sn=xtn+1−(tn+1−sn)vθ(xtn+1,tn+1).
再用小网络 gϕ 根据 U-Net 的中间特征 htn+1 与时间预测 sn,最终一步为
xtn≈xtn+1−Δtnvθ(x~sn,sn),sn=gϕ(htn+1,tn+1).(4)
训练 gϕ 时,先用高 NFE Solver 产生较精确的端点对,再最小化公式 (4) 的一步误差。教师提供的是局部轨迹端点,不是生成样本数据集;被训练的也只是很小的时间点预测器,所以蒸馏成本远低于把整个扩散模型蒸馏成生成器。
为什么高维中值近似仍可能有效
向量积分中值等式若严格成立,轨迹方向必须非常受限;一般曲线没有理由满足它。AMED 的经验依据是:高精度采样轨迹做 PCA 后,前一两个主成分已解释绝大多数变化,说明轨迹常落在近二维子空间,并且接近其中的一条直线。
这也与 ReFlow 一类方法的训练路径有关。训练时常用噪声与数据的线性插值作为伪轨迹,学出的速度场会受到“走直线”的偏置。于是 AMED 并非把标量定理直接推广到向量,而是利用扩散轨迹近低维、近直线这一额外结构。
NFE 与适用边界
AMED 一步需要端点和中间点两次主网络计算,因此 NFE 近似为 2;小网络 gϕ 的成本通常忽略。比较 Solver 时必须固定总 NFE,而不是固定步数。极低 NFE 下,二阶通用方法可能因为步长过大反而不如一阶 DDIM,AMED 的优势来自对扩散轨迹的定制,而不只是形式上的阶数。
第一步还可利用高噪声端速度近似可解析的 AFS 技巧省掉一次 NFE。它依赖模型参数化在噪声端的特性,并非对所有扩散 ODE 都自动成立。