Skip to content
huc
Go back

Shortcut Model:把步长作为生成条件

目录 1 / 4

《生成扩散模型漫谈(二十七):将步长作为条件输入》 —— 苏剑林

瞬时速度不能直接承担大步长

令噪声 x0p0x_0\sim p_0、数据 x1p1x_1\sim p_1,选直线路径

xt=(1t)x0+tx1,t[0,1].(1)x_t=(1-t)x_0+t x_1, \qquad t\in[0,1]. \tag{1}

Flow Matching 用平方损失学习瞬时速度:

LFM=E12vθ(xt,t,0)(x1x0)2.(2)\mathcal L_{\text{FM}} =\mathbb E\frac12 \left\|v_\theta(x_t,t,0)-(x_1-x_0)\right\|^2. \tag{2}

最后一个输入 00 表示零步长极限。推理只能用有限差分

xt+d=xt+dvθ(xt,t,0),x_{t+d}=x_t+d\,v_\theta(x_t,t,0),

它是欧拉近似,dd 越大误差越大。一步生成要求 d=1d=1,恰好处于瞬时速度最不可靠的区域。

网络输出改为区间平均速度

Shortcut Model 让网络同时接收步长 dd

xt+d=xt+dvθ(xt,t,d).(3)x_{t+d}=x_t+d\,v_\theta(x_t,t,d). \tag{3}

此时 vθ(xt,t,d)v_\theta(x_t,t,d) 的目标不再是 tt 点的瞬时速度,而是从 ttt+dt+d 的有效平均速度。d=0d=0 的锚点仍由公式 (2) 监督;有限步长则从组合一致性得到监督。

xtx_t 连走两次长度 dd 的步:

x~t+d=xt+dvθ(xt,t,d),x~t+2d=x~t+d+dvθ(x~t+d,t+d,d).\begin{aligned} \tilde x_{t+d}&=x_t+d\,v_\theta(x_t,t,d),\\ \tilde x_{t+2d}&=\tilde x_{t+d} +d\,v_\theta(\tilde x_{t+d},t+d,d). \end{aligned}

它应等于一次长度 2d2d 的更新。因此平均速度必须满足

vθ(xt,t,2d)=12[vθ(xt,t,d)+vθ(x~t+d,t+d,d)].(4)v_\theta(x_t,t,2d) =\frac12\left[ v_\theta(x_t,t,d) +v_\theta(\tilde x_{t+d},t+d,d) \right]. \tag{4}

用自举目标覆盖所有步长

将公式 (4) 变成损失:

LSC=Evθ(xt,t,2d)sg ⁣(vθ(xt,t,d)+vθ(x~t+d,t+d,d)2)2.(5)\mathcal L_{\text{SC}} =\mathbb E \left\| v_\theta(x_t,t,2d) -\operatorname{sg}\!\left( \frac{v_\theta(x_t,t,d)+v_\theta(\tilde x_{t+d},t+d,d)}{2} \right) \right\|^2. \tag{5}

stop-gradient 把两个小步分支作为自举教师,避免一次更新同时移动目标。训练在二进制步长层级上从小到大传播:Flow Matching 锚定 d=0d=0,较小 dd 的模型监督 2d2d,最终覆盖 d=1d=1。若小步模型本身有偏差,误差也会逐层累积,所以锚点损失与各尺度采样比例都不能省略。

一套模型支持一步与多步

训练后可直接用

x1=x0+vθ(x0,0,1)x_1=x_0+v_\theta(x_0,0,1)

一步生成;也可选 d=1/Nd=1/N 迭代 NN 次。与普通 Solver 不同,改变步数时网络会明确知道当前 dd,而不是拿同一个瞬时速度场硬套不同离散误差。

步长输入与“区间起点、终点”输入本质等价:给定 ttdd 就得到区间 [t,t+d][t,t+d]。该方法不需要预训练教师和对抗训练,但公式 (5) 依赖自举及 stop-gradient,理论保证不如直接监督瞬时 Flow 完整。