flow_matching
约 4483 字大约 15 分钟
2026-06-19
在前面的模型,我们构造了流模型和扩散模型,其中都提到了神经网络向量场utθ。但我们都没有讲如何训练,以及优化θ来让生成模型产生有意义的东西。
Flow matching
在这节中,将只关注 flow models
X0∼pinit,dXt=utθ(Xt)dt(10)
问题变成了,我们如何优化θ来让X1尽可能的接近真实数据分布pdata。
条件和边缘概率路径
流匹配的第一步是找到一个概率路径,直观感觉概率路径指明了噪声点到真实数据的一个插值路径。
我们定义的ODE 轨迹,满足X0∼pinit,t=0以及X1∼pdata,t=1。 但是当0<t<1时会发生什么?,这中间的空档实际上可以自由发挥了,这期间的路径,在数学上可以被称为概率路径。
下面介绍一些新的内容,对于数据点 z∈Rd,我们用 δz 表示 狄拉克δ(Dirac delta)“分布”。这是可以想象的最简单的分布:从 δz 中采样总是返回 z(即它是确定性的)。条件(插值)概率路径(Conditional (interpolating) probability path)是一组在 Rd 上的分布 pt(x∣z),使得:
p0(⋅∣z)=pinit,p1(⋅∣z)=δzfor all z∈Rd.
换句话说,一个概率路径,将初始化分布,转变成了一个单个数据点(高维空间里的一个“点”)。你可以把概率路径理解成空间中的一条轨迹。
每一个条件概率路径 pt(x∣z) 一起加起来产生出一个 marginal probability path(边缘概率路径)pt(x),其定义为通过首先从数据分布中采样一个数据点 z∼pdata,然后从 pt(⋅∣z) 中采样所获得的分布,在概率论中,“边缘化(Marginalization)”的意思就是“消去某个变量”:
z∼pdata,x∼pt(⋅∣z)⇒x∼ptpt(x)=∫pt(x∣z)pdata(z)dz▹从边缘路径采样▹边缘路径的密度
上面公式1 的视角: 这给出了一个两步采样的过程:
- 先从你的真实数据库里随机挑一张图 z(比如挑中了“狗”)。
- 再在时间 t,从针对“狗”的条件路径 pt(⋅∣z) 中采样一个中间点 x。这样得到的 x,就属于宏观的边缘概率路径 pt。
公式 2 的视角: pt(x)(边缘概率路径的密度):
- 这是在时间 t 时,全空间总的概率分布。
- 积分的含义:它是把所有可能的真实数据 z(从猫到狗到汽车)所对应的条件路径 pt(x∣z),按照它们在现实中出现的概率 pdata(z) 进行加权平均(叠加)。
上面的积分是无法进行计算的。
边缘概率路径 pt 在 pinit 和 pdata 之间进行插值:
p0=pinit和p1=pdata▹噪声-数据插值
下面以高斯条件路径来举例:
一种特别流行的概率路径是高斯概率路径(Gaussian probability path)。设 αt,βt 为 noise schedulers(噪声调度器):两个连续可微的单调函数,且满足 α0=β1=0 以及 α1=β0=1。我们随后定义条件概率路径
pt(⋅∣z)=N(αtz,βt2Id)◃高斯条件路径(15)
由我们对 αt 和 βt 施加的条件可知,上式满足:
p0(⋅∣z)=N(α0z,β02Id)=N(0,Id),且p1(⋅∣z)=N(α1z,β12Id)=δz,
其中我们利用了这样一个事实:均值为 z、方差为零的正态分布就是 δz。因此,这个 pt(x∣z) 的选择在 pinit=N(0,Id) 的情况下满足公式 (11),是一个有效的条件插值路径。我们可以将从边缘路径 pt 采样的过程表示为:
z∼pdata, ϵ∼pinit=N(0,Id)⇒x=αtz+βtϵ∼pt◃从边缘高斯路径采样
直观上,上述过程在较低的 t 时添加更多噪声,直到时间 t=0,此时只有噪声。
条件和边缘概率向量场
前面介绍了概率路径指明了Xt∼pt的分布,在 t 时刻,我们希望找到一个向量场来描述这个轨迹。
对于每个数据点 z∈Rd,令 uttarget(⋅∣z) 表示一个 conditional vector field(条件向量场)。这可以是任何满足相应常微分方程(ODE)能产生条件概率路径 pt(⋅∣z) 的向量场,即满足以下条件:
X0∼pinit,dtdXt=uttarget(Xt∣z)⇒Xt∼pt(⋅∣z)(0≤t≤1).
乍一看,条件向量场似乎没什么用,因为 ODE 的所有终点 X1 都会坍缩到 X1=z,即我们只是在重新生成已知的数据点 z。然而,条件向量场是构建真正能从 pdata 生成样本的向量场的基础模块:
设 uttarget(x∣z) 为一个条件向量场。那么,定义为如下形式的 marginal vector field(边缘向量场) uttarget(x):
uttarget(x)=∫uttarget(x∣z)pt(x)pt(x∣z)pdata(z)dz,
遵循边缘概率路径,即:
X0∼pinit,dtdXt=uttarget(Xt)⇒Xt∼pt(0≤t≤1).(19)
特别地,对于这个常微分方程(ODE),X1∼pdata,因此我们可以说 "uttarget 将噪声 pinit 转换为了数据 pdata"。
如果使用的是之前提到的高斯概率路径的话,他的条件高斯向量场,应该是如下的形式:
uttarget(x∣z)=(α˙t−βtβ˙tαt)z+βtβ˙tx
证明
下面是证明,让我们首先通过定义
ψttarget(x∣z)=αtz+βtx.
来构建一个条件流模型 ψttarget(x∣z)。
如果 Xt 是 ψttarget(⋅∣z) 的常微分方程(ODE)轨迹,且初始状态 X0∼pinit=N(0,Id),那么根据定义有:
Xt=ψttarget(X0∣z)=αtz+βtX0∼N(αtz,β2Id)=pt(⋅∣z).
我们由此得出轨迹的分布符合该条件概率路径(即满足了上面的公式)。接下来需要从 ψttarget(x∣z) 中提取出向量场 uttarget(x∣z)。根据流的定义(见 flow models),成立:
dtdψttarget(x∣z)=uttarget(ψttarget(x∣z)∣z)对所有 x,z∈Rd⇔(i)α˙tz+β˙tx=uttarget(αtz+βtx∣z)对所有 x,z∈Rd⇔(ii)α˙tz+β˙t(βtx−αtz)=uttarget(x∣z)对所有 x,z∈Rd⇔(iii)(α˙t−βtβ˙tαt)z+βtβ˙tx=uttarget(x∣z)对所有 x,z∈Rd
其中在 (i) 中我们使用了 ψttarget(x∣z) 的定义,在 (ii) 中进行了将X进行归一化 x→(x−αtz)/βt,在 (iii) 中只是进行了一些代数运算。注意,最后一个等式正是我们在公式中定义的条件高斯向量场。
其中在边缘概率场中的积分公式,值得注意的是
pt(x)pt(x∣z)pdata(z)="给定含噪数据 x 时,数据点 z 的后验分布"
其中 pdata(z) 是先验分布。边缘向量场则只是一个平均值:对于每个可能的数据点 z,它取速度 ut(x∣z) —— 即能将我们带到 z 的方向 —— 然后根据我们认为 x 来自 z 的程度来对该速度进行加权。通过对所有数据点取平均,我们得到了边缘向量场。
为了更严谨的说明这个公式 我们将使用 continuity equation(连续性方程),这是数学和物理学中的一个基本方程。定义 divergence(散度) 算子 div 为:
div(vt)(x)=i=1∑d∂xi∂vti(x)
其中 vti 是 vt 的第 i 个坐标分量。
连续性方程
让我们考虑一个带有向量场 uttarget 的流模型,其中初始状态 X0∼pinit=p0。那么对于所有 0≤t≤1,Xt∼pt 成立的充要条件是:
∂tpt(x)=−div(ptuttarget)(x)对所有 x∈Rd,0≤t≤1,
其中 ∂tpt(x)=dtdpt(x) 表示 pt(x) 对时间的导数。该公式被称为 连续性方程(continuity equation)。
数学证明过于复杂,我们可以直观的来理解,概率密度不会减少也不会增多,只是从一个点转移到另一个点。某点概率密度的变化 = 流进来的概率 - 流出去的概率
左边的式子代表固定点 x 在 x 处概率密度随时间的变化率
右边的式子代表概率流的散度,负的“概率流发散” → 即净流入率
现在我们证明边缘向量场的积分公式,满足连续性方程 我们必须证明如 公式 (18) 中定义的边缘向量场 uttarget 满足连续性方程。我们可以通过直接计算来完成这一点:
∂tpt(x)=(i)∂t∫pt(x∣z)pdata(z)dz=∫∂tpt(x∣z)pdata(z)dz=(ii)∫−div(pt(⋅∣z)uttarget(⋅∣z))(x)pdata(z)dz=(iii)−div(∫pt(x∣z)uttarget(x∣z)pdata(z)dz)=(iv)−div(pt(x)∫uttarget(x∣z)pt(x)pt(x∣z)pdata(z)dz)(x)=(v)−div(ptuttarget)(x),
至于(ii)这步已经使用了连续性方程,这是证明的前提,条件概率路径满足连续性方程。
从边缘概率向量场中学习
我们把时间t 定义为 0-1 的均匀分布Unif[0,1],利用均分误差来定义 flow matching loss
LFM(θ)=Et∼Unif,x∼pt[∥utθ(x)−uttarget(x)∥2]=(i)Et∼Unif,z∼pdata,x∼pt(⋅∣z)[∥utθ(x)−uttarget(x)∥2](24)
其中 pt(x)=∫pt(x∣z)pdata(z)dz 是边缘概率路径。直观地说,这个损失意味着:首先,抽取一个随机时间 t∈[0,1]。其次,从我们的数据集中抽取一个随机点 z,从 pt(⋅∣z) 中采样(例如,通过添加一些噪声),并计算 utθ(x)。最后,计算我们的神经网络输出与边缘向量场 uttarget(x) 之间的均方误差。不幸的是,我们在这里还没有完成。虽然我们确实通过之前的公式知道了 uttarget 的公式,但我们无法高效地计算它,因为该积分是难以处理的。相反,我们将利用conditional(条件)速度场 uttarget(x∣z) 是易于处理的事实。为此,让我们定义 conditional flow matching loss(条件流匹配损失):
LCFM(θ)=Et∼Unif,z∼pdata,x∼pt(⋅∣z)[∥utθ(x)−uttarget(x∣z)∥2]
定理
边缘流匹配损失等于条件流匹配损失加上一个常数。即,
LFM(θ)=LCFM(θ)+C,
其中 C 是与 θ 无关的常数。因此,它们的梯度是一致的:
∇θLFM(θ)=∇θLCFM(θ).
因此,使用例如随机梯度下降(SGD)来最小化 LCFM(θ) 等价于以同样的方式最小化 LFM(θ)。特别是,对于 LCFM(θ) 的最小值点 θ∗,将成立 utθ∗=uttarget,即神经网络将等于边缘向量场(假设参数化具有无限表达能力)。
对于定理的证明,如下:
LFM(θ)=(i)Et∼Unif,x∼pt[∥utθ(x)−uttarget(x)∥2]=(ii)Et∼Unif,x∼pt[∥utθ(x)∥2−2utθ(x)Tuttarget(x)+∥uttarget(x)∥2]=(iii)Et∼Unif,x∼pt[∥utθ(x)∥2]−2Et∼Unif,x∼pt[utθ(x)Tuttarget(x)]+=:C1Et∼Unif[0,1],x∼pt[∥uttarget(x)∥2]=(iv)Et∼Unif,z∼pdata,x∼pt(⋅∣z)[∥utθ(x)∥2]−2Et∼Unif,x∼pt[utθ(x)Tuttarget(x)]+C1
其中第i 步使用了定义。第ii 步使用了完全平方公式。第iii 步定义了一个常量C,因为这一项是完全不含θ的。 第iv 步使用了之前的采样来重写第一项。下面我们来重新表达下第二项的内容。
Et∼Unif,x∼pt[utθ(x)Tuttarget(x)]=(i)∫01∫pt(x)utθ(x)Tuttarget(x)dxdt=(ii)∫01∫pt(x)utθ(x)T[∫uttarget(x∣z)pt(x)pt(x∣z)pdata(z)dz]dxdt=(iii)∫01∫∫utθ(x)Tuttarget(x∣z)pt(x∣z)pdata(z)dzdxdt=(iv)Et∼Unif,z∼pdata,x∼pt(⋅∣z)[utθ(x)Tuttarget(x∣z)]
第一步将表达式展开为积分。第二步使用了之前的公式,将边缘向量场表示为条件向量场的积分形式。第三步利用了积分的线性规则,重新排列的积分顺序。 第四步将积分重新写成期望的形式。这是非常重要的一步证明,开始的时候是边缘向量场,结束的时候是条件向量场。我们把这个加到流匹配损失中去。
LFM(θ)=(i)Et∼Unif,z∼pdata,x∼pt(⋅∣z)[∥utθ(x)∥2−2Et∼Unif,z∼pdata,x∼pt(⋅∣z)[utθ(x)Tuttarget(x∣z)]+C1=(ii)Et∼Unif,z∼pdata,x∼pt(⋅∣z)[∥utθ(x)∥2−2utθ(x)Tuttarget(x∣z)+∥uttarget(x∣z)∥2−∥uttarget(x∣z)∥2]+C1=(iii)Et∼Unif,z∼pdata,x∼pt(⋅∣z)[∥utθ(x)−uttarget(x∣z)∥2]+C2Et∼Unif,z∼pdata,x∼pt(⋅∣z)[−∥uttarget(x∣z)∥2]+C1=(iv)LCFM(θ)+=:CC2+C1
第一步展开,是把第二项写成前面推导的形式。第二步是加一个,减一个。 第三步是利用了,完全平方公式的逆向使用。第四步,利用定义, 前面第一项就是条件流匹配损失,后面两项不含 θ 相当于常数了。
因此,流匹配训练归结为最小化条件流匹配损失关于该算法,有以下几个显著特点:
第一,我们在训练过程中实际上从不模拟任何常微分方程(ODE)。人们将算法的这一特性称为无模拟(simulation-free)。这使得训练成本极低,因为你无需在训练过程中展开ODE的轨迹(这需要很多步迭代)。
第二,训练目标是一个简单的回归目标——我们只是对目标向量场 (uttarget(x∣z))进行回归。因此,它本质上与监督学习没有太大区别。
最后,该算法极其简单——很难想象有比这更简单的训练目标了。
所有这些特点使得流匹配成为大规模机器学习模型中极具吸引力的方法。一旦(utθ) 训练完成,我们就可以通过例如算法1的方式模拟流模型
dXt=utθ(Xt)dt,X0∼pinit(27)
从而获得样本(X1∼pdata)。
总结
流匹配训练旨在学习边际向量场(marginal vector field)uttarget。为了构建它,我们选择满足条件 p0(⋅∣z)=pinit 和 p1(⋅∣z)=δz 的条件概率路径(conditional probability path)pt(x∣z)。接下来,我们寻找一个条件向量场(conditional vector field)uttarget(x∣z),使其对应的流(flow)ψttarget(x∣z) 满足:
X0∼pinit⇒Xt=ψttarget(X0∣z)∼pt(⋅∣z),
或者等价地,满足 uttarget 符合连续性方程(continuity equation)。然后,由以下公式定义的边际向量场(marginal vector field):
uttarget(x)=∫uttarget(x∣z)pt(x)pt(x∣z)pdata(z)dz
遵循边际概率路径,即:
X0∼pinit,dXt=uttarget(Xt)dt⇒Xt∼pt(0≤t≤1).
特别是对于该常微分方程(ODE),有 X1∼pdata,因此正如预期的那样,uttarget “将噪声转化为数据(converts noise into data)”。为了学习它,我们需要最小化条件流匹配损失(conditional flow matching loss):
LCFM(θ)=Et∼Unif,z∼pdata,x∼pt(⋅∣z)[∥utθ(x)−uttarget(x∣z)∥2].