从入门到入土 2:Flow Matching——如何训练一个生成模型
从入门到入土 2:Flow Matching——如何训练一个生成模型
在上一篇文章中,我们构建了流模型和扩散模型的基础平台:将生成问题转化为了一个微分方程模拟问题。但最核心的谜题尚未解开——如何训练那个驱动一切的核心神经网络?答案就是 Flow Matching。
在上一篇文章中,我们构建了流模型和扩散模型的基本框架:从简单的噪声分布 pinitp_{\text{init}}pinit(比如高斯噪声)出发,通过模拟一个由神经网络学得的 utθu_t^\thetautθ 定义的微分方程,最终到达复杂的数据分布 pdatap_{\text{data}}pdata。
然而,我们只明确了模型“该做什么”(模拟 ODE/SDE),却未解决“该怎么做”的问题,即如何优化参数 θ\thetaθ。换句话说,我们应如何调整 utθu_t^\thetautθ,才能确保其能将噪声转化为符合目标分布的样本?
在众多训练方法中,Flow Matching(流匹配) 因其简洁性与高效性脱颖而出,成为当前主流大规模图像和视频生成模型(如 Stable Diffusion 3)的核心训练范式。本文将拆解这一算法,阐述其如何实现“无中生有”的生成能力。
1. 核心问题:明确模型的学习目标
我们首先重新审视目标。对于一个已定义的流模型:
ddtXt=utθ(Xt),X0∼pinit \frac{d}{dt} X_t = u_t^\theta(X_t), \quad X_0 \sim p_{\text{init}} dtdXt=utθ(Xt),X0∼pinit
我们期望模拟该 ODE 至 t=1t=1t=1 时,X1X_1X1 的分布能够匹配数据分布 pdatap_{\text{data}}pdata。那么,理想的、目标性的向量场 uttargetu_t^{\text{target}}uttarget 应具备何种性质?
它应能够构建一条从噪声分布 pinitp_{\text{init}}pinit 到数据分布 pdatap_{\text{data}}pdata 的平滑、连续的概率路径。该路径在数学上被称为概率路径(Probability Path) ptp_tpt,满足 p0=pinitp_0 = p_{\text{init}}p0=pinit 和 p1=pdatap_1 = p_{\text{data}}p1=pdata。
问题可转化为:
我们期望学习一个神经网络 utθu_t^\thetautθ,使其尽可能逼近这条目标概率路径所对应的速度场 uttargetu_t^{\text{target}}uttarget。
2. 构建概率路径:从条件路径到边际路径
直接定义从复杂的 pdatap_{\text{data}}pdata 到简单的 pinitp_{\text{init}}pinit 的概率路径是困难的。Flow Matching 的核心思想在于:先为数据集中的每一个样本 zzz 定义一条从噪声到该特定数据点的条件概率路径,再通过边际化操作将这些路径整合起来。
这可以理解为分治的思想,很自然,而整合的过程相当于一个全概率公式。
2.1 条件概率路径:为每个数据点定制路径
对于数据集中的任意一个数据点 zzz,我们可以轻松定义一条连接噪声与该点的路径。一个常用且有效的选择是高斯条件概率路径:
pt(x∣z)=N(αtz,βt2Id) p_t(x|z) = \mathcal{N}(\alpha_t z, \beta_t^2 I_d) pt(x∣z)=N(αtz,βt2Id)
该公式表示:在时间 ttt,路径上的点 xxx 服从一个均值为 αtz\alpha_t zαtz、协方差为 βt2Id\beta_t^2 I_dβt2Id 的高斯分布。
- 当 t=0t=0t=0 时,令 α0=0,β0=1\alpha_0 = 0, \beta_0 = 1α0=0,β0=1,则 p0(⋅∣z)=N(0,Id)p_0(\cdot|z) = \mathcal{N}(0, I_d)p0(⋅∣z)=N(0,Id),即标准高斯噪声。
- 当 t=1t=1t=1 时,令 α1=1,β1=0\alpha_1 = 1, \beta_1 = 0α1=1,β1=0,则 p1(⋅∣z)=N(z,0)p_1(\cdot|z) = \mathcal{N}(z, 0)p1(⋅∣z)=N(z,0),即坍缩到点 zzz 的狄拉克分布 δz\delta_zδz,表示确定性地得到 zzz。
这里的 αt\alpha_tαt 和 βt\beta_tβt 称为噪声调度器(Noise Scheduler),控制从噪声到数据的插值方式。一个简单的线性调度是 αt=t,βt=1−t\alpha_t = t, \beta_t = 1-tαt=t,βt=1−t。此时,路径上的点可表示为 x=tz+(1−t)ϵx = t z + (1-t)\epsilonx=tz+(1−t)ϵ,其中 ϵ∼N(0,Id)\epsilon \sim \mathcal{N}(0, I_d)ϵ∼N(0,Id)。
为什么称为“条件”?因为该路径以特定的数据点 zzz 为条件。不同的 zzz 对应不同的、从噪声到该点的路径。
2.2 边际概率路径:整合所有路径
为每个数据点 zzz 定义了条件路径 pt(x∣z)p_t(x|z)pt(x∣z) 后,若我们从数据分布中随机抽取 zzz,再从 pt(x∣z)p_t(x|z)pt(x∣z) 中采样 xxx,则采样点 xxx 的分布即为边际概率路径 pt(x)p_t(x)pt(x):
pt(x)=∫pt(x∣z)pdata(z)dz p_t(x) = \int p_t(x|z) p_{\text{data}}(z) dz pt(x)=∫pt(x∣z)pdata(z)dz
这相当于将每个数据点对应的条件路径进行加权平均,形成一条覆盖整个数据分布的“总路径”。可以验证,该“总路径”的起点是噪声(p0=pinitp_0 = p_{\text{init}}p0=pinit),终点是数据分布(p1=pdatap_1 = p_{\text{data}}p1=pdata)。
3. 学习目标速度场:从已知中推断未知
至此,我们已构建了目标概率路径 pt(x∣z)p_t(x|z)pt(x∣z) 和 pt(x)p_t(x)pt(x),但仍未直接获得对应于边际概率路径 pt(x)p_t(x)pt(x) 的目标速度场 uttarget(x)u_t^{\text{target}}(x)uttarget(x)。
然而,关键在于:对于每个简单的条件概率路径 pt(x∣z)p_t(x|z)pt(x∣z),其对应的速度场 uttarget(x∣z)u_t^{\text{target}}(x|z)uttarget(x∣z) 可通过解析形式获得,无需依赖微分方程模拟。
特别地,对于高斯路径,该条件速度场为:
uttarget(x∣z)=(α˙t−β˙tβtαt)z+β˙tβtx u_t^{\text{target}}(x|z) = \left(\dot{\alpha}_t - \frac{\dot{\beta}_t}{\beta_t}\alpha_t\right)z + \frac{\dot{\beta}_t}{\beta_t} x uttarget(x∣z)=(α˙t−βtβ˙tαt)z+βtβ˙tx
其中 α˙t,β˙t\dot{\alpha}_t, \dot{\beta}_tα˙t,β˙t 为 αt,βt\alpha_t, \beta_tαt,βt 对时间的导数。该公式表明,若已知高斯分布的均值和方差随时间的演化规律(由 αt,βt\alpha_t, \beta_tαt,βt 描述),则可直接导出其漂移速度。
其意义在于:我们掌握了每条“条件路径”上任意点的精确速度信息,即“在给定数据点 zzz 和时间 ttt 的条件下,正确速度的解析表达式”。
3.1 条件流匹配:转化为回归任务
我们未知边际路径的速度 uttarget(x)u_t^{\text{target}}(x)uttarget(x),但已知条件路径的速度 uttarget(x∣z)u_t^{\text{target}}(x|z)uttarget(x∣z)。边际化技巧指出,边际速度场可通过条件速度场加权平均得到:
uttarget(x)=∫uttarget(x∣z)pt(x∣z)pdata(z)pt(x)dz u_t^{\text{target}}(x) = \int u_t^{\text{target}}(x|z) \frac{p_t(x|z) p_{\text{data}}(z)}{p_t(x)} dz uttarget(x)=∫uttarget(x∣z)pt(x)pt(x∣z)pdata(z)dz
后面的这一坨分式可以看作一个贝叶斯,等价于 pt(z∣x)p_t(z|x)pt(z∣x)。
基于此,Flow Matching 提出了一个简洁的训练策略:不再直接学习复杂的边际速度场,而是训练神经网络 utθu_t^\thetautθ 去拟合已知的、简单的条件速度场 uttarget(x∣z)u_t^{\text{target}}(x|z)uttarget(x∣z)。
这导向了条件流匹配(Conditional Flow Matching) 损失函数:
LCFM(θ)=Et,z,x∼pt(⋅∣z)[∥utθ(x)−uttarget(x∣z)∥2] \mathcal{L}_{\text{CFM}}(\theta) = \mathbb{E}_{t, z, x \sim p_t(\cdot|z)} \left[ \| u_t^\theta(x) - u_t^{\text{target}}(x|z) \|^2 \right] LCFM(θ)=Et,z,x∼pt(⋅∣z)[∥utθ(x)−uttarget(x∣z)∥2]
3.2 为何可行:损失函数的等价性
一个自然的问题是:训练模型去拟合条件速度场,如何能保证它学会边际速度场?
答案在于,当我们将所有条件速度场的回归目标进行平均时,其效果等价于直接回归到边际速度场(这可以通过数学严格证明,感兴趣的读者可自行搜索),即:
LFM(θ)=LCFM(θ)+C \mathcal{L}_{\text{FM}}(\theta) = \mathcal{L}_{\text{CFM}}(\theta) + C LFM(θ)=LCFM(θ)+C
其中 LFM(θ)\mathcal{L}_{\text{FM}}(\theta)LFM(θ) 是我们真正希望优化的边际速度场损失,CCC 是与参数 θ\thetaθ 无关的常数。
这意味着,最小化可计算的 LCFM\mathcal{L}_{\text{CFM}}LCFM 等价于最小化难以直接计算的 LFM\mathcal{L}_{\text{FM}}LFM。这一性质使得 Flow Matching 在理论上成立且易于实现。
4. 算法总结与实现
Flow Matching 的完整训练流程极其简洁:
- 采样数据点 zzz:从数据集中随机抽取一个样本。
- 采样噪声 ϵ\epsilonϵ:从标准高斯分布 N(0,Id)\mathcal{N}(0, I_d)N(0,Id) 中随机采样。
- 采样时间 ttt:从区间 [0,1][0, 1][0,1] 中均匀随机采样。
- 构造带噪样本 xtx_txt:根据选定的噪声调度器(如线性调度 αt=t,βt=1−t\alpha_t = t, \beta_t = 1-tαt=t,βt=1−t)混合数据和噪声:xt=αtz+βtϵx_t = \alpha_t z + \beta_t \epsilonxt=αtz+βtϵ。
- 计算目标速度 vtargetv_{\text{target}}vtarget:对于高斯路径,目标速度具有解析解:vtarget=α˙tz+β˙tϵv_{\text{target}} = \dot{\alpha}_t z + \dot{\beta}_t \epsilonvtarget=α˙tz+β˙tϵ。对于线性调度,该式简化为 vtarget=z−ϵv_{\text{target}} = z - \epsilonvtarget=z−ϵ。
- 计算损失:将 xtx_txt 和 ttt 输入神经网络 utθu_t^\thetautθ,得到预测速度,并计算其与 vtargetv_{\text{target}}vtarget 的均方误差。
- 参数更新:通过反向传播优化网络参数 θ\thetaθ。
算法:Flow Matching 训练过程(以线性调度为例)
输入:数据集,神经网络 utθu_t^\thetautθ
- 从数据集中采样一个数据点 zzz
- 采样随机时间 t∼Uniform[0,1]t \sim \text{Uniform}[0, 1]t∼Uniform[0,1]
- 采样噪声 ϵ∼N(0,Id)\epsilon \sim \mathcal{N}(0, I_d)ϵ∼N(0,Id)
- 构造带噪样本:x=t⋅z+(1−t)⋅ϵx = t \cdot z + (1-t) \cdot \epsilonx=t⋅z+(1−t)⋅ϵ
- 计算目标速度:vtarget=z−ϵv_{\text{target}} = z - \epsilonvtarget=z−ϵ
- 计算损失:L=∥utθ(x)−vtarget∥2\mathcal{L} = \| u_t^\theta(x) - v_{\text{target}} \|^2L=∥utθ(x)−vtarget∥2
- 更新 θ\thetaθ 以最小化 L\mathcal{L}L
Flow Matching 的训练过程是无模拟的(simulation-free),即无需像采样时那样逐步模拟整个 ODE 轨迹。它仅需学习将任意带噪样本“指向”数据点或目标速度的方向。训练完成后,即可通过欧拉法等数值方法模拟 ODE,从噪声生成新的样本。
小结
Flow Matching 的精妙之处在于,它将一个复杂的、涉及分布变换的生成问题,转化为一个简单的、基于数据点的回归问题。它无需引入变分下界或处理对数似然,仅依靠均方误差损失和明确的概率路径设计即可实现。这一框架已成为当前主流生成模型的核心训练范式。
在下一篇文章中,我们将探讨 Flow Matching 的关联方法——Score Matching(分数匹配),它从另一个视角揭示扩散模型的工作原理,并最终引出强大的条件生成与引导技术。
AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。
更多推荐



所有评论(0)