生成模型 (3.3):Flow Matching
引言
在上一章中,我们介绍了流模型。流模型通过模拟速度场对应的ODE,来得到参数的最大似然估计。
然而在实际优化过程中,计算这个最大似然以及其梯度需要非常精准的ODE模拟,只有完全正确的模拟结果才能得到无偏的梯度估计。
下面要介绍的流匹配,则是一种无需模拟的方法。
同样的,以下内容完全来自于Flow Matching Guide and Code。再次感谢FAIR的大佬们。
一、概览
给定一个初始分布
因此,流匹配主要有以下几步:
- 确定一个已知的初始分布
和一个未知的目标分布 。 - 构建一条概率路径
,满足 。 - 学习一个速度场
,使其能够生成 。 - 从学到的模型中采样,生成一个目标分布中的样本。
学习速度场时,我们一般最小化下面的回归损失:
其中,
这里有两个问题仍未解决:如何构建概率路径
二、构建概率路径
流匹配使用了一种【条件化】的策略,极大程度上简化了概率路径以及其对应向量场的构建难度。
具体来说,当我们只考虑目标分布上的某一个样本
这些条件概率都需要满足下面的约束:
其中
进一步我们可以通过对所有条件概率路径求期望得到【边缘概率路径】:
这个边缘概率路径满足:
正如在第一章中提到的,一种最常见的构建条件概率路径的方法是使用如下高斯分布:
三、构建速度场
得到边缘概率路径
由于
最终,我们得到的速度场为:
注意这里用的是后验概率
3.1. 合法性证明
下面我们将说明,在引入一些假设后,公式
我们所用到的工具是在第二章中介绍到的【质量守恒定律】。质量守恒定律是说,给定一个满足局部利普希茨性质的可积速度场
其中,
不失一般性地,我们可以把上面引入的目标样本
生成的速度场变为:
且满足:
为了说明我们的结论,首先需要引入下面的假设:
, , 具有有界支撑集,即满足 的样本位于某些无界集合中, 。
这四点假设并不难满足。首先,大多数常用的概率分布(如高斯分布)都是无限次平滑的,远远超出
有了这些准备,我们就可以证明下面的定理3(详细证明见附录)。
Theorem 3: (Marginalization Trick) 在满足上述假设的情况下,如果
这里条件可积是指:
四、Flow Matching损失函数
至此,我们证明了存在一个速度场
下面,我们就是要设计一个可微的损失函数,去学习一个速度场
然而,真实速度场
为了求解这个问题,我们需要引入一类特殊的损失函数,称为Bregman散度。Bregman散度可以利用条件速度
具体来说,Bregman散度衡量了两个向量
其中,
注意到,
当我们选择
Bregman散度一个非常重要的性质是:Bregman散度的梯度满足【仿射不变性】(证明见附录)。具体来说:
有了这种仿射不变性,我们就可以交换梯度和期望的位置:
Flow Matching的损失函数中使用了Bregman散度,用于度量
正如上面提到的,
可以证明,公式
标准的CFM损失中,时间一般是从均匀分布中采样的,即
但也有研究证明,在大规模的图像生成任务中,从其他的分布
比如说,在Stable Diffusion 3中,使用双峰分布 (Bimodal Distribution)来采样时间,显著增加接近0或1的时间步的采样概率。
附录
Apx1. 定理3的证明
首先我们证明
我们有:
其中第一、三行交换了积分和求导(散度)的顺序。这是可行的,因为
然后我们再证明
由于
第一个不等号的成立是根据Jensen不等式:
因此,我们证明了
根据质量守恒定律,
Apx2. Bregman散度的梯度仿射不变性
下面我们证明公式
首先我们推导Bregman散度的梯度。
最后一步利用了Hessian矩阵的对称性。
因此:
公式
公式
得证。
Apx3. FM Loss和CFM Loss的等价性
下面我们证明公式
这里面使用到的证明技巧:
- 第三、第六行用到了求导链式法则,变换了求导变量
- 第四行应用了公式
- 第五行应用了Bregman散度的梯度放射不变性,即公式
- 第七行应用了贝叶斯公式