Flow Matching 推导总结

整理自ChatGPT

1. 连续性方程(Continuity Equation)

设样本遵循 ODE:

dxtdt=vt(xt)\frac{dx_t}{dt} = v_t(x_t)

概率密度演化满足:

pt(x)t+x(pt(x)vt(x))=0\frac{\partial p_t(x)}{\partial t} + \nabla_x \cdot \big(p_t(x) v_t(x)\big) = 0

2. 构造分布路径(Probability Path)

对任意样本对 x0p0x_0 \sim p_0, x1p1x_1 \sim p_1,定义:

xt=(1t)x0+tx1x_t = (1-t)x_0 + t x_1

3. 条件向量场(Conditional Vector Field)

由上式可得:

ut(xtx0,x1)=x1x0u_t(x_t \mid x_0, x_1) = x_1 - x_0

目标:学习 vθ(xt,t)utv_\theta(x_t, t) \approx u_t

4. Flow Matching 损失

L(θ)=Et,x0,x1vθ(xt,t)(x1x0)22\mathcal{L}(\theta) = \mathbb{E}_{t, x_0, x_1} \| v_\theta(x_t, t) - (x_1 - x_0) \|_2^2

5. 采样(生成)

求解 ODE:

dxtdt=vθ(xt,t),t[0,1]\frac{dx_t}{dt} = v_\theta(x_t, t), \quad t \in [0,1]

6. 优势总结

  • 训练是监督回归,不是score matching或ELBO
  • 无需SDE,生成是确定ODE,收敛稳定
  • 与扩散模型兼容,可视为其deterministic版本

Log-Density Estimation in Flow Matching Models

本章节推导 Flow Matching / Rectified Flow 模型中终端对数密度(log-density)的离散近似公式

lnp^1(x1)=lnp0(x0)i=0K1tr ⁣[ZXvθ(ti,Xti)Z]Δti(1)\boxed{ \ln \hat p_1(x_1) = \ln p_0(x_0) - \sum_{i=0}^{K-1} \operatorname{tr}\!\big[ Z^\top \partial_X v_\theta(t_i, X_{t_i}) Z \big]\, \Delta t_i } \tag{1}

该公式用于流模型的 log-likelihood 估计,并能够在高维生成任务中无须显式构造 Jacobian,从而高效实现数值求解。

1. 生成流与连续性方程

考虑样本随向量场 vθv_\theta 生成的概率流(probability flow)ODE:

dxtdt=vθ(t,xt),t[0,1],(2)\frac{d x_t}{dt} = v_\theta(t, x_t), \qquad t \in [0,1], \tag{2}

其概率密度 pt(x)p_t(x) 满足连续性方程(Continuity / Liouville equation):

pt(x)t+x(pt(x)vθ(t,x))=0.(3)\frac{\partial p_t(x)}{\partial t} + \nabla_x \cdot \big( p_t(x) v_\theta(t,x) \big) = 0 . \tag{3}

展开散度项(公式4中前者是梯度后者是散度):

x(ptvθ)=vθxpt+ptxvθ.(4)\nabla_x \cdot (p_t v_\theta) = v_\theta^\top \nabla_x p_t + p_t\, \nabla_x \cdot v_\theta . \tag{4}

将其代回 (3):

ptt+vθxpt+ptxvθ=0.(5)\frac{\partial p_t}{\partial t} + v_\theta^\top \nabla_x p_t + p_t\, \nabla_x \cdot v_\theta = 0 . \tag{5}

2. Log-Density 沿流轨迹的演化

考虑密度沿真实轨迹 xtx_t 的全导数(pt(xt)=p(t,x);x=x(t)p_t(x_t)=p(t,x);x=x(t)):

ddtlnpt(xt)=1pt(xt)(ptt+vθxpt).(6)\frac{d}{dt} \ln p_t(x_t) = \frac{1}{p_t(x_t)} \Big( \frac{\partial p_t}{\partial t} + v_\theta^\top \nabla_x p_t \Big). \tag{6}

使用 (5) 替换分子:

ptt+vθxpt=ptxvθ.\frac{\partial p_t}{\partial t} + v_\theta^\top \nabla_x p_t = - p_t \, \nabla_x \cdot v_\theta.

代回 (6) 得:

ddtlnpt(xt)=xvθ(t,xt)=tr(xvθ(t,xt))(7)\boxed{ \frac{d}{dt}\ln p_t(x_t) = - \nabla_x \cdot v_\theta(t,x_t) = - \operatorname{tr}(\partial_x v_\theta(t,x_t)) } \tag{7}

t[0,1]t\in[0,1] 积分:

lnp1(x1)=lnp0(x0)01tr(xvθ(t,xt))dt(8)\boxed{ \ln p_1(x_1) = \ln p_0(x_0) - \int_0^1 \operatorname{tr}(\partial_x v_\theta(t, x_t))\, dt } \tag{8}

3. 时间离散化(Riemann 求和)

建立时间网格:

0=t0<t1<<tK=1,Δti=ti+1ti,0 = t_0 < t_1 < \cdots < t_K = 1, \qquad \Delta t_i = t_{i+1} - t_i,

则积分 (8) 的 Riemann 近似:

lnp1(x1)lnp0(x0)i=0K1tr(xvθ(ti,xti))Δti(9)\boxed{ \ln p_1(x_1) \approx \ln p_0(x_0) - \sum_{i=0}^{K-1} \operatorname{tr}(\partial_x v_\theta(t_i, x_{t_i}))\, \Delta t_i } \tag{9}

4. Hutchinson 迹估计

为了高效估计 trace(Jacobian),需要一种只用 向量乘法就能估计 trace 的方法。

对任意方阵 ARd×dA \in \mathbb{R}^{d\times d}

tr(A)=EzN(0,I)[zAz],(10)\operatorname{tr}(A) = \mathbb{E}_{z \sim \mathcal{N}(0,I)}[z^\top A z], \tag{10}

对于满足 E[zz]=I\mathbb{E}[z z^\top] = I 的随机向量,成立:

E[zAz]=tr(A)\mathbb{E}[z^\top A z] = \operatorname{tr}(A)

E[zAz]=E[tr(zAz)]=E[tr(Azz)]=tr(AE[zz])=tr(AI)=tr(A).\begin{aligned} \mathbb{E}[z^\top A z] &= \mathbb{E}[\operatorname{tr}(z^\top A z)] \\ &= \mathbb{E}[\operatorname{tr}(A z z^\top)] \\ &= \operatorname{tr}(A\, \mathbb{E}[z z^\top]) \\ &= \operatorname{tr}(A I) \\ &= \operatorname{tr}(A). \end{aligned}

使用 mm 个探针向量 Z=[z1,,zm]Rd×mZ = [z_1,\dots,z_m] \in \mathbb{R}^{d\times m}

tr(A)tr(ZAZ)=j=1mzjAzj.(11)\operatorname{tr}(A) \approx \operatorname{tr}(Z^\top A Z) = \sum_{j=1}^m z_j^\top A z_j. \tag{11}

将其代入离散 log-density 公式 (9) 得:

lnp^1(x1)=lnp0(x0)i=0K1tr ⁣[ZXvθ(ti,Xti)Z]Δti(12)\boxed{ \ln \hat p_1(x_1) = \ln p_0(x_0) - \sum_{i=0}^{K-1} \operatorname{tr}\!\big[ Z^\top \partial_X v_\theta(t_i, X_{t_i}) Z \big]\, \Delta t_i } \tag{12}

5. 实现注记(Jacobian-Free 计算)

y=f(x)y = f(x),Jacobian Jf(x)=yxJ_f(x) = \frac{\partial y}{\partial x}

名称 符号 数学形式
JVP(Jacobian–Vector Product) Jf(x)vJ_f(x) v ddϵf(x+ϵv)\frac{d}{d\epsilon} f(x+\epsilon v)
VJP(Vector–Jacobian Product) uJf(x)u^\top J_f(x) ddϵ(uf(x+ϵI))\frac{d}{d\epsilon} (u^\top f(x+\epsilon I))

Hutchinson 迹估计:

tr(Jf(x))zJf(x)z\operatorname{tr}(J_f(x)) \approx z^\top J_f(x) z

PyTorch:

1
2
3
4
5
def divergence(f, x):
z = torch.randn_like(x)
fz = (f(x) * z).sum()
grad = torch.autograd.grad(fz, x, create_graph=True)[0] # VJP
return (grad * z).sum() # zᵀ (J z)

tr(ZxvZ)\operatorname{tr}(Z^\top \partial_x v Z) 不需要显式计算 Jacobian,可用 JVP/VJP:

Az=xvθ(t,x)zA z = \partial_x v_\theta(t,x)\,z

可由自动微分库(PyTorch/JAX)高效实现。

6. 总结

连续理论结论:

ddtlnpt(xt)=tr(xvθ)\frac{d}{dt} \ln p_t(x_t) = -\operatorname{tr}(\partial_x v_\theta)

离散可计算形式:

lnp^1(x1)=lnp0(x0)i=0K1tr ⁣[ZXvθ(ti,Xti)Z]Δti\ln \hat p_1(x_1) = \ln p_0(x_0) - \sum_{i=0}^{K-1} \operatorname{tr}\!\big[ Z^\top \partial_X v_\theta(t_i, X_{t_i}) Z \big]\, \Delta t_i

特点:

优势 原因
无需显式 Jacobian 使用 Hutchinson 估计
ODE 确定性推断 比 SDE 扩散模型更稳定
适用于 Flow Matching / Rectified Flow 与生成 ODE 完全一致
可用于 log-likelihood / bits-per-dim 支持可评估生成模型

参考文献

  • Chen et al., 2022 — Flow Matching for Generative Modeling
  • Lipman et al., 2023 — Flow Matching in Probability Space
  • Liu et al., 2024 — Rectified Flow