生成模型 (3.3):Flow Matching

生成模型 #生成模型 #流匹配 #Bregman 散度 约 24 分钟 · 8378 字

引言

在上一章中,我们介绍了流模型。流模型通过模拟速度场对应的ODE,来得到参数的最大似然估计。

然而在实际优化过程中,计算这个最大似然以及其梯度需要非常精准的ODE模拟,只有完全正确的模拟结果才能得到无偏的梯度估计。

下面要介绍的流匹配,则是一种无需模拟的方法。

同样的,以下内容完全来自于Flow Matching Guide and Code。再次感谢FAIR的大佬们。

一、概览

给定一个初始分布 和目标分布 ,流匹配的终极目的是训练一个流模型 ,使得 能够生成概率路径 ,且 。

因此,流匹配主要有以下几步:

  1. 确定一个已知的初始分布 和一个未知的目标分布 。
  2. 构建一条概率路径 ,满足 。
  3. 学习一个速度场 ,使其能够生成 。
  4. 从学到的模型中采样,生成一个目标分布中的样本。

学习速度场时,我们一般最小化下面的回归损失:

(1)

其中, 是一个定义在 上的度量,常用的如L2范数。

这里有两个问题仍未解决:如何构建概率路径 、如何构建真实的速度场 。下面我们逐个说明。

二、构建概率路径

流匹配使用了一种【条件化】的策略,极大程度上简化了概率路径以及其对应向量场的构建难度。

具体来说,当我们只考虑目标分布上的某一个样本 ,我们就可以生成一些【条件概率路径】 。

这些条件概率都需要满足下面的约束:

(2)

其中 是delta函数。

进一步我们可以通过对所有条件概率路径求期望得到【边缘概率路径】:

(3)

这个边缘概率路径满足:。

正如在第一章中提到的,一种最常见的构建条件概率路径的方法是使用如下高斯分布:

(4)

三、构建速度场

得到边缘概率路径 后,我们可以构建出能够生成 的一个速度场 。

由于 是由多个条件概率路径 组合而成,因此 也是由多个【条件速度场】 组合而成,且满足: 能够生成 。

最终,我们得到的速度场为:

(5)

注意这里用的是后验概率 ,表示当前样本 能够生成目标样本 的概率。由贝叶斯公式,这个后验概率可以用下面的公式来计算:

(6)

3.1. 合法性证明

下面我们将说明,在引入一些假设后,公式得到的速度场能够生成公式中的概率路径。

我们所用到的工具是在第二章中介绍到的【质量守恒定律】。质量守恒定律是说,给定一个满足局部利普希茨性质的可积速度场 以及一条概率路径 , 能够生成 当且仅当满足下面的公式:

(7)

其中, 称为散度矩阵。公式也被称为连续性方程。

不失一般性地,我们可以把上面引入的目标样本 推广为任意的随机向量 ,则对应的概率路径变为:

(8)

生成的速度场变为:

(9)

且满足: 能够生成 。

为了说明我们的结论,首先需要引入下面的假设:

  • ,
  • ,
  • 具有有界支撑集,即满足 的样本位于某些无界集合中,
  • 。

这四点假设并不难满足。首先,大多数常用的概率分布(如高斯分布)都是无限次平滑的,远远超出 的要求。其次,真实世界的数据(如图片、声音、文本)几乎总是有界的。最后,我们可以让条件 满足 且 ,就能满足第四条假设。

有了这些准备,我们就可以证明下面的定理3(详细证明见附录)。

Theorem 3: (Marginalization Trick) 在满足上述假设的情况下,如果 是【条件可积】的,且能够生成条件概率路径 ,则 能够生成 。

这里条件可积是指:

(10)

四、Flow Matching损失函数

至此,我们证明了存在一个速度场 能够生成从 到 的一条概率路径 。

下面,我们就是要设计一个可微的损失函数,去学习一个速度场 ,使其尽可能接近真实的速度场 。

然而,真实速度场 非常难以计算。因为它需要对整个训练集中的目标样本 进行积分(见公式)。

为了求解这个问题,我们需要引入一类特殊的损失函数,称为Bregman散度。Bregman散度可以利用条件速度 对学习 的梯度进行无偏的估计。

具体来说,Bregman散度衡量了两个向量 之间的距离:

(11)

其中, 是一个定义在凸集 上的严格凸函数。

注意到, 实际上就是在求 在 上的一阶Taylor近似,因此Bregman散度实际上度量的是函数 与其一阶Taylor估计的距离。

当我们选择 时,Bregman散度就变为我们熟知的欧式距离:。

Bregman散度一个非常重要的性质是:Bregman散度的梯度满足【仿射不变性】(证明见附录)。具体来说:

(12)

有了这种仿射不变性,我们就可以交换梯度和期望的位置:

(13)

Flow Matching的损失函数中使用了Bregman散度,用于度量 和 的距离:

(14)

正如上面提到的, 是难解的。因此我们考虑下面的条件损失:

(15)

可以证明,公式和公式有着相同的梯度。因此二者对于优化问题来说是等价的。证明见附录。

标准的CFM损失中,时间一般是从均匀分布中采样的,即 。

但也有研究证明,在大规模的图像生成任务中,从其他的分布 中采样时间往往会得到更好的效果。此时CFM损失变为:

(16)

比如说,在Stable Diffusion 3中,使用双峰分布 (Bimodal Distribution)来采样时间,显著增加接近0或1的时间步的采样概率。

(17)

附录

Apx1. 定理3的证明

首先我们证明 和 满足连续性方程。

我们有:

(18)

其中第一、三行交换了积分和求导(散度)的顺序。这是可行的,因为 和 都是 函数,且具有有界的支撑集。第二行成立是因为 能够生成 。第四行成立是因为 。第五行成立是因为贝叶斯公式。第六行成立则是代入了公式。

然后我们再证明 是一个满足局部利普希茨性质的可积速度场。

由于 函数满足局部利普希茨条件,因此我们只需说明 是一个 函数。这是成立的,因为 和 都是 函数且 。

的可积性是因为 的条件可积性:

(19)

第一个不等号的成立是根据Jensen不等式:

(20)

因此,我们证明了 是一个满足局部利普希茨性质的可积速度场,且 和 满足连续性方程。

根据质量守恒定律, 能够生成 。

Apx2. Bregman散度的梯度仿射不变性

下面我们证明公式 。

首先我们推导Bregman散度的梯度。

(21)

最后一步利用了Hessian矩阵的对称性。

因此:

(22)

公式 的等式左边

(23)

公式 的等式右边

(24)

得证。

Apx3. FM Loss和CFM Loss的等价性

下面我们证明公式和公式有着相同的梯度。

(25)

这里面使用到的证明技巧:

  • 第三、第六行用到了求导链式法则,变换了求导变量
  • 第四行应用了公式
  • 第五行应用了Bregman散度的梯度放射不变性,即公式
  • 第七行应用了贝叶斯公式