深入理解 Flow Matching:生成建模的统一理论框架

5235 字
26 分钟
深入理解 Flow Matching:生成建模的统一理论框架

1. 为什么要理解 Flow Matching?#

1.1 扩散模型的”理论黑箱”问题#

在阅读 Diffusion Transformer 文章MMDiT 文章 时,你可能已经注意到一个细节:两篇文章都提到了 Rectified Flow,但它们都没有从最底层解释——

为什么”把噪声和样本用直线连接”的轨迹,就能训练出一个生成模型?

答案藏在 Flow Matching(流匹配)理论里——这是 Yaron Lipman 等人在 2022 年提出的万有理论框架,它把 DDPM、Score Matching、Rectified Flow、Nerfies、Dynablock 等看似不同的生成方法统一在一个优雅的数学框架下

1.2 一句话概括#

Flow Matching 的核心思想:用一条精心设计的 连续路径(Flow) 把噪声分布连接到数据分布,然后训练一个神经网络去学”这条路径的速度场”——学成之后,沿 ODE 反向走,就能从噪声生成样本。

1.3 直观理解#

想象一条河流系统:

起点 (噪声 x_1) ───────→ ←─── ←─── ←─── ←─── ←─── ←─── ←─── ←───
←─── ←─── ←─── ←─── ←─── ←─── ←─── ←───
←─── ←─── ←─── ←─── ←─── ←─── ←─── ←───
终点 (数据 x_0)
每条路径代表一个样本从噪声到数据的"旅程"
Flow Matching = 学出整条河流的速度场
→ 知道了速度场,就可以从任意点出发
→ 逆流而上就能从噪声到达数据

1.4 Flow Matching 的历史脉络#

2015-2019: Neural ODE (Chen et al., 2018) — "连续化神经网络"
│ 把离散残差网络 → 微分方程
2019-2021: Score Matching → Noise Conditional Score Network (NCSN)
│ 学 ∇log p(x) — 朗之万采样
2020: DDPM (Ho et al.) — 加噪 → 去噪, 但理论不统一
2022: Flow Matching (Lipman et al., Jun 2022)
│ ★ 统一了: DDPM / NCSN / 任意路径
2022: Conditional Flow Matching (CFM) — 条件版本
2023: Rectified Flow (Liu et al.) — 最优传输路径
│ ★ 被 SD3 / FLUX / MMDiT 采用
2024+: Flow Matching 成为扩散模型训练的理论基础

2. 概率论基础:从噪声到数据的路径#

2.1 目标是什么?#

我们有一堆真实数据点 {x0(i)}\{x_0^{(i)}\},它们服从某个未知分布 pdata(x)p_{\text{data}}(x)

我们想学一个可逆变换 f:zxf: z \mapsto x,使得:

  • zN(0,I)z \sim \mathcal{N}(0, I)(噪声)
  • x=f(z)pdata(x)x = f(z) \sim p_{\text{data}}(x)(生成数据)

可逆性:如果知道如何从 x0x_0 走到 x1x_1,那么我就能从 x1x_1 反推回 x0x_0——这就是”生成”的本质。

2.2 连续时间路径#

Flow Matching 把从 x0x_0(数据)到 x1x_1(噪声)的过程参数化为时间连续的路径 xt:[0,1]Rdx_t: [0,1] \rightarrow \mathbb{R}^d

x0=数据x1=噪声x_0 = \text{数据} \quad x_1 = \text{噪声}

对于任意时刻 t[0,1]t \in [0,1],路径 xtx_t 给出该时刻的”中间状态”。

# 路径的直观例子
import numpy as np
t = np.linspace(0, 1, 100) # 100 个时间点
# 1) DDPM 路径 (概率守恒路径)
x_t_ddpm = alpha_bar[t]**0.5 * x0 + (1 - alpha_bar[t])**0.5 * noise
# 2) Rectified Flow 路径 (直线)
x_t_rf = (1 - t) * x0 + t * noise
# 3) 子空间线性路径
x_t_linear = (1 - t) * x0 + t * noise

2.3 三种经典路径对比#

路径类型公式形状代表模型
DDPM (VP-SDE)xt=αˉtx0+1αˉtϵx_t = \sqrt{\bar\alpha_t}x_0 + \sqrt{1-\bar\alpha_t}\epsilon弧线(高斯混合)DDPM, SD 1/2/xl
Rectified Flowxt=(1t)x0+tϵx_t = (1-t)x_0 + t\epsilon直线SD3, FLUX, MMDiT
Sub-VP (次扩散)xt=1e2tsomethingx_t = \sqrt{1 - e^{-2t}} \cdot \text{something}弧线(更快衰减)改进型 SD
x_t
│ DDPM: 弧线
│ ╭─╮
│ ╱ ╲
│ ╱ RF ╲ ← 直线,最短路径
│ ╱ ╲ ╲
│╱ ────╲
└────────────── t
0 1

2.4 边际分布#

对于固定的 tt,所有样本的均值和方差决定边际分布 pt(x)p_t(x)

pt(x)=pdata(x0)p(xtx0)dx0p_t(x) = \int p_{\text{data}}(x_0) \cdot p(x_t | x_0) \, dx_0

好的路径设计让 ptp_ttt 平滑变化——不能有任何 tt 出现奇点。

3. 常微分方程 (ODE) 与概率流#

3.1 核心:从路径到向量场#

Flow Matching 的关键洞察是:每条路径 xtx_t 都是某个向量场 vt(x)v_t(x) 的积分曲线

dxtdt=vt(xt),x0=xstart,x1=xend\frac{dx_t}{dt} = v_t(x_t), \quad x_0 = x_{\text{start}}, \quad x_1 = x_{\text{end}}

反过来,给定向量场 vtv_t,我就可以解 ODE 得到整条路径。

把生成建模问题重新表述为

  1. 找一个向量场 vt(x)v_t(x) 使得 ODE dx/dt=vt(x)dx/dt = v_t(x) 把噪声分布 N(0,I)\mathcal{N}(0,I) 推到数据分布 pdatap_{\text{data}}
  2. 学这个 vtv_t
  3. 采样时,从 zN(0,I)z \sim \mathcal{N}(0,I) 开始,反向积分 ODE(tt 从 1 降到 0)

3.2 向量场的定义#

对于条件路径 xt(x0)x_t(x_0)(以数据点 x0x_0 为起点),定义条件向量场

vt(xt(x0))=ddtxt(x0)v_t(x_t(x_0)) = \frac{d}{dt} x_t(x_0)

整体的边际向量场是条件向量场的加权平均:

vt(x)=Ex0pdata,  xtx0[vt(xt(x0))]v_t(x) = \mathbb{E}_{x_0 \sim p_{\text{data}}, \; x_t|x_0} \left[ v_t(x_t(x_0)) \right]
def conditional_velocity(x0, noise, t):
"""Rectified Flow 的条件向量场 = ε - x0(速度场)。"""
return noise - x0
def marginal_velocity(x, dataset):
"""边际向量场 = 所有条件向量场的期望。"""
v = 0.0
for x0 in dataset:
noise = sample_noise()
t = sample_time_uniform()
x_t = (1 - t) * x0 + t * noise # RF 路径
v += conditional_velocity(x0, noise, t)
return v / len(dataset)

3.3 概率流方程(连续性方程)#

一个深刻的事实:边际向量场 vtv_t 必须满足连续性方程——保证概率质量守恒:

ptt=(vtpt)\frac{\partial p_t}{\partial t} = -\nabla \cdot (v_t \cdot p_t)

这意味着:

  • vtv_t 的散度 vt\nabla \cdot v_t 控制概率密度如何变化
  • 如果 vtv_t 无散(vt=0\nabla \cdot v_t = 0),则 ptp_t 不随 tt 变化

为什么这很重要:因为 Flow Matching 证明,可以完全绕过 ptp_t 直接学习 vtv_t,而 vtv_t 学好后,ptp_t 自动满足连续性方程。

3.4 ODE 求解(采样)#

给定学好的向量场 vθ(x,t)v_\theta(x, t),生成样本只需要反向求解 ODE

@torch.no_grad()
def ode_sample(model, num_steps=50, device="cuda"):
"""
x_{t-1} = x_t - (1/num_steps) * v_θ(x_t, t)
t 从 1 → 0 (反向)
"""
x = torch.randn(batch_size, dim, device=device)
t_grid = torch.linspace(1.0, 0.0, num_steps + 1, device=device)
for i in range(num_steps):
t = t_grid[i]
t_batch = t.expand(x.shape[0])
v = model(x, t_batch) # 学好的向量场
dt = t_grid[i] - t_grid[i + 1]
# Euler 积分
x = x - dt * v
return x # 近似从 p_1 逆流到 p_0

这里的 tt反向时间101 \to 0):t=1t=1 是纯噪声,t=0t=0 是数据。

4. Flow Matching 定理#

4.1 条件流匹配目标 (Conditional Flow Matching, CFM)#

这是 Flow Matching 理论的核心。设 p(xtx0)p(x_t|x_0) 是以 x0x_0 为起点的条件分布,vt(xtx0)v_t(x_t|x_0) 是对应的条件向量场。

定理(Lipman et al., 2022)

训练目标 LCFM\mathcal{L}_{\text{CFM}} 等价于 边际流匹配目标 LFM\mathcal{L}_{\text{FM}} ,且梯度相同。

LCFM(θ)=Et,x0,x1vθ(xt,t)vt(xtx0)2\mathcal{L}_{\text{CFM}}(\theta) = \mathbb{E}_{t, x_0, x_1} \left\| v_\theta(x_t, t) - v_t(x_t | x_0) \right\|^2

其中 xtp(xtx0)x_t \sim p(x_t|x_0)

换句话说:你只需要让模型去拟合”以真实起点 x0x_0 为条件的向量场”,自动保证边际目标也是最优的。

4.2 证明的直觉#

边际目标: ∫ || v_θ(x) - 𝔼[v_t(x|x_0)] ||² p_t(x) dx
这是 x 的函数,难以优化
CFM 目标: 𝔼 || v_θ(x_t) - v_t(x_t|x_0) ||²
给定了 x_0,就知道 v_t(x_t|x_0) = dx_t/dt
因为 x_t = f(x_0, t) 是已知的!
通过边际分解: p_t(x) = ∫ p(x_t|x_0) p_data(x_0)
可以证明两个目标在最优解处等价 (差一个常数)

关键洞察:条件路径 xt(x0)x_t(x_0)已知的(因为我们设计它),所以 vt(xtx0)=ddtxt(x0)v_t(x_t|x_0) = \frac{d}{dt}x_t(x_0)可以解析计算的

def cfm_loss(model, x0, noise, t):
"""
条件流匹配损失 = MSE(预测速度, 真实速度)
Rectified Flow: v_t = ε - x0
DDPM VP-SDE: v_t = -σ² ∇_x log p_t(x|...)
"""
# 1) 构造条件路径
xt = (1 - t) * x0 + t * noise # RF 直线
# 2) 解析计算条件向量场 (速度)
true_v = noise - x0 # RF: d/dt[(1-t)x0 + tε] = ε - x0
# 3) 模型预测
pred_v = model(xt, t)
# 4) MSE
return ((pred_v - true_v) ** 2).mean()

4.3 为什么这个定理如此重要?#

传统方法问题CFM 的优势
Score Matching需要计算 Hessian 或 log-likelihood不需要
变分推断需要 ELBO、KL散度不需要
原始 DDPM目标函数复杂只需 MSE(预测, 真实速度)

CFM 目标 = 简单的 MSE + 解析已知的目标向量场——这让它成为最简单的扩散类训练目标。

4.4 对 DDPM 的重新解释#

用 Flow Matching 框架,DDPM 的训练目标可以重新推导

DDPM 的路径:xt=αˉtx0+1αˉtϵx_t = \sqrt{\bar\alpha_t} x_0 + \sqrt{1-\bar\alpha_t} \epsilon

tt 求导:

dxtdt=dαˉtdtx0+d1αˉtdtϵ\frac{dx_t}{dt} = \frac{d\sqrt{\bar\alpha_t}}{dt} x_0 + \frac{d\sqrt{1-\bar\alpha_t}}{dt} \epsilon

整理后得到速度场形式

vt(xtx0)=12σt2xtlogp(xtx0)v_t(x_t | x_0) = -\frac{1}{2} \sigma_t^2 \nabla_{x_t} \log p(x_t | x_0)

这揭示了 DDPM 与 Score Matching 的内在联系——两者学的其实是同一个东西!

def ddpm_velocity_from_eps(xt, alpha_bar, noise):
"""DDPM 速度场 = α' x0 + σ' ε
但用噪声预测网络 ε_θ 来表示。
"""
sqrt_alpha = np.sqrt(alpha_bar)
sqrt_1_alpha = np.sqrt(1 - alpha_bar)
alpha_prime = (d/dt) sqrt_alpha_bar # 与 DDPM 超参有关
sigma_prime = (d/dt) sqrt(1 - alpha_bar)
# 重新组织
x0_pred = (xt - sqrt_1_alpha * noise) / sqrt_alpha
v = alpha_prime * x0_pred + sigma_prime * noise
return v

5. 最优传输与 Rectified Flow#

5.1 什么是”最优传输路径”?#

在所有可能的路径中,最优传输(Optimal Transport, OT)路径有一个极其重要的性质:

沿最优传输路径训练的模型,其 ODE 采样 更直、更短——这意味着可以用更少的步数生成高质量样本。

定义:最优传输路径是最小化总体路径弯曲能量的那条:

minpt01Ex0pdata,ϵNdxtdt2dt\min_{p_t} \int_0^1 \mathbb{E}_{x_0\sim p_{\text{data}}, \epsilon\sim\mathcal{N}} \left\| \frac{dx_t}{dt} \right\|^2 dt

结论:对于高斯分布之间的传输,这个最小能量路径就是直线

5.2 Rectified Flow 的最优传输证明#

# Rectified Flow 是 Gaussian-to-Gaussian 的最优传输路径
def rf_optimal_transport_path(x0, epsilon):
"""
x_t = (1 - t) * x0 + t * epsilon
d/dt x_t = epsilon - x0 = - (x0 - epsilon)
能量: ||dx_t/dt||² = ||x0 - epsilon||²
这与 t 无关,是常数!
→ 所以 Rectified Flow 是平坦路径(零曲率)
→ 走这条路径做 ODE 采样时,误差最小
"""
return epsilon - x0

直观理解:DDPM 的路径是”弧线”,每一步都在”减速”(因为高斯方差在变化)。Rectified Flow 的路径是”直线”,模型只需学”匀速前进”——但从 x0x_0ϵ\epsilon 的匀速运动,就是恒定的速度场,最容易学习。

5.3 最速下降流 (Gradient Flow)#

另一种重要路径是梯度流——沿数据分布的对数密度梯度方向:

dxtdt=logpdata(xt)\frac{dx_t}{dt} = \nabla \log p_{\text{data}}(x_t)

这是 Stein Variational Gradient Descent 的连续时间版本。

def gradient_flow_loss(model, x0, noise, t):
"""
梯度流: 学 log p_data 的梯度
但 log p_data 未知!
所以在实际中,梯度流通常用
JKO 离散化或 MCMC 近似。
"""
pass

5.4 路径设计空间#

Flow Matching 路径设计空间:
所有路径 = {x_t = φ(x_0, ε, t) | φ 是任意光滑插值函数}
约束:
- x_0 = data
- x_1 = noise
- p_t 随 t 平滑变化
最优路径?
- 能量最小 → 直线 (Rectified Flow)
- 熵最小 → DDPM VP-SDE
- 混合 → sub-VP, VE-SDE, ...

6. 扩展:随机微分方程 (SDE) 视角#

6.1 为什么需要 SDE?#

ODE 路径是确定性的——给定起点,路径完全确定。但在 DDPM 中,加噪过程是随机的

dx=f(x,t)dt+g(t)dwdx = f(x, t) dt + g(t) dw

这引入了随机性。Flow Matching 可以同时处理 ODE(确定性)和 SDE(随机性) 两种情况。

6.2 VP-SDE(Variance Preserving)和 VE-SDE#

VP-SDE(方差保持)——DDPM 的连续化:

dx=12β(t)xdt+β(t)dwdx = -\frac{1}{2} \beta(t) x \, dt + \sqrt{\beta(t)} \, dw

VE-SDE(方差爆炸)——高噪声方差的情况:

dx=ddt[σ2(t)]dwdx = \sqrt{\frac{d}{dt}[\sigma^2(t)]} \, dw

6.3 ODE vs SDE 采样#

# ODE 采样 (Flow Matching 标准)
@torch.no_grad()
def ode_sample(model, xt, num_steps=50):
"""确定性反向 ODE 积分。"""
dt = 1.0 / num_steps
for i in reversed(range(num_steps)):
t = torch.full((xt.shape[0],), i / num_steps, device=xt.device)
v = model(xt, t)
xt = xt - dt * v # 反向积分
return xt
# SDE 采样 (DDPM Langevin 类型)
@torch.no_grad()
def sde_sample(model, xt, num_steps=50, beta_fn=None):
"""随机微分方程采样(DDPM 风格)。"""
dt = 1.0 / num_steps
for i in reversed(range(num_steps)):
t = torch.full((xt.shape[0],), i / num_steps, device=xt.device)
# 预测噪声
eps = model(xt, t)
# 漂移项 + 扩散项
drift = -beta_fn(t) / 2 * xt
diffusion = torch.sqrt(beta_fn(t)) * torch.randn_like(xt)
xt = xt - drift * dt + diffusion * np.sqrt(dt)
return xt

6.4 从 SDE 到 ODE 的技巧#

概率流对应:任何 SDE 都有一个对应的 ODE(去掉随机项后的期望轨迹)——这就是 Flow Matching 所用的。

SDE: dx=f(x,t)dt+g(t)dww=0ODE: dx=f(x,t)dt\text{SDE: } dx = f(x,t)dt + g(t)dw \quad \xRightarrow{w=0} \quad \text{ODE: } dx = f(x,t)dt

反过来:可以先训练 SDE 版本(更稳定的训练),然后用对应的 ODE 采样(更快,更确定性)。

7. 训练与采样完整流程#

7.1 完整训练代码#

import torch
import torch.nn as nn
from torch.utils.data import DataLoader
def flow_matching_train_step(model, batch, sigma_min=1e-3, sigma_max=50.0):
"""
通用 Flow Matching 训练步。
支持任意路径 (通过 path_fn 指定)。
"""
x0 = batch # (B, D) 真实数据
# 1. 采样时间步 t ∈ [0, 1]
t = torch.rand(x0.shape[0], device=x0.device)
# 2. 采样噪声
eps = torch.randn_like(x0)
# 3. 构造路径 (可替换为任意 φ)
# 默认: Rectified Flow (直线)
xt = (1 - t.view(-1, 1)) * x0 + t.view(-1, 1) * eps
# 4. 解析速度场 (由 φ 的定义决定)
# RF: v = ε - x0
true_v = eps - x0
# 5. 模型预测
pred_v = model(xt, t)
# 6. MSE 损失
loss = ((pred_v - true_v) ** 2).mean()
return loss
def train_flow_matching(model, dataloader, epochs=100, lr=1e-4):
optimizer = torch.optim.AdamW(model.parameters(), lr=lr)
for epoch in range(epochs):
for batch in dataloader:
batch = batch.to(device)
loss = flow_matching_train_step(model, batch)
optimizer.zero_grad()
loss.backward()
optimizer.step()
print(f"Epoch {epoch}, Loss: {loss.item():.4f}")
# 使用不同路径的示例
class FlowMatchingPath:
"""可配置的 Flow Matching 路径。"""
@staticmethod
def rectified_flow(x0, eps, t):
"""直线: x_t = (1-t)x_0 + t ε"""
return (1 - t) * x0 + t * eps, eps - x0
@staticmethod
def ve_sde(x0, eps, t, sigma_min=0.01, sigma_max=50.0):
"""方差爆炸 SDE 的路径。"""
sigma_t = sigma_min * (sigma_max / sigma_min) ** t
xt = x0 + sigma_t * eps
# 速度场: dσ/dt * ε
dsigma = sigma_min * (sigma_max / sigma_min) ** t * torch.log(
sigma_max / sigma_min
)
true_v = dsigma.view(-1, 1) * eps
return xt, true_v
@staticmethod
def goeginger(x0, eps, t):
"""GoGGinger 路径: 更稳定的插值。"""
lambda_t = torch.sinh(3 * t) / torch.sinh(3)
xt = (1 + lambda_t) / 2 * x0 + (1 - lambda_t) / 2 * eps
true_v = 3 / 2 * (torch.cosh(3 * t) / torch.sinh(3)) * (x0 - eps)
return xt, true_v

7.2 推理加速技术#

7.2.1 高阶 ODE 求解器#

Euler 积分简单但慢。可以用更高阶的求解器:

def heun_sampler(model, num_steps=20):
"""
Heun 二阶方法,比 Euler 精度高。
适合 RF 的直线 ODE——因为局部截断误差 O(dt²)。
"""
x = torch.randn(batch_size, dim, device=device)
dt = 1.0 / num_steps
for i in reversed(range(num_steps)):
t = i / num_steps
t_batch = torch.full((x.shape[0],), t, device=device)
# Euler 预测
v1 = model(x, t_batch)
x_mid = x - dt * v1
# Heun 校正 (用中点斜率)
t_mid = (i - 0.5) / num_steps
t_mid_batch = torch.full((x.shape[0],), t_mid, device=device)
v2 = model(x_mid, t_mid_batch)
# 加权组合
x = x - dt * ((v1 + v2) / 2)
return x

7.2.2 自适应步长#

from scipy.integrate import solve_ivp
def adaptive_sampler(model, rtol=1e-5, atol=1e-6):
"""用 scipy 的自适应 ODE 求解器。"""
def ode_fn(t, x_flat):
x = torch.tensor(x_flat.reshape(1, -1), device=device)
t_batch = torch.full((1,), 1.0 - t, device=device) # 反向时间
v = model(x, t_batch)
return -v.cpu().numpy().flatten() # dt/dt = -1 (反向积分)
# 从噪声开始
x0 = torch.randn(1, dim).cpu().numpy().flatten()
sol = solve_ivp(
ode_fn, t_span=(0, 1),
y0=x0,
method="RK45",
rtol=rtol, atol=atol,
)
return torch.tensor(sol.y[:, -1].reshape(1, -1), device=device)

7.3 一致性模型蒸馏#

Flow Matching 的直线 ODE 让一致性蒸馏变得极为简单:

一致性模型 (Consistency Model) 的核心观察: 在直线路径上,同一条直线上的所有点应该映射到同一个终点x0x_0)。

def consistency_distillation_step(model, student, x0, num_steps=6):
"""
从教师模型蒸馏到少步学生模型。
适用于 RF 直线路径。
"""
eps = torch.randn_like(x0)
t = torch.rand(x0.shape[0])
# 噪声图
xt = (1 - t) * x0 + t * eps
# 教师输出 (teacher 是多步模型)
with torch.no_grad():
v_teacher = model(xt, t)
x0_teacher = xt - t.view(-1, 1) * v_teacher
# 学生输出 (student 是少步模型,直接预测 x0)
x0_student = student(xt, t)
# 蒸馏损失: 学生应该预测和教师一样的 x0
loss = ((x0_student - x0_teacher) ** 2).mean()
return loss

8. Flow Matching 与 Score Matching 的关系#

8.1 核心区别#

维度Flow MatchingScore Matching
学什么向量场 vtv_t对数密度梯度 logpt\nabla \log p_t
关系vt=dxtdtv_t = \frac{dx_t}{dt}vt=σt2logptv_t = -\sigma_t^2 \nabla \log p_t
采样ODE 反向积分朗之万动力学 / ODE
训练目标MSE(简单)需要 Hutchinson 估计(复杂)
方差稳定取决于噪声调度

8.2 数学上的精确关系#

vt(xtx0)=ddtlogαtxt+ddtlogσtϵddtlogσtσtlogp(xtx0)v_t(x_t | x_0) = \frac{d}{dt} \log \alpha_t \cdot x_t + \frac{d}{dt} \log \sigma_t \cdot \epsilon - \frac{d}{dt} \log \sigma_t \cdot \sigma_t \nabla \log p(x_t | x_0)

在 Rectified Flow 中,αt=1t,σt=t\alpha_t = 1-t, \sigma_t = t,上式化简为:

vt(xtx0)=ϵx0v_t(x_t | x_0) = \epsilon - x_0

logp(xtx0)=ϵx0t\nabla \log p(x_t | x_0) = -\frac{\epsilon - x_0}{t},与 vtv_t 成正比。

9. Flow Matching 的应用版图#

9.1 图像生成#

Flow Matching 在图像生成中的位置:
Rectified Flow → SD3 / FLUX / MMDiT (核心训练目标)
- 论文: "Scaling Rectified Flow Transformers"
- 4 步采样 (FLUX Schnell)
- 比 DDPM 快 10 倍以上

9.2 音频生成#

# AudioLDM 2 / MusicLDM 使用 Flow Matching
# 声音 = 频谱图 (STFT) → 2D Flow Matching
# 音乐 = 多通道频谱图 + 时序依赖

9.3 3D 生成#

Point-E (Nichol et al., 2022):
- 3D 点云 → Flow Matching
- 从随机点云 → 目标形状
DreamBooth3D:
- 个性化 3D 资产生成

9.4 科学计算#

分子生成:
- GPS: "Generative Flow Networks" (GFlowNets)
- 用 Flow Matching 做贝叶斯推断
强化学习:
- Q-Flow, Decision Flow Matching

10. Flow Matching 训练的理论细节#

10.1 损失函数的推导#

从 CFM 目标出发:

LCFM(θ)=Et,x0,ϵvθ(xt,t)vt(xtx0)2\mathcal{L}_{\text{CFM}}(\theta) = \mathbb{E}_{t, x_0, \epsilon} \left\| v_\theta(x_t, t) - v_t(x_t | x_0) \right\|^2

其中 xt=ϕ(x0,ϵ;t)x_t = \phi(x_0, \epsilon; t) 是设计的路径。

def cfm_loss_detailed(model, x0, eps, t):
"""
完整推导的 CFM 损失,包含权重选项。
"""
# 可选: 重要性采样时间步 (与 σ² 相关)
lambda_t = 1.0 # 或 lambda(t)
# 路径
xt = (1 - t) * x0 + t * eps
# 速度
true_v = eps - x0
# 预测
pred_v = model(xt, t)
# 加权 MSE
return (lambda_t * (pred_v - true_v) ** 2).mean()

10.2 训练分布的等价性#

一个深刻结论:边际流匹配目标的梯度等于条件流匹配目标的期望

θLFM=Epdata[θLCFM]\nabla_\theta \mathcal{L}_{\text{FM}} = \mathbb{E}_{p_{\text{data}}} \left[ \nabla_\theta \mathcal{L}_{\text{CFM}} \right]

这意味着可以直接用批量平均来近似期望——这就是标准的小批量 SGD。

10.3 时间采样策略#

时间采样公式适合场景
均匀tU[0,1]t \sim \mathcal{U}[0,1]通用(Rectified Flow 默认)
重要性加权tσt2t \propto \sigma_t^2DDPM 风格路径
噪声调度ttαt\alpha_t 相关高分辨率生成

11. 完整实现:从零到 SD3#

11.1 完整的 Flow Matching 模型#

import math
import torch
import torch.nn as nn
import torch.nn.functional as F
class FlowMatchingUNet(nn.Module):
"""简化的 Flow Matching UNet (与 DDPM 相同架构, 不同目标)。"""
def __init__(self, dim=64, channels=3):
super().__init__()
self.time_embed = nn.Sequential(
nn.Linear(64, dim * 4),
nn.SiLU(),
nn.Linear(dim * 4, dim * 4),
)
# ... 标准 UNet 组件
self.final = nn.Conv2d(dim, channels, 3, padding=1)
def forward(self, xt, t, cond=None):
"""
xt: (B, C, H, W) — 任意时间步的中间状态
t: (B,) — 时间步 ∈ [0, 1]
"""
# 1) 时间嵌入
t_emb = self._timestep_embedding(t * 1000, 64)
t_emb = self.time_embed(t_emb)
# 2) 噪声图
h = self.init_conv(xt)
# 3) 下采样 + 中间层 + 上采样 (标准 UNet)
# ... (省略中间层)
# 4) 输出速度场 v_θ(x_t, t)
return self.final(h)
class FlowMatchingTransformer(nn.Module):
"""基于 DiT 的 Flow Matching Transformer。"""
def __init__(self, hidden_size=1024, depth=12):
super().__init__()
self.time_embed = TimestepEmbedder(hidden_size)
self.blocks = nn.ModuleList([TransformerBlock(hidden_size) for _ in range(depth)])
def forward(self, xt, t, cond=None):
"""
返回速度场 v_θ(xt, t) = ε - x0
与 DiT/MMDiT 完全兼容!
"""
t_emb = self.time_embed(t)
h = self.patch_embed(xt)
for block in self.blocks:
h = block(h, t_emb)
return self.unpatchify(self.final_layer(h))

11.2 与 MMDiT 的对应#

# MMDiT (SD3/FLUX) 中的 Flow Matching
# 1) backbone: DiT → MMDiT (双流)
# 2) 路径: DDPM → Rectified Flow (直线)
# 3) 目标: 噪声 ε → 速度 v = ε - x0
# MMDiT 的 forward 就是:
def mmdit_forward(xt, t, txt_tokens):
v_pred = mmdit(xt, t, txt_tokens) # 预测速度场
return v_pred
def mmdit_training_loss(xt, x0, eps, t, txt_tokens):
true_v = eps - x0 # 速度场 (解析)
pred_v = mmdit_forward(xt, t, txt_tokens)
return F.mse_loss(pred_v, true_v) # Flow Matching 目标

一图总结

Flow Matching (理论框架)
├── 路径选择 ──→ Rectified Flow (直线, OT 最优)
├── 网络架构 ──→ DiT / MMDiT
└── 训练目标 ──→ MSE(预测速度, 真实速度)
SD3 / FLUX = Rectified Flow + MMDiT + CFG + 蒸馏

12. 局限与开放问题#

12.1 当前局限#

问题描述
直线陷阱RF 直线路径在高维空间未必是全局最优传输
分布外泛化学到的向量场在训练分布外可能不可靠
** ODE 数值误差**少步采样时 ODE 积分误差累积
条件生成复杂度CFG 仍需双倍计算(cond + uncond)

12.2 开放研究方向#

理论:
├── 最优传输路径的全局最优性证明
├── 高维最优传输的计算高效近似
└── Flow Matching 与 W-GAN 的理论联系
实践:
├── 任意到任意的分布传输 (任意起点 → 任意终点)
├── 多模态 Flow Matching (同时建模多个分布)
├── Flow Matching + RLHF (对齐)
└── 时空 Flow (视频、3D)

13. 总结#

13.1 核心要点#

维度关键要点
核心思想学向量场 vθv_\theta,通过 ODE 把噪声推向数据
核心定理条件流匹配 (CFM) = 边际流匹配 (FM),无需 KL 散度
最优传输Rectified Flow 是 OT 最优路径 → 最短 ODE 轨迹
目标函数L=Evθvt2\mathcal{L} = \mathbb{E} \| v_\theta - v_t \|^2(简单 MSE)
ODE vs SDEODE 快速采样,SDE 稳定训练
与 DDPMvtv_tϵθ\epsilon_\theta 成正比,Rectified Flow 让关系最简
与 MMDiTMMDiT + RF = SD3 / FLUX 的训练基础

13.2 一句话总结#

Flow Matching 用”向量场 + ODE”的语言,把扩散模型、统一在最简洁的 MSE 目标下——而 Rectified Flow 的直线路径,则是这条理论框架下最优传输路径的工程实现,也是 SD3/FLUX/MMDiT 高质量、快推理的数学保证。

13.3 推荐学习资源#

论文:
- Flow Matching (Lipman et al., 2022): "Flow Matching for Causal Inference"
- Conditional Flow Matching (Tong et al., 2023)
- Rectified Flow (Liu et al., 2022): "Flow Straight and Fast"
- RF-LR (2024): "Scaling Rectified Flow for Image Understanding" (SD3 论文)
- Consistency Models (Song et al., 2023)
代码:
- facebookresearch/flow_matching (官方实现)
- stabilityai/sd3-ref (Rectified Flow + MMDiT)
- black-forest-labs/FLUX (FLUX.1)
扩展:
- GFlowNets: 离散的 Flow Matching
- Diffusion Schrödinger Bridge: 最优传输的另一种视角

一句话总结:Flow Matching 把”扩散模型”从工程技巧升华为优雅理论——它的核心定理证明”简单 MSE = 最优生成”,而 Rectified Flow 的直线路径让这一理论在 SD3/FLUX 中落地成工程现实。

文章分享

如果这篇文章对你有帮助,欢迎分享给更多人!

深入理解 Flow Matching:生成建模的统一理论框架
https://aiattnstudio.link/posts/flow-matching/
作者
Federico
发布于
2026-07-16
许可协议
CC BY-NC-SA 4.0
Profile Image of the Author

Federico

AI Research Lab

Hello, I'm Federico.

关于实验室 / About
公告

欢迎来到Federico的个人博客

分类
标签
站点统计
57文章
7分类
404标签
1
1. 为什么要理解 Flow Matching?
1.1 扩散模型的”理论黑箱”问题
1.2 一句话概括
1.3 直观理解
1.4 Flow Matching 的历史脉络
2
2. 概率论基础:从噪声到数据的路径
2.1 目标是什么?
2.2 连续时间路径
2.3 三种经典路径对比
2.4 边际分布
3
3. 常微分方程 (ODE) 与概率流
3.1 核心:从路径到向量场
3.2 向量场的定义
3.3 概率流方程(连续性方程)
3.4 ODE 求解(采样)
4
4. Flow Matching 定理
4.1 条件流匹配目标 (Conditional Flow Matching, CFM)
4.2 证明的直觉
4.3 为什么这个定理如此重要?
4.4 对 DDPM 的重新解释
5
5. 最优传输与 Rectified Flow
5.1 什么是”最优传输路径”?
5.2 Rectified Flow 的最优传输证明
5.3 最速下降流 (Gradient Flow)
5.4 路径设计空间
6
6. 扩展:随机微分方程 (SDE) 视角
6.1 为什么需要 SDE?
6.2 VP-SDE(Variance Preserving)和 VE-SDE
6.3 ODE vs SDE 采样
6.4 从 SDE 到 ODE 的技巧
7
7. 训练与采样完整流程
7.1 完整训练代码
7.2 推理加速技术
7.2.1 高阶 ODE 求解器
7.2.2 自适应步长
7.3 一致性模型蒸馏
8
8. Flow Matching 与 Score Matching 的关系
8.1 核心区别
8.2 数学上的精确关系
9
9. Flow Matching 的应用版图
9.1 图像生成
9.2 音频生成
9.3 3D 生成
9.4 科学计算
10
10. Flow Matching 训练的理论细节
10.1 损失函数的推导
10.2 训练分布的等价性
10.3 时间采样策略
11
11. 完整实现:从零到 SD3
11.1 完整的 Flow Matching 模型
11.2 与 MMDiT 的对应
12
12. 局限与开放问题
12.1 当前局限
12.2 开放研究方向
13
13. 总结
13.1 核心要点
13.2 一句话总结
13.3 推荐学习资源