Flow Matching 推导总结
整理自ChatGPT
1. 连续性方程(Continuity Equation)
设样本遵循 ODE:
dtdxt=vt(xt)
概率密度演化满足:
∂t∂pt(x)+∇x⋅(pt(x)vt(x))=0
2. 构造分布路径(Probability Path)
对任意样本对 x0∼p0, x1∼p1,定义:
xt=(1−t)x0+tx1
3. 条件向量场(Conditional Vector Field)
由上式可得:
ut(xt∣x0,x1)=x1−x0
目标:学习 vθ(xt,t)≈ut
4. Flow Matching 损失
L(θ)=Et,x0,x1∥vθ(xt,t)−(x1−x0)∥22
5. 采样(生成)
求解 ODE:
dtdxt=vθ(xt,t),t∈[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=0∑K−1tr[Z⊤∂Xvθ(ti,Xti)Z]Δti(1)
该公式用于流模型的 log-likelihood 估计,并能够在高维生成任务中无须显式构造 Jacobian,从而高效实现数值求解。
1. 生成流与连续性方程
考虑样本随向量场 vθ 生成的概率流(probability flow)ODE:
dtdxt=vθ(t,xt),t∈[0,1],(2)
其概率密度 pt(x) 满足连续性方程(Continuity / Liouville equation):
∂t∂pt(x)+∇x⋅(pt(x)vθ(t,x))=0.(3)
展开散度项(公式4中前者是梯度后者是散度):
∇x⋅(ptvθ)=vθ⊤∇xpt+pt∇x⋅vθ.(4)
将其代回 (3):
∂t∂pt+vθ⊤∇xpt+pt∇x⋅vθ=0.(5)
2. Log-Density 沿流轨迹的演化
考虑密度沿真实轨迹 xt 的全导数(pt(xt)=p(t,x);x=x(t)):
dtdlnpt(xt)=pt(xt)1(∂t∂pt+vθ⊤∇xpt).(6)
使用 (5) 替换分子:
∂t∂pt+vθ⊤∇xpt=−pt∇x⋅vθ.
代回 (6) 得:
dtdlnpt(xt)=−∇x⋅vθ(t,xt)=−tr(∂xvθ(t,xt))(7)
对 t∈[0,1] 积分:
lnp1(x1)=lnp0(x0)−∫01tr(∂xvθ(t,xt))dt(8)
3. 时间离散化(Riemann 求和)
建立时间网格:
0=t0<t1<⋯<tK=1,Δti=ti+1−ti,
则积分 (8) 的 Riemann 近似:
lnp1(x1)≈lnp0(x0)−i=0∑K−1tr(∂xvθ(ti,xti))Δti(9)
4. Hutchinson 迹估计
为了高效估计 trace(Jacobian),需要一种只用 向量乘法就能估计 trace 的方法。
对任意方阵 A∈Rd×d:
tr(A)=Ez∼N(0,I)[z⊤Az],(10)
对于满足 E[zz⊤]=I 的随机向量,成立:
E[z⊤Az]=tr(A)
E[z⊤Az]=E[tr(z⊤Az)]=E[tr(Azz⊤)]=tr(AE[zz⊤])=tr(AI)=tr(A).
使用 m 个探针向量 Z=[z1,…,zm]∈Rd×m:
tr(A)≈tr(Z⊤AZ)=j=1∑mzj⊤Azj.(11)
将其代入离散 log-density 公式 (9) 得:
lnp^1(x1)=lnp0(x0)−i=0∑K−1tr[Z⊤∂Xvθ(ti,Xti)Z]Δti(12)
5. 实现注记(Jacobian-Free 计算)
设 y=f(x),Jacobian Jf(x)=∂x∂y。
| 名称 |
符号 |
数学形式 |
| JVP(Jacobian–Vector Product) |
Jf(x)v |
dϵdf(x+ϵv) |
| VJP(Vector–Jacobian Product) |
u⊤Jf(x) |
dϵd(u⊤f(x+ϵI)) |
Hutchinson 迹估计:
tr(Jf(x))≈z⊤Jf(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] return (grad * z).sum()
|
tr(Z⊤∂xvZ) 不需要显式计算 Jacobian,可用 JVP/VJP:
Az=∂xvθ(t,x)z
可由自动微分库(PyTorch/JAX)高效实现。
6. 总结
连续理论结论:
dtdlnpt(xt)=−tr(∂xvθ)
离散可计算形式:
lnp^1(x1)=lnp0(x0)−i=0∑K−1tr[Z⊤∂Xvθ(ti,Xti)Z]Δti
特点:
| 优势 |
原因 |
| 无需显式 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