《生成扩散模型漫谈(二十七):将步长作为条件输入》 —— 苏剑林
瞬时速度不能直接承担大步长
令噪声 x0∼p0、数据 x1∼p1,选直线路径
xt=(1−t)x0+tx1,t∈[0,1].(1)
Flow Matching 用平方损失学习瞬时速度:
LFM=E21∥vθ(xt,t,0)−(x1−x0)∥2.(2)
最后一个输入 0 表示零步长极限。推理只能用有限差分
xt+d=xt+dvθ(xt,t,0),
它是欧拉近似,d 越大误差越大。一步生成要求 d=1,恰好处于瞬时速度最不可靠的区域。
网络输出改为区间平均速度
Shortcut Model 让网络同时接收步长 d:
xt+d=xt+dvθ(xt,t,d).(3)
此时 vθ(xt,t,d) 的目标不再是 t 点的瞬时速度,而是从 t 到 t+d 的有效平均速度。d=0 的锚点仍由公式 (2) 监督;有限步长则从组合一致性得到监督。
从 xt 连走两次长度 d 的步:
x~t+dx~t+2d=xt+dvθ(xt,t,d),=x~t+d+dvθ(x~t+d,t+d,d).
它应等于一次长度 2d 的更新。因此平均速度必须满足
vθ(xt,t,2d)=21[vθ(xt,t,d)+vθ(x~t+d,t+d,d)].(4)
用自举目标覆盖所有步长
将公式 (4) 变成损失:
LSC=Evθ(xt,t,2d)−sg(2vθ(xt,t,d)+vθ(x~t+d,t+d,d))2.(5)
stop-gradient 把两个小步分支作为自举教师,避免一次更新同时移动目标。训练在二进制步长层级上从小到大传播:Flow Matching 锚定 d=0,较小 d 的模型监督 2d,最终覆盖 d=1。若小步模型本身有偏差,误差也会逐层累积,所以锚点损失与各尺度采样比例都不能省略。
一套模型支持一步与多步
训练后可直接用
x1=x0+vθ(x0,0,1)
一步生成;也可选 d=1/N 迭代 N 次。与普通 Solver 不同,改变步数时网络会明确知道当前 d,而不是拿同一个瞬时速度场硬套不同离散误差。
步长输入与“区间起点、终点”输入本质等价:给定 t 与 d 就得到区间 [t,t+d]。该方法不需要预训练教师和对抗训练,但公式 (5) 依赖自举及 stop-gradient,理论保证不如直接监督瞬时 Flow 完整。