生成模型 (3.1):Flow-based Method
引言
最近看到Qwen-Edit,感觉效果特别帅。借此机会想系统学习一下Flow Matching的数学原理,这里记录一下。
后面内容完全是在翻译这篇论文:Flow Matching Guide and Code。感谢FAIR的大佬们,这种从1+1开始教的self-contained教程对我这种外行人真的非常重要。
PS:学物理的放过咱们cs吧,是真难懂啊。。。
一开始以为Flow Matching的作者 Ricky Tian Qi Chen 就是xgBoost的陈天奇,但发现似乎不是哈哈
一、大白话理解 (也没那么白。。。)
Flow Matching (FM) 在干什么?
- Flow Matching的终极任务是学习一个向量场 (vector field);
- 这个向量场通过ODE的形式,定义了一个流 (flow)
; - 一个流是定义在
上的一系列可逆变换,这些变换是在时间上连续的; - Flow Matching就是在学习一个流,将初始分布(一般是简单分布如
)中的某个样本 映射到目标分布上,即 。
二、FM的形式化定义
给定一个定义在
为了完成这个任务,Flow Matching搭建了一条概率路径
更加具体地说,我们训练的是一个能够描述样本瞬时速度的神经网络。这个网络后面被用于沿着概率路径
在训练完成后,我们 (1) 从
在形式上,向量场是一个依赖于时间的函数
其中
我们称向量场
其中
三、FM的步骤
- 第一步:从目标分布
中收集一些样本作为训练集。 - 第二步:设计一条连续时间的概率路径
,满足 且 。 - 第三步:使用回归 (Regression) 的方式,训练一个参数为
的向量场 。 - 第四步:从
中采样一个新样本 ,根据向量场所定义的概率路径,得到 上的一个样本 。
四、实际场景下
假设初始分布为
这个路径也被称为条件最优运输 (conditional optimal transport)。
我们可以通过线性插值,定义从
在训练时,我们的目标是去最小化当前的向量场
其中
然而这个loss非常难实现,因为目标向量场
具体来说,我们需要求解下面这个ODE:
其中,
这个ODE的解为如下的条件向量场(证明见附录):
我们定义如下的条件FM损失:
将ODE的解代入公式
其中
附录
公式(8)的证明
从公式
由公式
代入上式得:
Q.E.D.