概率路径

概率路径与条件概率路径

流匹配 Flow Matching 的起点是 概率路径 Probability Paths。概率路径是随时间变化的分布序列 (pt)0t1{(p_t)}_{0 \le t \le 1},用于把简单先验分布连接到真实数据分布。

实际训练中无法直接观测 pdatap_{data},也无法直接构造 pinitp_{init}pdatap_{data} 的全局路径。

可行做法是对每个样本 zpdataz \sim p_{data} 构造从噪声到该样本的 条件概率路径 Conditional Probability Path pt(xz)p_t(x|z),并满足:

  1. t=0t=0 时为初始噪声分布,通常设为 p0(z)=pinit=N(0,Id)p_0(\cdot|z) = p_{init} = \mathcal{N}(0, I_d)
  2. t=1t=1 时塌缩到数据点 zz,即 p1(z)=δzp_1(\cdot|z) = \delta_z,其中 δz\delta_z 是以 zz 为中心的 Dirac delta 分布。

这样即可定义从噪声到样本的条件演化路径。

边缘概率路径

给定条件路径 pt(xz)p_t(x |z) 后,可通过边缘化得到 边缘概率路径 Marginal Probability Path pt(x)p_t(x)

pt(x)=pt(xz)pdata(z)dzp_t(x) = \int p_t(x|z) p_{data}(z) dz

边缘分布是所有条件分布的加权平均,权重由 pdata(z)p_{data}(z) 给出。学习大量条件路径即可间接学习整体分布的演化。

由于条件路径满足边界条件,边缘路径也满足 p0=pinitp_0 = p_{init}p1=pdatap_1 = p_{data}。因此,只要找到驱动粒子沿 ptp_t 演化的机制,就能从噪声生成数据。

高斯条件概率路径

常用的中间路径是高斯条件路径:

pt(xz)=N(x;αtz,βt2Id)p_t(x|z) = \mathcal{N}(x; \alpha_t z, \beta_t^2 I_d)

其中 αt\alpha_tβt\beta_t 是噪声调度函数,满足 α0=0,β0=1\alpha_0 = 0, \beta_0 = 1α1=1,β1=0\alpha_1 = 1, \beta_1 = 0。该路径通过线性变换和噪声缩放,使分布从标准高斯逐渐收敛到 zz。常见调度如下:

路径类型 αt\alpha_t βt\beta_t 几何特征
最优传输路径 tt 1t1-t 严格直线演化,速度恒定
方差保持路径 1σt2\sqrt{1-\sigma_t^2} σt\sigma_t 曲线演化,沿球面移动
方差爆炸路径 1 σ(t)\sigma(t) 均值不变,仅方差扩张

调整 αt,βt\alpha_t,\beta_t 可得到不同几何路径。最优传输路径对应线性插值,轨迹直接、速度形式简单。

从条件路径采样可写为:

zpdataϵpinit=N(0,Id)    x=αtz+βtϵptz \sim p_{data} \quad \epsilon \sim p_{init} = \mathcal{N}(0, I_d) \implies x=\alpha_tz+\beta_t\epsilon \sim p_t

向量场

条件向量场

概率路径给出分布如何变化;向量场给出粒子在该路径上的速度。

对每个数据点 zz,定义 条件向量场 Conditional Vector Field

X0pinitddtXt=uttarget(Xtz)    Xtpt(z)(0t1)X_0 \sim p_{init} \quad \frac{d}{dt} X_t = u_t^{target}(X_t|z) \implies X_t \sim p_t(\cdot | z) \quad (0 \le t \le 1)

其中 uttarget(Xtz)u_t^{target}(X_t|z) 是目标条件向量场。

X0pinitX_0 \sim p_{init} 是噪声起点。给定初始点后,ODE 定义的流 ψt(x0)\psi_t(x_0) 确定粒子位置。向量场、ODE 和流分别对应速度规则、局部变化和全局轨迹。

连续性方程:向量场改变概率分布的基础是连续性方程,它描述概率质量在向量场下的守恒演化:

tpt(x)+div(ptut)(x)=0\frac{\partial}{\partial t}p_t(x) + \text{div}(p_t u_t)(x) = 0

其中 div(ptut)\text{div}(p_t u_t) 是概率流散度,tpt(x)\partial_t p_t(x) 是密度变化率。若向量场满足该方程,则概率质量只发生连续搬运,不凭空产生或消失。

对高斯条件路径 pt(xz)=N(x;αtz,βt2Id)p_t(x|z) = \mathcal{N}(x; \alpha_t z, \beta_t^2 I_d),可用流图 ψt(x0z)=αtz+βtx0\psi_t(x_0|z) = \alpha_t z + \beta_t x_0 求条件向量场:

ddtψt(x0z)=α˙tz+β˙tx0=uttarget(ψt(x0z)z)\frac{d}{dt} \psi_t(x_0|z) = \dot{\alpha}_t z + \dot{\beta}_t x_0 = u_t^{target}(\psi_t(x_0|z) | z)

x0x_0 替换为 (xαtz)/βt(x - \alpha_t z) / \beta_t,得到闭式解:

uttarget(xz)=(α˙tβ˙tβtαt)z+β˙tβtxu_t^{target}(x|z) = (\dot{\alpha}_t - \frac{\dot{\beta}_t}{\beta_t} \alpha_t) z + \frac{\dot{\beta}_t}{\beta_t} x

该速度由数据点 zz 相关项和尺度变化项组成。对 CondOT 路径(αt=t,βt=1t\alpha_t=t, \beta_t=1-t),可简化为 ut(xz)=zϵu_t(x|z) = z - \epsilon,其中 ϵ\epsilon 是初始噪声。

边缘向量场

条件向量场可由条件路径直接构造。对所有可能的 zz 做后验加权平均,即得到边缘向量场:

uttarget(x)=uttarget(xz)pt(xz)pdata(z)pt(x)dzu_t^{target}(x) = \int u_t^{target}(x|z) \frac{p_t(x|z) p_{data}(z)}{p_t(x)} dz

权重 pt(xz)pdata(z)pt(x)\frac{p_t(x|z) p_{data}(z)}{p_t(x)} 是后验概率 pt(zx)p_t(z|x)。因此,位置 xx 的边缘速度等于所有条件速度的后验期望:

uttarget(x)=Ezpdata(x)[uttarget(xz)]u_t^{target}(x) = \mathbb{E}_{z \sim p_{data}(\cdot|x)} [u_t^{target}(x|z)]

直观上,uttarget(xz)u_t^{target}(x|z) 表示“若终点是 zz,当前位置应如何移动”,pt(zx)p_t(z|x) 表示当前位置最终对应各数据点的概率。二者加权积分得到当前位置的平均运动方向。

流匹配损失

定义概率路径和向量场后,目标是学习边缘向量场,即总体分布的演化速度。

理想的 流匹配损失 FM Loss 要求神经网络向量场 vθ(x,t)v_\theta(x, t) 逼近真实边缘向量场 uttarget(x)u_t^{target}(x)

LFM(θ)=EtUnif,xpt(x)[vθ(x,t)uttarget(x)2]\mathcal{L}_{FM}(\theta) = \mathbb{E}_{t \sim \text{Unif}, x \sim p_t(x)} [\|v_\theta(x, t) - u_t^{target}(x)\|^2]

uttarget(x)u_t^{target}(x) 依赖未知的 pdatap_{data},不能直接计算。因此使用可计算的 条件流匹配损失 CFM Loss

LCFM(θ)=Et,zpdata,xpt(xz)[vθ(x,t)uttarget(xz)2]\mathcal{L}_{CFM}(\theta) = \mathbb{E}_{t, z \sim p_{data}, x \sim p_t(x|z)} [\|v_\theta(x, t) - u_t^{target}(x|z)\|^2]

其中 uttarget(xz)u_t^{target}(x|z)pt(xz)p_t(x|z) 可直接计算。可以证明 FM 与 CFM 对参数 θ\theta 的梯度等价:

θLFM(θ)=θLCFM(θ)\nabla_\theta \mathcal{L}_{FM}(\theta) = \nabla_\theta \mathcal{L}_{CFM}(\theta)

因此,训练时匹配条件速度即可得到正确的边缘向量场。

FM/CFM 等价性证明

  1. 展开 FM 损失:

LFM(θ)=Et,x\mathcal{L}_{FM}(\theta) = \mathbb{E}_{t, x}

其中第三项 uttarget2\|u_t^{target}\|^2 不含 θ\theta,优化时可视为常数。

  1. 处理交叉项:利用边缘化公式 uttarget(x)=uttarget(xz)pt(xz)pdata(z)pt(x)dzu_t^{target}(x) = \int u_t^{target}(x|z) \frac{p_t(x|z) p_{data}(z)}{p_t(x)} dz

Expt=pt(x)vθ(x,t)T(uttarget(xz)pt(xz)pdata(z)pt(x)dz)dx\mathbb{E}_{x \sim p_t} = \int p_t(x) v_\theta(x, t)^T \left( \int u_t^{target}(x|z) \frac{p_t(x|z) p_{data}(z)}{p_t(x)} dz \right) dx

=vθ(x,t)Tuttarget(xz)pt(xz)pdata(z)dxdz\dots = \int \int v_\theta(x, t)^T u_t^{target}(x|z) p_t(x|z) p_{data}(z) dx dz

这等价于 Ezpdata,xpt(xz)\mathbb{E}_{z \sim p_{data}, x \sim p_t(x|z)}

  1. 合并回归目标:同理可证明第一项 vθ2\|v_\theta\|^2 的期望一致。因此两个损失关于 θ\theta 的一阶导数一致。

确定损失后,训练流程如下: