LLM WIKI · PAPER

MIT 6.S184 · LECTURE NOTES · 2026 · 中文全文译稿

An Introduction to
Flow Matching and
Diffusion Models

Peter Holderrieth · Ezra ErivesMIT Class 6.S184Generative AI With Stochastic Differential Equations
PDF
84 页
Figures
22
Algorithms
8
Equations
152
References
51

绪论

从数据中创建噪声很容易;从噪声中创建数据是生成建模。Song 等人 [43]

概览

近年来,我们都见证了人工智能(AI)的巨大革命。像 Nano BananaStable Diffusion 3 这样的图像生成器可以生成各种风格的照片级逼真和艺术图像,像 Meta 的 VEO-3 这样的视频模型可以生成高度逼真的电影片段,而像 ChatGPT 这样的大型语言模型可以对文本提示生成看似人类水平的响应。这场革命的核心是 AI 系统的一项新能力:生成对象的能力。虽然前几代 AI 系统主要用于预测,但这些新的 AI 系统具有创造性:它们根据用户指定的输入进行想象或生成新对象。这种生成式 AI 系统是近期 AI 革命的核心。

本课程的目标是教授你两种最广泛使用的生成式 AI 算法:Denoising Diffusion Model [43] 和 Flow Matching [25, 27, 1, 26]。这些模型是最佳图像、音频和视频 Generative Model(例如,Nano BananaFLUXVEO-3)的支柱,并且最近已成为蛋白质结构等科学应用中的最先进技术(例如,AlphaFold3 是一个 Diffusion Model)。毫无疑问,理解这些模型确实是一项非常有用的技能。

所有这些 Generative Model 都通过迭代地将噪声转换为数据来生成对象。这种从噪声到数据的演化是通过模拟常微分或 Stochastic Differential Equation(ODE/SDE) 来实现的。Flow Matching 和 Denoising Diffusion Model 是一系列技术,使我们能够使用深度神经网络大规模地构建、训练和模拟此类 ODE/SDE。虽然这些模型实现起来相当简单,但 SDE 的技术性质可能使这些模型难以理解。在本课程中,我们提供了关于微分方程的必要数学工具的自包含介绍,使您能够系统地理解这些模型。然后,我们逐步解释最先进的图像和视频生成器的现代技术栈。除了广泛适用之外,我们相信 Flow Model 和 Diffusion Model 背后的理论本身就很优雅。因此,最重要的是,我们希望这门课程能给您带来很多乐趣。

备注 1(补充资源)

虽然这些讲义是自包含的,但我们鼓励您使用另外两个资源:

  1. 讲座录音: 这些以讲座形式引导您学习每个部分。
  2. 实验: 这些指导您从零开始实现自己的 Diffusion Model。我们强烈建议您“动手实践”并编写代码。

您可以在我们的课程网站上找到这些:https://diffusion.csail.mit.edu/

课程结构

我们简要概述本文档。

  • 第 1 节,生成建模作为采样: 我们形式化“生成”图像、视频、蛋白质等的含义。我们将例如“如何生成一张狗的图片?”的问题转化为从概率分布中采样的更精确问题。
  • 第 2 节,流和 Diffusion Model: 我们解释生成的机制。正如您从本课程名称中猜到的那样,这种机制包括模拟常微分和 Stochastic Differential Equation。我们介绍微分方程,并解释如何使用它们来构建 Generative Model。
  • 第 3 节,Flow Matching: 接下来,我们解释并推导 Flow Matching,这是一种简单且可扩展的算法,位于所有上述大规模 Generative Model(如 Stable Diffusion、Nano Banana 或 SORA)的核心。
  • 第 4 节,Score Matching: 我们研究 score functions 以及如何通过 Score Matching 学习它们。这不仅是 Diffusion Model 的训练算法,而且解锁了 SDE 采样和引导。
  • 第 5 节,引导: 我们学习如何根据 prompt(例如“一只猫的图像”)来条件化我们的样本,以及如何通过 classifier-free guidance 强制遵循这样的 prompt。
  • 第 6 节,latent 空间,神经网络架构: 我们讨论如何构建大规模图像和视频生成器,如 Nano Banana。这包括常见的神经网络架构以及如何在 latent space 中构建事物。我们还调查了最先进的模型。
  • 第 7 节(可选),Discrete Diffusion Model: 我们学习如何将 Diffusion Model 的原理从欧几里得空间转换到离散数据(如语言)。这使得使用 Diffusion Model 的原理构建大型语言模型成为可能。

所需背景

由于本主题的技术性质,我们建议具备一定的基础数学成熟度,特别是对概率论有一定的熟悉。因此,我们在附录 A 中包含了一个关于概率论的简要提醒部分。如果那里有些概念对您来说不熟悉,请不要担心。

生成建模即采样

让我们从思考我们可能遇到的各种数据类型或数据模态开始,以及我们将如何用数值表示它们:

  1. 图像:考虑具有H×WH \times W 像素的图像,其中HH 描述图像的高度,WW 描述宽度,每个像素有三个颜色通道(RGB)。对于每个像素和每个颜色通道,我们给定一个R\mathbb{R} 中的强度值。因此,图像可以用元素zRH×W×3\dap\in\mathbb{R}^{H \times W \times 3} 表示。
  2. 视频:视频只是时间上的一系列图像。如果我们有TT 个时间点或,那么视频将由元素zRT×H×W×3\dap\in\mathbb{R}^{T\times H \times W \times 3} 表示。
  3. 分子结构:一种简单的方法是用矩阵z=(z1,,zN)R3×Nz=(z^1,\dots,z^N)\in\mathbb{R}^{3\times N} 表示分子的结构,其中NN 是分子中的原子数,每个ziR3z^i\in\mathbb{R}^3 描述该原子的位置。当然,还有其他更复杂的方法来表示这样的分子。

在上述所有示例中,我们想要生成的对象在数学上都可以表示为向量(可能在展平之后)。因此,在本文档中,我们将有:

核心思想 1(对象作为向量)

我们将生成的对象视为向量zRdz \in \mathbb{R}^d

上述情况的一个显著例外是文本数据,它通常被语言模型(如ChatGPT)建模为离散对象。虽然连续数据zRdz\in \R^d 是我们的主要关注点,但我们也在第 7 节中研究文本生成。

生成作为采样

让我们定义“生成”某物的含义。例如,假设我们想要生成一张狗的图片。自然,有很多可能的狗图片我们都会满意。特别是,没有单一的“最佳”狗图片。相反,有一系列图片适合程度不同。在机器学习中,通常将这种可能图像的多样性实现为图像空间上的概率分布。我们称这样的分布为数据分布,并将其表示为pdata\pdata。在数学上,可以将pdata\pdata 视为概率密度,即一个函数pdata:RdR0\pdata:\R^d\to\mathbb{R}_{\geq 0},它为每个可能的对象zRdz\in \R^d 分配一个似然pdata(z)0\pdata(z)\geq 0。在狗图像的示例中,该分布将给看起来更像狗的图像zz 更高的似然pdata(z)\pdata(z)。因此,图像/视频/分子匹配的“好坏”——一个相当主观的陈述——被替换为它在数据分布pdata\pdata 下的“可能性”有多大。这样,我们可以将生成任务数学地表达为从(未知的)分布pdata\pdata 中采样:

核心思想 2(生成即采样)

生成对象zz 被建模为从数据分布zpdataz\sim \pdata 中采样。

Generative Model是一种机器学习模型,允许我们从pdata\pdata 生成样本。在机器学习中,我们需要数据来训练模型。在生成建模中,我们通常假设可以访问从pdata\pdata 独立采样的有限数量的示例,这些示例共同作为真实分布的代理。

核心思想 3(dataset)

一个 dataset 由有限数量的样本z1,,zNpdataz_1, \dots, z_N \sim \pdata 组成。

对于图像,我们可以通过从互联网上收集公开可用的图像来构建 dataset。对于视频,我们可能类似地考虑使用 YouTube。对于蛋白质结构,像 RCSB 蛋白质数据银行(PDB)这样的来源提供了数十万个实验解析的结构。随着我们的 dataset 规模变得非常大,它越来越成为底层分布pdata\pdata 的更好表示。

引导/条件生成

在许多情况下,我们希望生成一个以某些数据yy 为条件的对象。例如,我们可能希望生成一张以y=y=“一只狗在覆盖着雪的山坡上奔跑,背景是山”为条件的图像。我们可以将其重新表述为从条件分布中采样:

核心思想 4(引导生成)

引导生成涉及从zpdata(y)z\sim \pdata(\cdot | y) 中采样,其中yy 是一个条件变量。

我们称pdata(y)\pdata(\cdot|y)引导数据分布。引导生成建模任务通常涉及学习以任意的(而非固定的)yy 选择为条件。使用我们之前的例子,我们可能希望以不同的文本 prompt 为条件,例如y=y=“一张猫吹生日蜡烛的照片级真实图像”。因此,我们寻求一个可以以任何这样的yy 选择为条件的单一模型。事实证明,无条件生成的技术很容易推广到条件情况。因此,在前三节中,我们将几乎完全专注于无条件情况(同时记住条件生成是我们最终要实现的目标)。

Generative Model

抽象地说,Generative Model 是一种从zpdataz\sim\pdata 返回样本(或至少近似)的算法。如果pdata\pdata 是狗的图像分布,该算法将返回随机的狗图像。在本课程中,我们将专注于使用 Flow Model 或 Diffusion Model 构建 Generative Model 的具体构造,因为这些代表了当前的最先进技术。然而,重要的是要记住,许多其他 Generative Model 已经被开发出来(也许未来还会发现更多)。

小结 2(采样生成)

我们总结本节的发现:

  1. 在这项工作中,我们主要考虑生成表示为向量zRdz\in\mathbb{R}^d 的对象,如图像、视频和分子结构。
  2. 生成是从概率分布pdata\pdata 中生成样本的任务,在训练期间可以访问样本的 dataset z1,,zNpdataz_1,\dots,z_N\sim \pdata
  3. 引导生成假设我们以标签yy 为条件,并且我们希望在训练期间访问数据对(z1,y),(zN,y)(z_1,y)\dots,(z_N,y) 的情况下从pdata(y)\pdata(\cdot|y) 中采样。
  4. 我们的目标是构建一个 Generative Model,即在训练后返回pdata\pdata 样本的模型。

Flow Model 与 Diffusion Model

在上一节中,我们将生成建模形式化为从数据分布 pdata\pdata 中采样。进一步,我们形式化了我们的目标:构建一个 Generative Model,即一个返回样本 zpdataz\sim \pdata 的算法。在本节中,我们描述如何通过模拟适当构造的微分方程来构建 Generative Model。例如,Flow Matching 和 Diffusion Model 分别涉及模拟Ordinary Differential Equation(ODE)和Stochastic Differential Equation(SDE)。因此,本节的目标是定义并构建这些 Generative Model,因为它们将在本笔记的其余部分中使用。具体来说,我们首先定义 ODE 和 SDE,并讨论它们的模拟。其次,我们描述如何使用深度神经网络参数化 ODE/SDE。这引出了 Flow Model 和 Diffusion Model 的定义,以及从这些模型采样的基本算法。在后面的章节中,我们将探讨如何训练这些模型。

Flow Model

我们从定义Ordinary Differential Equation(ODE)开始。ODE 的解由一条轨迹定义,即一个如下形式的函数

X:[0,1]Rd,tXt,X: [0,1] \to \R^d, \quad t \mapsto X_t,

它将时间 tt 映射到空间 Rd\mathbb{R}^d 中的某个位置。每个 ODE 由一个Vector Field uu 定义,即一个如下形式的函数

u:Rd×[0,1]Rd,(x,t)ut(x),u:\mathbb{R}^d\times [0,1]\to \R^d,\quad (x,t)\mapsto u_t(x),

即对于每个时间 tt 和位置 xx,我们得到一个向量 ut(x)Rdu_t(x)\in\R^d,指定空间中的速度(见 Figure 1)。ODE 对轨迹施加了一个条件:我们想要一条轨迹 XX,它“沿着 Vector Field utu_t 的线”移动,从点 x0x_0 开始。我们可以将这样的轨迹形式化为以下方程的解:

ddtXt=ut(Xt)ODE(1a)\begin{aligned}\frac{\dd}{\dd t}X_{t} &= u_t(X_t) &&\blacktriangleright\,\,\text{ODE}\end{aligned}\tag{1a}
X0=x0initial conditions(1b)\begin{aligned}X_0&= x_0 &&\blacktriangleright\,\,\text{initial conditions}\end{aligned}\tag{1b}

式(1a–1b)要求XtX_t 的导数由utu_t 给出的方向指定。式(1a–1b)要求我们在时间t=0t=0x0x_0 开始。我们现在可以问:如果我们在时间t=0t=0X0=x0X_0 = x_0 开始,那么在时间tt 我们在哪里(XtX_t 是什么)?这个问题由一个称为的函数回答,它是 ODE 的解。

ψ:Rd×[0,1]Rd,(x0,t)ψt(x0)(2a)\begin{aligned}\psi:\R^d\times [0,1]\to& \R^d,\quad (x_0,t)\mapsto \psi_t(x_0)\end{aligned}\tag{2a}
ddtψt(x0)=ut(ψt(x0))flow ODE(2b)\begin{aligned}\frac{\dd}{\dd t}\psi_{t}(x_0) &= u_t(\psi_{t}(x_0)) &&\blacktriangleright\,\, \text{flow ODE}\end{aligned}\tag{2b}
ψ0(x0)=x0flow initial conditions(2c)\begin{aligned}\psi_{0}(x_0) &= x_0 &&\blacktriangleright\,\,\text{flow initial conditions}\end{aligned}\tag{2c}

对于给定的初始条件 X0=x0X_0=x_0,ODE 的轨迹通过 Xt=ψt(X0)X_t = \psi_t(X_0) 恢复。因此,Vector Field、ODE 和流直观上是同一对象的三种描述:Vector Field 定义 ODE,其解是流。与每个方程一样,我们应该问自己关于 ODE 的问题:解是否存在,如果存在,是否唯一?数学中的一个基本结果是“是的!”两者都成立,只要我们对 utu_t 施加弱假设:

定理 3(流的存在性和唯一性)

如果u:Rd×[0,1]Rdu:\R^d\times[0,1]\to\R^d 是连续可微的且具有有界导数,那么式(2a–2c)中的 ODE 具有由流ψt\psi_t 给出的唯一解。在这种情况下,对于所有ttψt\psi_t 是一个微分同胚,即ψt\psi_t 是连续可微的,且具有连续可微的逆ψt1\psi_t^{-1}

请注意,流的存在性和唯一性所需的假设在机器学习中几乎总是满足的,因为我们使用神经网络来参数化 ut(x)u_t(x),而它们总是具有有界导数。因此,第 2.1 节不应成为你的担忧,而应是个好消息:在我们的关注情况下,流存在且是 ODE 的唯一解。 证明可在 [32, 9] 中找到。

Figure 1ψt:RdRd\psi_t:\Real^d\too \Real^d(红色方格网格)由速度场 ut:RdRdu_t :\Real^d\too\Real^d(用蓝色箭头可视化)定义,该速度场规定了其在所有位置的瞬时运动(此处为 d=2d=2)。我们展示了三个不同的时间 tt。可以看出,流是一个“扭曲”空间的微分同胚。图来自 [26]。
例 4(线性 Vector Field)

让我们考虑一个简单的 Vector Field ut(x)u_t(x) 的例子,它是 xx 的简单线性函数,即对于 θ>0\theta>0ut(x)=θxu_t(x)=-\theta x。那么函数

ψt(x0)=exp(θt)x0(3)\psi_t(x_0) = \exp\left(-\theta t\right)x_0 \tag{3}

定义了一个流 ψ\psi,它求解式(2a–2c)中的 ODE。你可以通过检查 ψ0(x0)=x0\psi_0(x_0)=x_0 并计算来自己验证这一点。

ddtψt(x0)=式(3)ddt(exp(θt)x0)=(i)θexp(θt)x0=式(3)θψt(x0)=ut(ψt(x0)),\frac{\dd}{\dd t}\psi_t(x_0) \overset{\text{式(3)}}{=}\frac{\dd}{\dd t}\left(\exp\left(-\theta t\right)x_0\right) \overset{(i)}{=}-\theta\exp\left(-\theta t\right)x_0\overset{\text{式(3)}}{=}-\theta\psi_t(x_0)=u_t(\psi_t(x_0)),

其中在 (i) 中我们使用了链式法则。在 Figure 3 中,我们可视化了这种形式的流,它指数级地收敛到 00

模拟一个 ODE

一般来说,如果 utu_t 不像前一个例子中那样简单,就不可能显式计算流 ψt\psi_t。在这些情况下,人们使用数值方法来模拟 ODE。幸运的是,这是数值分析中一个经典且研究充分的课题,存在大量强大的方法 [21]。最简单和最直观的方法之一是 Euler method。在 Euler method 中,我们初始化 X0=x0X_0=x_0 并通过以下方式更新。

Xt+h=Xt+hut(Xt)(t=0,h,2h,3h,,1h)(4)X_{t+h} = X_t + h u_t(X_t)\quad (t=0,h,2h,3h,\dots,1-h) \tag{4}

其中 h=n1>0h=n^{-1}>0步长nNn \in \Nat 是模拟步数。对于本课程,Euler method 将足够好。为了让你领略更复杂的方法,让我们考虑通过更新规则定义的 Heun's method

Xt+h=Xt+hut(Xt)initial guess of new state (same as Euler step)Xt+h=Xt+h2(ut(Xt)+ut+h(Xt+h))update with average u at current and guessed state\begin{aligned}X_{t+h}'&=X_t+hu_t(X_t)\quad &&\blacktriangleright\,\, \text{initial guess of new state (same as Euler step)}\\ X_{t+h} &= X_{t} + \frac{h}{2}(u_t(X_t)+u_{t+h}(X_{t+h}'))\quad &&\blacktriangleright\,\,\text{update with average }u\text{ at current and guessed state}\end{aligned}

直观地说,Heun's method 如下:它先对下一步可能是什么做一个初步猜测 Xt+hX_{t+h}',但通过更新的猜测来修正最初采取的方向。

Flow Model

我们现在可以通过使 Vector Field 成为neural-network Vector Field utθu_t^\theta 来构建一个基于 ODE 的 Generative Model。目前,我们仅指 utθu_t^\theta 是一个参数化函数 utθ:Rd×[0,1]Rdu_t^\theta:\mathbb{R}^d\times [0,1]\to\mathbb{R}^d,参数为 θ\theta。稍后,我们将讨论神经网络架构的特定选择。记住我们的目标是从分布 pdata\pdata 生成样本 zpdataz\sim \pdata。特别是,这些样本必须是随机的。但请注意,ODE 本身不是随机的,而是完全确定性的。为了注入一些随机性,我们只需使初始条件 X0X_0 随机。具体来说,我们选择一个初始分布 pinit\pinit。在大多数情况下,我们将 pinit=N(0,Id)\pinit=\mathcal{N}(0,I_d) 设置为简单的 standard Gaussian distribution。最重要的是,无论你选择什么分布,它必须是在推理时容易采样的。一个Flow Model然后由 ODE 描述。

X0pinitrandom initializationddtXt=utθ(Xt)ODE\begin{aligned}X_0 &\sim \pinit &&\blacktriangleright\,\,\text{random initialization} \\ \frac{\dd}{\dd t}X_t &= u_t^\theta(X_t) &&\blacktriangleright\,\,\text{ODE}\end{aligned}

我们的目标是使轨迹的终点 X1X_1 具有分布 pdata\pdata,即

X1pdataψ1θ(X0)pdataX_1 \sim \pdata \quad \Leftrightarrow \quad \psi_1^\theta(X_0)\sim \pdata

其中 ψtθ\psi_t^\theta 描述了由 utθu_t^\theta 诱导的流。但请注意:尽管它被称为Flow Model神经网络参数化的是 Vector Field,而不是流。为了计算流,我们需要模拟 ODE。在 Algorithm 1 中,我们总结了如何从 Flow Model 中采样的过程。

Algorithm 1
算法 1 · 使用 Euler method 从 Flow Model 采样
输入: neural-network Vector Field utθu_t^\theta,步数 nn

- 设置 t=0t=0

- 设置步长 h=1nh=\frac{1}{n}

- 抽取样本 X0pinitX_0\sim \pinit

- 循环: i=1,,ni=1,\dots,n

- Xt+h=Xt+hutθ(Xt)X_{t+h} = X_{t} + h u_t^\theta(X_t)

- 更新 tt+ht\leftarrow t+h

- 结束循环

- 返回: X1X_1

Diffusion Model

Stochastic Differential Equation (SDE) 将 ODE 的确定性轨迹扩展为随机轨迹。随机轨迹通常称为随机过程 (Xt)0t1(X_t)_{0\leq t\leq 1},由下式给出。

Xt is a random variable for every 0t1X:[0,1]Rd,tXt is a random trajectory for every draw of X\begin{aligned}X_t \text{ is a random variable for every } 0\leq t\leq 1\\ X: [0,1] \to \R^d, \quad t \mapsto X_t\text{ is a random trajectory for every draw of }X\end{aligned}

特别是,当我们两次模拟同一个随机过程时,可能会得到不同的结果,因为其动力学被设计为随机的。

Brownian motion

SDE 是通过 Brownian motion 构建的——这是一个源自物理扩散过程研究的基本随机过程。你可以将 Brownian motion 想象为一种连续随机游走。

Figure 2使用 式(5) 模拟的 Brownian motion WtW_t 在维度 d=1d=1 中的示例轨迹。

让我们来定义它:Brownian motion W=(Wt)0t1W = (W_t)_{0\leq t\leq 1} 是一个随机过程,满足 W0=0W_0=0,轨迹 tWtt\mapsto W_t 是连续的,并且以下两个条件成立:

  1. 正态增量: WtWsN(0,(ts)Id)W_{t}-W_{s}\sim \mathcal{N}(0,(t-s)I_d) 对所有 0s<t0\leq s<t 成立,即增量服从 Gaussian distribution,其方差随时间线性增加(IdI_d 是单位矩阵)。
  2. 独立增量: 对于任意 0t0<t1<<tn=10\leq t_0<t_1<\dots <t_n=1,增量 Wt1Wt0,,WtnWtn1W_{t_1}-W_{t_0},\dots,W_{t_n}-W_{t_{n-1}} 是独立的随机变量。

Brownian motion 也被称为 Wiener 过程,这就是我们用“WW”表示它的原因。〔脚注〕 我们可以通过设置 W0=0W_0=0 并以步长 h>0h>0 更新来近似模拟 Brownian motion。

Wt+h=Wt+hϵt,ϵtN(0,Id)(t=0,h,2h,,1h)(5)\begin{aligned}W_{t+h} =& W_{t} + \sqrt{h}\epsilon_t,\quad \epsilon_t\sim\mathcal{N}(0,I_d)\quad (t=0,h,2h,\dots,1-h)\end{aligned}\tag{5}

在图 2 中,我们绘制了 Brownian motion 的几个示例轨迹。Brownian motion 在随机过程研究中的核心地位,正如 Gaussian distribution 在概率分布研究中的核心地位。从金融到统计物理再到流行病学,Brownian motion 的研究在机器学习之外有着广泛的应用。例如,在金融领域,Brownian motion 被用于对复杂金融工具的价格进行建模。同样,仅作为数学构造,Brownian motion 也令人着迷:例如,虽然 Brownian motion 的路径是连续的(因此你可以不抬笔地画出它),但它们无限长(因此你永远不会停止画)。

从 ODE 到 SDE

SDE 的思想是通过添加由 Brownian motion 驱动的随机动力学,来扩展 ODE 的确定性动力学。由于一切都是随机的,我们可能不再像式(1a–1b)那样求导。因此,我们需要找到一种不使用导数的 ODE 等价表述。为此,让我们将 ODE 的轨迹 (Xt)0t1(X_t)_{0\leq t\leq 1} 重写如下:

Figure 3Ornstein-Uhlenbeck 过程(式(8))在维度 d=1d=1 中的图示,其中 θ=0.25\theta=0.25σ\sigma 取不同值(从左到右递增)。对于 σ=0\sigma=0,我们恢复了一个流(光滑、确定性的轨迹),随着 tt \to \infty 收敛到原点。对于 σ>0\sigma>0,我们有随机路径,随着 tt\to\infty 收敛到 Gaussian N(0,σ22θ)\mathcal{N}(0,\frac{\sigma^2}{2\theta})
ddtXt=ut(Xt)expression via derivatives(i)1h(Xt+hXt)=ut(Xt)+Rt(h)Xt+h=Xt+hut(Xt)+hRt(h)expression via infinitesimal updates\begin{aligned}\frac{\dd}{\dd t} X_t &= u_t(X_t) \quad &&\blacktriangleright\,\,\text{expression via derivatives}\\ \overset{(i)}{\Leftrightarrow} \quad \frac{1}{h}\left(X_{t+h}-X_{t}\right)&=u_t(X_t) + R_t(h)&&\\ \Leftrightarrow \quad X_{t+h} &= X_{t}+hu_t(X_t) + hR_t(h)\quad &&\blacktriangleright\,\,\text{expression via infinitesimal updates}\end{aligned}

其中 Rt(h)R_t(h) 描述了一个对于小的 hh 可忽略的函数,即满足 limh0Rt(h)=0\lim\limits_{h\to 0}R_t(h)=0,而在 (i)(i) 中我们仅使用了导数的定义。上述推导只是重述了我们已知的内容:一个 ODE 的轨迹 (Xt)0t1(X_t)_{0 \le t \le 1} 在每个 timestep 沿方向 ut(Xt)u_t(X_t) 迈出一小步。我们现在可以修改最后一个方程使其具有随机性:一个 SDE 的轨迹 (Xt)0t1(X_t)_{0 \le t \le 1} 在每个 timestep 沿方向 ut(Xt)u_t(X_t) 迈出一小步,加上来自 Brownian motion 的某些贡献:

Xt+h=Xt+hut(Xt)deterministic+σt(Wt+hWt)stochastic+hRt(h)error term(6)X_{t+h} = X_{t}+\underbrace{hu_t(X_t)}_{\text{deterministic}} + \sigma_t\underbrace{(W_{t+h}-W_{t})}_{\text{stochastic}}+\underbrace{hR_t(h)}_{\text{error term}} \tag{6}

其中 σt0\sigma_t\geq 0 描述扩散系数Rt(h)R_t(h) 描述一个随机误差项,使得标准差 E[Rt(h)2]1/20\mathbb{E}[\|R_t(h)\|^2]^{1/2}\to 0h0h\to 0 时趋于零。上述描述了一个Stochastic Differential Equation (SDE)。通常用以下符号表示:

dXt=ut(Xt)dt+σtdWtSDE(7a)\begin{aligned}\dd X_t &= u_t(X_t)\dd t + \sigma_t\dd W_t &&\blacktriangleright\,\,\text{SDE}\end{aligned}\tag{7a}
X0=x0initial condition(7b)\begin{aligned}X_0 &= x_0 &&\blacktriangleright\,\,\text{initial condition}\end{aligned}\tag{7b}

然而,请始终记住,上述“dXt\dd X_t”记号只是式(6)的一种非正式记号。不幸的是,SDE 不再具有流映射 ϕt\phi_t。这是因为值 XtX_t 不再完全由 X0pinitX_0\sim \pinit 决定,因为演化本身是随机的。尽管如此,与 ODE 类似,我们有:

定理 5(SDE 解的存在性和唯一性)

如果u:Rd×[0,1]Rdu:\R^d\times[0,1]\to\R^d 连续可微且有界导数,且σt\sigma_t 连续,则式(7a–7b)中的 SDE 具有由满足式(6)的唯一随机过程(Xt)0t1(X_t)_{0\leq t\leq 1} 给出的解。

如果这是一门随机微积分课程,我们会花几节课来证明这个定理,并以完全严谨的数学方式构造 SDE,即从基本原理构造 Brownian motion,并通过随机积分构造过程 XtX_{t}。由于本课程侧重于机器学习,我们参考 [29] 以获得更技术性的处理。最后,注意每个 ODE 也是一个 SDE——只需令扩散系数 σt=0\sigma_t=0 为零。因此,在本课程剩余部分,当我们谈论 SDE 时,我们将 ODE 视为一种特殊情况

例 6(奥恩斯坦-乌伦贝克过程)

让我们考虑一个常数扩散系数 σt=σ0\sigma_t=\sigma\geq 0 和一个常数线性漂移 ut(x)=θxu_t(x)=-\theta x 对于 θ>0\theta>0,得到 SDE

dXt=θXtdt+σdWt.(8)\dd X_t = -\theta X_t\dd t + \sigma \dd W_t. \tag{8}

上述 SDE 的一个解 (Xt)0t1(X_t)_{0 \le t \le 1} 被称为Ornstein-Uhlenbeck (OU) 过程。我们在图 3 中将其可视化。Vector Field θx-\theta x 将过程推回其中心 00(因为漂移总是指向当前位置的相反方向),而扩散系数 σ\sigma 总是添加更多噪声。如果我们模拟 tt\to \infty,该过程收敛到 Gaussian distribution N(0,σ2/(2θ))\mathcal{N}(0,\sigma^2/(2\theta))。注意对于 σ=0\sigma=0,我们有一个线性 Vector Field 的流,我们在式(3)中已经研究过。

模拟一个 SDE

如果你到目前为止对 SDE 的抽象定义感到困惑,不用担心。一个更直观的思考 SDE 的方式是回答这个问题:我们如何模拟一个 SDE?最简单的此类方案被称为Euler-Maruyama 方法,它本质上相当于 SDE 中的 Euler method 之于 ODE。使用 Euler-Maruyama 方法,我们初始化 X0=x0X_0=x_0 并迭代更新:

Xt+h=Xt+hut(Xt)+hσtϵt,ϵtN(0,Id)(9)X_{t+h} = X_{t}+hu_t(X_t) + \sqrt{h}\sigma_t\epsilon_t,\quad \quad \epsilon_t \sim \mathcal{N}(0,I_d) \tag{9}

其中 h=n1>0h=n^{-1}>0nNn \in \Nat 的步长超参数。换句话说,要使用 Euler-Maruyama 方法进行模拟,我们沿 ut(Xt)u_t(X_t) 方向迈出一小步,并添加一些按 hσt\sqrt{h}\sigma_t 缩放的 Gaussian noise。在本课程中模拟 SDE 时(例如在配套实验室中),我们通常坚持使用 Euler-Maruyama 方法。

Diffusion Model

我们现在可以通过与 ODE 相同的方式,利用 SDE 构建 Generative Model。记住我们的目标是将简单分布pinit\pinit 转换为复杂分布pdata\pdata。与 ODE 类似,用X0pinitX_0\sim \pinit 随机初始化 SDE 的模拟是这种转换的自然选择。为了参数化这个 SDE,我们可以简单地通过神经网络utθu_t^\theta 参数化其核心成分——Vector Fieldutu_t。因此,Diffusion Model由下式给出:

X0pinitrandom initializationdXt=utθ(Xt)dt+σtdWtSDE\begin{aligned}X_0 &\sim \pinit &&\blacktriangleright\,\,\text{random initialization}\\ \dd X_t &= u_t^\theta(X_t)\dd t + \sigma_t \dd W_t &&\blacktriangleright\,\,\text{SDE}\end{aligned}
Algorithm 2
算法 2 · 从 Diffusion Model 采样(Euler-Maruyama 方法)
输入: 神经网络utθu_t^\theta,步数nn,扩散系数σt\sigma_t

- 设置t=0t=0

- 设置步长h=1nh=\frac{1}{n}

- 抽取样本X0pinitX_0\sim \pinit

- 循环: i=1,,ni=1,\dots,n

- 抽取样本ϵN(0,Id)\epsilon\sim \mathcal{N}(0,I_d)

- Xt+h=Xt+hutθ(Xt)+σthϵX_{t+h} = X_{t} + h u_t^\theta(X_t)+\sigma_t\sqrt{h}\epsilon

- 更新tt+ht\leftarrow t+h

- 结束循环

- 返回: X1X_1

在算法 2 中,我们描述了使用 Euler-Maruyama 方法从 Diffusion Model 采样的过程。我们将本节的结果总结如下。

小结 7(SDEGenerative Model)

在本文档中,Diffusion Model由一个参数为θ\theta 的神经网络utθu_t^\theta 组成,该网络参数化一个 Vector Field,并具有固定的扩散系数σt\sigma_t

Neural network: uθ:Rd×[0,1]Rd,(x,t)utθ(x) with parameters θFixed: σt:[0,1][0,),tσt\begin{aligned}\textbf{\sffamily Neural network: }&u^\theta:\R^d\times [0,1]\to \R^d,\,\, (x,t)\mapsto u_t^\theta(x)\text{ with parameters }\theta\\ \textbf{\sffamily Fixed: }&\sigma_t:[0,1]\to [0,\infty),\,\, t\mapsto \sigma_t\end{aligned}

为了从我们的 SDE 模型中获得样本(即生成对象),过程如下:

Initialization:X0pinitInitialize with simple distribution, e.g. a GaussianSimulation:dXt=utθ(Xt)dt+σtdWtSimulate SDE from 0 to 1Goal:X1pdataGoal is to make X1 have distribution pdata\begin{aligned}\textbf{\sffamily Initialization:}\quad X_0&\sim\pinit \quad &&\blacktriangleright\,\,\text{Initialize with simple distribution, e.g. a Gaussian}\\ \textbf{\sffamily Simulation:}\quad \dd X_t &= u_t^\theta(X_t)\dd t + \sigma_t\dd W_t\quad &&\blacktriangleright\,\,\text{Simulate SDE from 0 to 1}\\ \textbf{\sffamily Goal:}\quad X_1 &\sim \pdata \quad &&\blacktriangleright\,\,\text{Goal is to make }X_1\text{ have distribution }\pdata\end{aligned}

具有σt=0\sigma_t=0 的 Diffusion Model 是一个Flow Model

Flow Matching

在上一节中,我们将 Flow Model 和 Diffusion Model 构建为由 neural-network Vector Fieldutθu_t^\theta 参数化的 Generative Model。然而,我们尚未讨论如何训练它们,即如何优化参数θ\theta,使得 Generative Model 返回合理的结果,例如好看的图像或令人兴奋的视频。接下来,我们讨论Flow Matching [25, 1, 27],这是一种训练utθu_t^\theta 的算法,它简单、可扩展,并且代表了当前的最先进水平。

在本节中,我们仅限于 Flow Model,即我们有一个神经网络utθu_t^\theta,并通过模拟 ODE 从 Generative Model 中获得样本:

X0pinit,dXt=utθ(Xt)dt(Flow model)(10)\begin{aligned}X_0\sim&\pinit,\quad \dd X_t = u_t^\theta(X_t)\dd t& \text{(Flow model)}\end{aligned}\tag{10}

并使用端点X1X_1(来自t=1t=1)作为样本。正如我们所讨论的,我们的目标是X1X_1 服从数据分布pdata\pdata,即X1pdataX_1\sim \pdata。因此,“如何训练”神经网络的问题实际上是以下问题:我们如何优化θ\theta,使得模拟式(10)中的 Flow Model 产生来自数据分布X1pdataX_1\sim \pdata 的样本?

Figure 4通过 Gaussian Conditional Probability Path 从噪声到数据的逐步插值,用于一组图像。请注意,每个图像是维度d=32×32d=32\times 32 的数据点,因此我们绘制的是 Probability Path 的单个样本,而在 Figure 5 中,我们将分布绘制为二维直方图。

Conditional and Marginal Probability Paths

Flow Matching 的第一步是指定一个Probability Path。直观上,Probability Path 指定了噪声pinit\pinit 和数据pdata\pdata 之间的逐渐插值(见图 4)。但为什么我们需要这个呢?记住我们期望的 ODE 轨迹满足X0pinitX_0\sim \pinit 对于t=0t=0X1pdataX_1\sim \pdata 对于t=1t=1。但介于开始和结束之间的时间0<t<10<t<1 呢?事实证明,我们有一些自由来选择中间应该发生什么,这就是在 Probability Path 中数学形式化的内容。

在下文中,对于数据点zRdz\in\mathbb{R}^d,我们用δz\delta_{z} 表示狄拉克δ“分布”。这是能想象到的最简单的分布:从δz\delta_{z} 采样总是返回zz(即它是确定性的)。一个conditional (interpolation) probability pathpt(xz)p_t(x|z)Rd\mathbb{R}^d 上的一组分布,使得:

p0(z)=pinit,p1(z)=δz for all zRd.(11)p_0(\cdot|z)=\pinit, \quad p_1(\cdot|z)=\delta_{z}\quad \text{ for all }z\in\R^d. \tag{11}

换句话说,Conditional Probability Path 逐渐将初始分布pinit\pinit 转换为单个数据点(例如见图 4)。你可以将 Probability Path 视为分布空间中的轨迹。

每个 Conditional Probability Pathpt(xz)p_t(x|z) 都诱导一个Marginal Probability Pathpt(x)p_t(x),定义为通过首先从数据分布中采样数据点zpdataz\sim \pdata,然后从pt(z)p_t(\cdot|z) 中采样而获得的分布:

zpdata,xpt(z)xptsampling from marginal path(12)\begin{aligned}z&\sim\pdata, \quad x\sim p_t(\cdot|z)\quad \Rightarrow x\sim p_t &&\blacktriangleright\,\,\text{sampling from marginal path}\end{aligned}\tag{12}
pt(x)=pt(xz)pdata(z)dzdensity of marginal path(13)\begin{aligned}p_t(x) &= \int p_t(x|\dap) \pdata (z) \dd z &&\blacktriangleright\,\,\text{density of marginal path}\end{aligned}\tag{13}

注意,我们知道如何从ptp_t 采样,但我们不知道密度值pt(x)p_t(x),因为积分是难以处理的(即我们实际上可以计算式(12)但不能计算式(13))。请自行验证,由于式(11)中对pt(z)p_t(\cdot|z) 的条件,Marginal Probability Pathptp_tpinit\pinitpdata\pdata 之间插值:

p0=pinitandp1=pdata.noise-data interpolation(14)p_0 =\pinit\quad\text{and}\quad p_1=\pdata.\quad\quad\quad\quad \blacktriangleright\,\,\text{noise-data interpolation} \tag{14}

迄今为止最重要的 Probability Path 示例是 Gaussian Probability Path——因此,我们强烈建议仔细阅读下一个示例。

Figure 5条件(顶部)和边缘(底部)Probability Path 的图示。这里,我们绘制了一个 Gaussian Probability Path,其中αt=t,βt=1t\alpha_t=t,\beta_t=1-t。Conditional Probability Path 在单个数据点zz 的 Gaussianpinit=N(0,Id)\pinit=\mathcal{N}(0,I_d)pdata=δz\pdata=\delta_{z} 之间插值。Marginal Probability Path 在 Gaussian 和数据分布pdata\pdata 之间插值(这里,pdata\pdata 是维度d=2d=2 中的玩具分布,由棋盘图案表示)。
例 8(Gaussian Conditional Probability Path)

一个特别流行的 Probability Path 是Gaussian Probability Path。这是大多数最先进模型使用的 Probability Path。设αt,βt\alpha_t,\beta_tnoise schedules:两个连续可微、单调的函数,满足α0=β1=0\alpha_0=\beta_1=0α1=β0=1\alpha_1=\beta_0=1。然后我们定义 Conditional Probability Path

pt(z)=N(αtz,βt2Id)Gaussian conditional path(15)\begin{aligned}p_t(\cdot|\dap) &= \mathcal{N}(\alpha_t \dap,\beta_t^2 I_d) & \blacktriangleright \,\, \text{Gaussian conditional path}\end{aligned}\tag{15}

根据我们对αt\alpha_tβt\beta_t 施加的条件,这满足

p0(z)=N(α0z,β02Id)=N(0,Id),andp1(z)=N(α1z,β12Id)=δz,\begin{aligned}p_0(\cdot|\dap) &= \mathcal{N}(\alpha_0 \dap,\beta_0^2 I_d) = \mathcal{N}(0,I_d),\quad \text{and}\quad p_1(\cdot|\dap) = \mathcal{N}(\alpha_1 \dap,\beta_1^2 I_d) = \delta_{\dap},\end{aligned}

其中我们使用了方差为零且均值为z\dap 的正态分布就是δz\delta_{\dap} 这一事实。因此,这种pt(xz)p_t(x|\dap) 的选择满足式(11)对于pinit=N(0,Id)\pinit=\mathcal{N}(0,I_d),因此是一个有效的条件插值路径。

在图 4 中,我们展示了其在图像上的应用。我们可以将从边缘路径ptp_t 采样表示为:

zpdata,ϵpinit=N(0,Id)x=αtz+βtϵptsampling from marginal Gaussian path(16)\begin{aligned}z\sim&\pdata,\,\epsilon\sim\pinit = \mathcal{N}(0,I_d) \,\Rightarrow\, x=\alpha_tz+\beta_t \epsilon\sim p_t \quad &\blacktriangleright \,\,\text{sampling from marginal Gaussian path}\end{aligned}\tag{16}

直观上,上述过程在时间 tt 之前添加更多噪声,直到时间 t=0t=0,此时只剩下噪声。在图 5 中,我们绘制了这样一个插值路径的示例。

Conditional and Marginal Vector Fields

Probability Path (pt)0t1(p_t)_{0\leq t\leq 1} 指定了轨迹上的点 XtX_t 应该具有的分布 XtptX_t\sim p_t。此时,这只是我们“希望”的情况。但我们如何找到一个 Vector Field,使得轨迹 XtX_t 遵循该 Probability Path?Flow Matching 显式地构造了这样一个 Vector Field——“Marginal Vector Field”——我们将在本节中解释。

对于每个数据点 zRdz\in \mathbb{R}^d,让 uttarget(z)\uref_t(\cdot|z) 表示一个Conditional Vector Field。这可以是任何 Vector Field,使得相应的 ODE 产生 Conditional Probability Path pt(z)p_t(\cdot|z),即满足

X0pinit,ddtXt=uttarget(Xtz)Xtpt(z)(0t1).(17)\begin{aligned}X_0&\sim\pinit,\quad \frac{\dd}{\dd t}X_t =\uref_t(X_t|z)\quad \Rightarrow \quad X_t\sim p_t(\cdot|z)\quad (0\leq t\leq 1).\end{aligned}\tag{17}

我们通常可以通过手工分析(即我们自己做一些代数运算)找到 Conditional Vector Field uttarget(z)\uref_t(\cdot|z)。我们通过为 Gaussian Probability Path 的示例推导 Conditional Vector Field ut(xz)u_t(x|z) 来说明这一点,见示例 第 3.2 节。

乍一看,Conditional Vector Field 似乎无用,因为 ODE X1X_1 的所有端点都会坍缩到 X1=zX_1=z,即我们只是在重新生成已知的数据点 zz。然而,Conditional Vector Field 是生成来自 pdata\pdata 的实际样本的 Vector Field 的构建块:

定理 9(边缘化技巧)

uttarget(xz)\uref_t(x|z) 是一个 Conditional Vector Field(式(17))。那么Marginal Vector Field uttarget(x)\uref_t(x) 定义为

uttarget(x)=uttarget(xz)pt(xz)pdata(z)pt(x)dz,(18)\uref_t(x) = \int \uref_t(x|z)\frac{p_t(x|z)\pdata(z)}{p_t(x)}\dd z, \tag{18}

遵循 Marginal Probability Path,即

X0pinit,ddtXt=uttarget(Xt)Xtpt(0t1).(19)\begin{aligned}X_0&\sim\pinit,\quad \frac{\dd}{\dd t}X_t =\uref_t(X_t)\quad \Rightarrow \quad X_t\sim p_t\quad (0\leq t\leq 1).\end{aligned}\tag{19}

特别是,对于这个 ODE,X1pdataX_1\sim \pdata,因此我们可以说“uttarget\uref_t 将噪声 pinit\pinit 转换为数据 pdata\pdata

例 10
Figure 6第 3.2 节的插图。用 ODE 模拟 Probability Path。数据分布 pdata\pdata 在蓝色背景中。Gaussian pinit\pinit 在红色背景中。顶行:Conditional Probability Path。左:来自条件路径 pt(z)p_t(\cdot|z) 的真实样本。中:随时间变化的 ODE 样本。右:通过用 式(20) 中的 uttarget(xz)\uref_t(x|z) 模拟 ODE 得到的轨迹。底行:模拟 Marginal Probability Path。左:来自 ptp_t 的真实样本。中:随时间变化的 ODE 样本。右:通过用 Marginal Vector Field utflow(x)\uflow_t(x) 模拟 ODE 得到的轨迹。可以看出,Conditional Vector Field 遵循 Conditional Probability Path,Marginal Vector Field 遵循 Marginal Probability Path。

[Gaussian Probability Path 的目标 ODE]

如前所述,设 pt(z)=N(αtz,βt2Id)p_t(\cdot|\dap) = \mathcal{N}(\alpha_t \dap,\beta_t^2 I_d) 为 noise schedules αt,βt\alpha_t,\beta_t(见式(15))。令 α˙t=tαt\dot{\alpha}_t=\partial_t\alpha_tβ˙t=tβt\dot{\beta}_t=\partial_t\beta_t 分别表示 αt\alpha_tβt\beta_t 对时间的导数。这里,我们想要证明由下式给出的条件 GaussianVector Field

uttarget(xz)=(α˙tβ˙tβtαt)z+β˙tβtx(20)\uref_t(x|z) = \left(\dot{\alpha}_t-\frac{\dot{\beta}_t}{\beta_t}\alpha_t\right)z+\frac{\dot{\beta}_t}{\beta_t}x \tag{20}

是一个有效的 Conditional Vector Field 模型,符合第 3.2 节的定义:其 ODE 轨迹 XtX_t 满足 Xtpt(z)=N(αtz,βt2Id)X_t\sim p_t(\cdot|z)=\mathcal{N}(\alpha_t z,\beta_t^2I_d),如果 X0N(0,Id)X_0\sim \mathcal{N}(0,I_d)。在图 6 中,我们通过将 Conditional Probability Path(真实值)的样本与模拟的该流的 ODE 轨迹的样本进行视觉对比来确认这一点。如你所见,分布匹配。我们现在将证明这一点。

证明

首先,我们通过定义来构造一个条件 Flow Model ψttarget(xz)\psiref_t(x|z)

ψttarget(xz)=αtz+βtx.(21)\psiref_t(x|z) = \alpha_t z + \beta_t x. \tag{21}

如果 XtX_tψttarget(z)\psiref_t(\cdot|z) 的 ODE 轨迹,且 X0pinit=N(0,Id)X_0\sim\pinit = \mathcal{N}(0,I_d),那么根据定义

Xt=ψttarget(X0z)=αtz+βtX0N(αtz,β2Id)=pt(z).X_t = \psiref_t(X_0|z) = \alpha_t z + \beta_t X_0 \sim \mathcal{N}(\alpha_t z,\beta^2 I_d) = p_t(\cdot|z).

我们得出结论,轨迹的分布与 Conditional Probability Path 一致(即式(17)成立)。接下来需要从ψttarget(xz)\psiref_t(x|z) 中提取 Vector Fielduttarget(xz)\uref_t(x|z)。根据流的定义(式(2a–2c)),有

ddtψttarget(xz)=uttarget(ψttarget(xz)z) for all x,zRd(i)α˙tz+β˙tx=uttarget(αtz+βtxz) for all x,zRd(ii)α˙tz+β˙t(xαtzβt)=uttarget(xz) for all x,zRd(iii)(α˙tβ˙tβtαt)z+β˙tβtx=uttarget(xz) for all x,zRd\begin{aligned}\frac{\dd}{\dd t}\psiref_{t}(x|z) &= \uref_t(\psiref_t(x|z)|z)\quad \text{ for all }x,z\in\mathbb{R}^d\\ \overset{(i)}{\Leftrightarrow} \quad \dot{\alpha}_tz+\dot{\beta}_tx &= \uref_t(\alpha_tz+\beta_t x|z)\quad \text{ for all }x,z\in\mathbb{R}^d\\ \overset{(ii)}{\Leftrightarrow} \quad \dot{\alpha}_tz+\dot{\beta}_t\left(\frac{x-\alpha_tz}{\beta_t}\right)&= \uref_t(x|z)\quad \text{ for all }x,z\in\mathbb{R}^d\\ \overset{(iii)}{\Leftrightarrow} \quad\left(\dot{\alpha}_t-\frac{\dot{\beta}_t}{\beta_t}\alpha_t\right)z+\frac{\dot{\beta}_t}{\beta_t}x &= \uref_t(x|z)\quad \text{ for all }x,z\in\mathbb{R}^d\end{aligned}

其中,在 (i)(i) 中我们使用了 ψttarget(xz)\psiref_t(x|z) 的定义(式(21)),在 (ii)(ii) 中我们重新参数化了 x(xαtz)/βtx\rightarrow (x-\alpha_t z)/\beta_t,在 (iii)(iii) 中我们仅进行了一些代数运算。注意,最后一个等式正是我们在式(20)中定义的条件 GaussianVector Field。这证明了该命题。〔脚注〕

参见图 6 以了解第 3.2 节的说明。让我们对 Marginal Vector Field 获得一些直觉。统计学中的Bayes' rule表明,以下项描述了一个后验分布

pt(xz)pdata(z)pt(x)="posterior over data points z given noisy data x"\frac{p_t(x|z)\pdata(z)}{p_t(x)} = \text{"posterior over data points }z\text{ given noisy data }x\text{"}

其中 pdata(z)\pdata(z) 是先验分布。Marginal Vector Field 则简单地是一个平均:对于每个可能的数据点 zz,它取速度 ut(xz)u_t(x|z)——即能将我们带到 zz 的方向——然后根据我们对 xx 来自 zz 的相信程度来加权该速度。对所有数据点取平均,我们得到 Marginal Vector Field。

本节的其余部分将使这一直觉严谨化,并证明第 3.2 节。作为主要的数学工具,我们将使用连续性方程,这是数学和物理学中的一个基本方程。定义散度算子 div\divv

div(vt)(x)=i=1dxivti(x)(22)\begin{aligned}\divv(v_t)(x)=&\sum\limits_{i=1}^{d}\frac{\partial}{\partial x_i}v_t^i(x)\end{aligned}\tag{22}

其中 vtiv_t^ivtv_t 的第 ii 个坐标。

定理 11(连续性方程)

让我们考虑一个 Vector Field 为 uttarget\uref_t 且满足 X0pinit=p0X_0\sim\pinit=p_0 的 Flow Model。那么,对于所有 0t10\leq t\leq 1,有 XtptX_t\sim p_t 当且仅当

tpt(x)=div(ptuttarget)(x) for all xRd,0t1,(23)\partial_t p_t(x)=-\divv (p_t\uref_t)(x)\quad \text{ for all }x\in\R^d, 0\leq t\leq 1, \tag{23}

其中 tpt(x)=ddtpt(x)\partial_tp_t(x)= \frac{\dd}{\dd t}p_t(x) 表示 pt(x)p_t(x) 的时间导数。式(23) 被称为连续性方程

对于有数学兴趣的读者,我们在 第 B 节 中给出了连续性方程的自包含证明。在继续之前,让我们尝试直观地理解连续性方程。左边 tpt(x)\partial_tp_t(x) 描述了在 xx 处的概率 pt(x)p_t(x) 随时间的变化量。直观上,这个变化应该对应于概率质量的净流入。对于 Flow Model,粒子 XtX_t 沿着 Vector Field uttarget\uref_t 运动。你可能从物理学中记得,散度衡量的是 Vector Field 的某种净流出。因此,负散度衡量净流入。将其乘以当前位于 xx 的总概率质量,我们得到净 div(ptut)-\divv (p_tu_t) 衡量概率质量的总流入。由于概率质量是守恒的(始终积分为 1),方程的左边和右边应该相同!我们现在继续证明第 3.2 节中的边缘化技巧。

证明(第 3.2 节的证明)

根据第 3.2 节,我们必须证明 Marginal Vector Field uttarget\uref_t(如式(18)中所定义)满足连续性方程。我们可以通过直接计算来做到这一点:

tpt(x)=(i)tpt(xz)pdata(z)dz=tpt(xz)pdata(z)dz=(ii)div(pt(z)uttarget(z))(x)pdata(z)dz=(iii)div(pt(xz)uttarget(xz)pdata(z)dz)=(iv)div(pt(x)uttarget(xz)pt(xz)pdata(z)pt(x)dz)(x)=(v)div(ptuttarget)(x),\begin{aligned}\partial_t p_t(x) \overset{(i)} {=} \partial_t\int p_t(x|\dap) \pdata (z) \dd z &= \int \partial_t p_t(x|\dap) \pdata (z) \dd z\\ &\overset{(ii)}{=} \int -\divv (p_t(\cdot|z)\uref_t(\cdot|z))(x) \pdata (z) \dd z\\ &\overset{(iii)}{=} -\divv \left(\int p_t(x|z) \uref_t(x|z)\pdata(z) \dd z\right)\\ &\overset{(iv)}{=} -\divv \left(p_t(x)\int \uref_t(x|z) \frac{p_t(x|z)\pdata(z)}{p_t(x)}\dd z\right)(x)\\ &\overset{(v)}{=} -\divv \left(p_t\uref_t\right)(x),\end{aligned}

其中在 (i)(i) 中我们使用了式(12)中 pt(x)p_t(x) 的定义,在 (ii)(ii) 中我们使用了 Conditional Probability Path pt(z)p_t(\cdot|z) 的连续性方程,在 (iii)(iii) 中我们使用式(22)交换了积分和散度算子,在 (iv)(iv) 中我们乘以并除以 pt(x)p_t(x),在 (v)(v) 中我们使用了式(18)。上述等式链的开头和结尾表明 uttarget\uref_t 满足连续性方程。根据第 3.2 节,这足以推出式(19),我们就完成了。

学习 Marginal Vector Field

现在,我们准备描述训练算法。Flow Matching 的目标是训练神经网络 utθu_t^\theta,使其等于 Marginal Vector Field uttarget\uref_t。如果这成立,我们知道根据第 3.2 节,端点 X1pdataX_1\sim \pdata 具有期望的分布。在下文中,我们用 Unif=Unif[0,1]\text{Unif}=\text{Unif}_{[0,1]} 表示区间 [0,1][0,1] 上的均匀分布,用 E\mathbb{E} 表示随机变量的期望值。获得 utθuttargetu_t^\theta\approx \uref_t 的一种直观方法是使用均方误差,即使用Flow Matching 损失,定义为

LFM(θ)=EtUnif,xpt[utθ(x)uttarget(x)2](24)\begin{aligned}\Lmarg(\theta)&=\mathbb{E}_{t\sim\text{Unif},x\sim p_t}[\|u_t^\theta(x) - \uref_t(x)\|^2]\end{aligned}\tag{24}
=(i)EtUnif,zpdata,xpt(z)[utθ(x)uttarget(x)2],(25)\begin{aligned}&\overset{(i)}{=} \mathbb{E}_{t\sim\text{Unif},z\sim \pdata, x\sim p_t(\cdot|z)}[\|u_t^\theta(x) - \uref_t(x)\|^2],\end{aligned}\tag{25}

其中 pt(x)=pt(xz)pdata(z)dzp_t(x)=\int p_t(x|z)\pdata(z) \dd z 是 Marginal Probability Path,在 (i)(i) 中我们使用了式(12)给出的采样过程。直观上,这个损失说:首先,抽取一个随机时间 t[0,1]t \in [0,1]。其次,从我们的数据集中抽取一个随机点 zz,从 pt(z)p_t(\cdot|z) 中采样(例如,通过添加一些噪声),并计算 utθ(x)u_t^\theta(x)。最后,计算我们神经网络的输出与 Marginal Vector Field uttarget(x)\uref_t(x) 之间的均方误差。不幸的是,我们还没有完成。虽然我们通过第 3.2 节知道 uttarget\uref_t 的公式,但我们无法高效计算它,因为积分是难以处理的。相反,我们将利用条件速度场 uttarget(xz)\uref_t(x|z) 是可处理的事实。
为此,让我们定义conditional flow matching 损失

LCFM(θ)=EtUnif,zpdata,xpt(z)[utθ(x)uttarget(xz)2].(26)\Lcond(\theta) = \mathbb{E}_{t\sim \Unif, z\sim \pdata, x\sim p_t(\cdot|\dap)}[\|u_t^\theta(x) - \uref_t(x|\dap)\|^2]. \tag{26}

注意与式(24)的区别:我们使用 Conditional Vector Field uttarget(xz)\uref_t(x|z) 而不是边缘向量 uttarget(x)\uref_t(x)。由于我们有 uttarget(xz)\uref_t(x|z) 的解析公式,我们可以轻松地最小化上述损失。但是等等,如果我们关心的是 Marginal Vector Field,那么回归 Conditional Vector Field 有什么意义呢?事实证明,通过显式地回归可处理的 Conditional Vector Field,我们隐式地回归了难以处理的 Marginal Vector Field。下一个结果使这种直觉变得精确。

定理 12

边缘 Flow Matching 损失等于 conditional flow matching 损失加上一个常数。也就是说,

LFM(θ)=LCFM(θ)+C,\Lmarg(\theta) = \Lcond(\theta) + C,

其中 CC 独立于 θ\theta。因此,它们的梯度一致:

θLFM(θ)=θLCFM(θ).\nabla_\theta \Lmarg(\theta) = \nabla_\theta \Lcond(\theta).

因此,使用例如随机梯度下降(SGD)最小化 LCFM(θ)\Lcond(\theta) 等同于以相同方式最小化 LFM(θ)\Lmarg(\theta)。特别地,对于 θ\theta^* 的最小化器 LCFM(θ)\Lcond(\theta),将有 utθ=uttargetu_t^{\theta^*}=\uref_t,即神经网络将等于 Marginal Vector Field(假设无限表达力的参数化)。

证明(直接证明)

证明通过将均方误差展开为三个分量并移除常数来完成:

LFM(θ)=(i)EtUnif,xpt[utθ(x)uttarget(x)2]=(ii)EtUnif,xpt[utθ(x)22utθ(x)Tuttarget(x)+uttarget(x)2]=(iii)EtUnif,xpt[utθ(x)2]2EtUnif,xpt[utθ(x)Tuttarget(x)]+EtUnif[0,1],xpt[uttarget(x)2]=:C1=(iv)EtUnif,zpdata,xpt(z)[utθ(x)2]2EtUnif,xpt[utθ(x)Tuttarget(x)]+C1\begin{aligned}\Lmarg(\theta)&\overset{(i)}{=}\mathbb{E}_{t\sim\text{Unif},x\sim p_t}[\|u_t^\theta(x) - \uref_t(x)\|^2]\\ &\overset{(ii)}{=}\mathbb{E}_{t\sim\text{Unif},x\sim p_t}[\|u_t^\theta(x)\|^2 - 2u_t^\theta(x)^T\uref_t(x) + \|\uref_t(x)\|^2]\\ &\overset{(iii)}{=}\mathbb{E}_{t\sim\text{Unif},x\sim p_t}\left[\|u_t^\theta(x)\|^2\right] - 2\mathbb{E}_{t\sim\text{Unif},x\sim p_t}[u_t^\theta(x)^T\uref_t(x)] + \underbrace{\mathbb{E}_{t\sim\text{Unif}_{[0,1]}, x\sim p_t}[\|\uref_t(x)\|^2]}_{=:C_1}\\ &\overset{(iv)}{=}\mathbb{E}_{t\sim\text{Unif},z\sim\pdata, x\sim p_t(\cdot|z)}[\|u_t^\theta(x)\|^2] - 2\mathbb{E}_{t\sim\text{Unif},x\sim p_t}[u_t^\theta(x)^T\uref_t(x)] + C_1\end{aligned}

其中 (i)(i) 由定义成立,在 (ii)(ii) 中我们使用了公式 ab2=a22aTb+b2\|a-b\|^2=\|a\|^2-2a^Tb+\|b\|^2,在 (iii)(iii) 中我们定义常数 C1C_1,在 (iv)(iv) 中我们使用了 ptp_t 的采样过程,由式(12)给出。让我们重新表达第二项:

EtUnif,xpt[utθ(x)Tuttarget(x)]=(i)01pt(x)utθ(x)Tuttarget(x)dxdt=(ii)01pt(x)utθ(x)T[uttarget(xz)pt(xz)pdata(z)pt(x)dz]dxdt=(iii)01utθ(x)Tuttarget(xz)pt(xz)pdata(z)dzdxdt=(iv)EtUnif,zpdata,xpt(z)[utθ(x)Tuttarget(xz)]\begin{aligned}\mathbb{E}_{t\sim\text{Unif},x\sim p_t}[u_t^\theta(x)^T\uref_t(x)]&\overset{(i)}{=}\int\limits_{0}^{1}\int p_t(x)u_t^\theta(x)^T\uref_t(x)\,\dd x\, \dd t\\ &\overset{(ii)}{=}\int\limits_{0}^{1}\int p_t(x)u_t^\theta(x)^T\left[\int \uref_t(x|z)\frac{ p_t(x|z)\pdata(z)}{p_t(x)}\dd z\right]\dd x\, \dd t\\ &\overset{(iii)}{=}\int\limits_{0}^{1}\int\int u_t^\theta(x)^T\uref_t(x|z) p_t(x|z)\pdata(z)\,\dd z\,\dd x\, \dd t\\ &\overset{(iv)}{=}\mathbb{E}_{t\sim\text{Unif},z\sim \pdata, x\sim p_t(\cdot|z)}[u_t^\theta(x)^T\uref_t(x|z)]\end{aligned}

其中在 (i)(i) 中我们将期望值表示为积分,在 (ii)(ii) 中我们使用式(18),在 (iii)(iii) 中我们使用积分是线性的这一事实,在 (iv)(iv) 中我们将积分表示为期望值。注意这确实是证明的关键步骤:等式的开始使用了 Marginal Vector Field uttarget(x)\uref_t(x),而结尾使用了 Conditional Vector Field uttarget(xz)\uref_t(x|z)。我们将其代入 LFM\Lmarg 的方程得到:

LFM(θ)=(i)EtUnif,zpdata,xpt(z)[utθ(x)2]2EtUnif,zpdata,xpt(z)[utθ(x)Tuttarget(xz)]+C1=(ii)EtUnif,zpdata,xpt(z)[utθ(x)22utθ(x)Tuttarget(xz)+uttarget(xz)2uttarget(xz)2]+C1=(iii)EtUnif,zpdata,xpt(z)[utθ(x)uttarget(xz)2]+EtUnif,zpdata,xpt(z)[uttarget(xz)2]C2+C1=(iv)LCFM(θ)+C2+C1=:C\begin{aligned}\Lmarg(\theta)&\overset{(i)}{=}\mathbb{E}_{t\sim\text{Unif},z\sim \pdata, x\sim p_t(\cdot|z)}[\|u_t^\theta(x)\|^2]-2\mathbb{E}_{t\sim\text{Unif},z\sim \pdata, x\sim p_t(\cdot|z)}[u_t^\theta(x)^T\uref_t(x|z)] +C_1\\ &\overset{(ii)}{=}\mathbb{E}_{t\sim\text{Unif},z\sim \pdata, x\sim p_t(\cdot|z)}[\|u_t^\theta(x)\|^2-2u_t^\theta(x)^T\uref_t(x|z)+\|\uref_t(x|z)\|^2-\|\uref_t(x|z)\|^2] +C_1\\ &\overset{(iii)}{=}\mathbb{E}_{t\sim\text{Unif},z\sim \pdata, x\sim p_t(\cdot|z)}[\|u_t^\theta(x)-\uref_t(x|z)\|^2]+\underbrace{\mathbb{E}_{t\sim\text{Unif},z\sim \pdata, x\sim p_t(\cdot|z)}[-\|\uref_t(x|z)\|^2]}_{C_2} +C_1\\ &\overset{(iv)}{=}\Lcond(\theta) + \underbrace{C_2+C_1}_{=:C}\end{aligned}

其中在 (i)(i) 中我们代入了推导出的方程,在 (ii)(ii) 中我们加上并减去相同的值,在 (iii)(iii) 中我们再次使用公式 ab2=a22aTb+b2\|a-b\|^2=\|a\|^2-2a^Tb+\|b\|^2,在 (iv)(iv) 中我们在 θ\theta 中定义了一个常数。这完成了证明。

因此,Flow Matching 训练包括最小化 conditional flow matching 损失。 训练过程总结在算法 3 中,并在图 Figure 7 中可视化。注意该算法有几个显著特点:首先,我们在训练期间从未实际模拟任何 ODE。人们称这一算法特性为无模拟。这使得训练极其廉价,因为你不必在训练期间展开 ODE 的轨迹(这需要很多步骤)。其次,训练是一个简单的回归目标——我们只是对 uttarget(xz)\uref_t(x|z) 进行回归。所以毕竟与监督学习没有太大区别。最后,该算法极其简单——很难想到更简单的训练目标。所有这些使得 Flow Matching 成为大规模机器学习模型极具吸引力的方法。一旦 utθu_t^{\theta} 被训练好,我们可以模拟 Flow Model

dXt=utθ(Xt)dt,X0pinit(27)\dd X_t = u_t^\theta(X_t)\, \dd t,\quad\quad X_0\sim\pinit \tag{27}

例如通过算法 1 获得样本 X1pdataX_1\sim \pdata。整个流程在文献中被称为 Flow Matching [25, 27, 1, 26]。
现在让我们为 Gaussian Probability Path 实例化 conditional flow matching 损失:

例 13(Gaussian Conditional Probability Path 的 Flow Matching)
Algorithm 3
算法 3 · Flow Matching 训练过程(Gaussian CondOT 路径 pt(xz)=N(tz,(1t)2Id)p_t(x\mid z)=\mathcal{N}(tz,(1-t)^2I_d)
输入: 样本数据集 zpdataz\sim\pdata,neural-network Vector Field utθu_t^\theta

- 循环: 遍历每个数据 mini-batch
- 从数据集中采样一个样本 zz
- 采样随机时间 tUnif[0,1]t\sim\mathrm{Unif}[0,1]
- 采样噪声 ϵN(0,Id)\epsilon\sim\mathcal{N}(0,I_d)
- 构造训练点:
x=tz+(1t)ϵ(一般情形: xpt(z))x=t z+(1-t)\epsilon \qquad (\text{一般情形: }x\sim p_t(\cdot\mid z))
- 计算损失:
L(θ)=utθ(x)(zϵ)2(一般情形: utθ(x)uttarget(xz)2)\mathcal{L}(\theta)=\lVert u_t^\theta(x)-(z-\epsilon)\rVert^2 \qquad (\text{一般情形: }\lVert u_t^\theta(x)-\uref_t(x\mid z)\rVert^2)
- 更新 θgrad_update(L(θ))\theta\gets\mathrm{grad\_update}(\mathcal{L}(\theta))
- 结束循环

让我们回到 Gaussian Probability Path pt(z)=N(αtz;βt2Id)p_t(\cdot|z)=\mathcal{N}(\alpha_t z; \beta_t^2 I_d) 的例子,其中我们可以通过以下方式从条件路径采样

ϵN(0,Id)xt=αtz+βtϵN(αtz,βt2Id)=pt(z).(28)\epsilon\sim\mathcal{N}(0,I_d)\quad \Rightarrow\quad x_t = \alpha_t z + \beta_t \epsilon \sim \mathcal{N}(\alpha_tz,\beta_t^2I_d)=p_t(\cdot|z). \tag{28}

如我们在式(20)中推导的,Conditional Vector Field uttarget(xz)\uref_t(x|z) 由下式给出

uttarget(xz)=(α˙tβ˙tβtαt)z+β˙tβtx,(29)\begin{aligned}\uref_t(x|z) =& \left(\dot{\alpha}_t-\frac{\dot{\beta}_t}{\beta_t}\alpha_t\right)z+\frac{\dot{\beta}_t}{\beta_t}x,\end{aligned}\tag{29}

其中 α˙t=tαt\dot{\alpha}_t=\partial_t\alpha_tβ˙t=tβt\dot{\beta}_t=\partial_t\beta_t 分别是相应的时间导数。将此公式代入,conditional flow matching 损失变为

LCFM(θ)=EtUnif,zpdata,xN(αtz,βt2Id)[utθ(x)(α˙tβ˙tβtαt)zβ˙tβtx2](30)\begin{aligned}\Lcond(\theta) &= \mathbb{E}_{t\sim \text{Unif},z\sim \pdata, x\sim \mathcal{N}(\alpha_tz,\beta_t^2I_d)}[\lVert u_t^\theta(x)-\left(\dot{\alpha}_t-\frac{\dot{\beta}_t}{\beta_t}\alpha_t\right)z-\frac{\dot{\beta}_t}{\beta_t}x\rVert^2]\end{aligned}\tag{30}
=(i)EtUnif,zpdata,ϵN(0,Id)[utθ(αtz+βtϵ)(α˙tz+β˙tϵ)2](31)\begin{aligned}&\overset{(i)}{=}\mathbb{E}_{t\sim\Unif,z\sim \pdata, \epsilon\sim \mathcal{N}(0,I_d)}[\|u_t^\theta(\alpha_tz+\beta_t\epsilon)-(\dot{\alpha}_tz+\dot{\beta}_t\epsilon)\|^2]\end{aligned}\tag{31}

其中在 (i)(i) 中我们代入了式(28),并将 xx 替换为 αtz+βtϵ\alpha_tz+\beta_t\epsilon。注意 LCFM\Lcond 的简洁性:我们采样一个数据点 zz,采样一些噪声 ϵ\epsilon,然后计算均方误差。让我们针对 αt=t\alpha_t=tβt=1t\beta_t=1-t 的特殊情况使其更具体。相应的概率 pt(xz)=N(tz,(1t)2)p_t(x|z)=\mathcal{N}(tz,(1-t)^2) 有时被称为(Gaussian)CondOT Probability Path。那么我们有 α˙t=1,β˙t=1\dot{\alpha}_t=1,\dot{\beta}_t=-1,因此

Lcfm(θ)=EtUnif,zpdata,ϵN(0,Id)[utθ(tz+(1t)ϵ)(zϵ)2]\begin{aligned}\mathcal{L}_{\text{cfm}}(\theta)=&\mathbb{E}_{t\sim\Unif,z\sim \pdata, \epsilon\sim \mathcal{N}(0,I_d)}[\|u_t^\theta(tz+(1-t)\epsilon)-(z-\epsilon)\|^2]\end{aligned}

许多著名的先进模型都使用这种简单而有效的程序进行训练,例如 Stable Diffusion 3、Meta 的 Movie Gen Video,以及可能许多其他专有模型。在图 Figure 7 中,我们在一个简单示例中对其进行了可视化,并在算法 3 中总结了训练过程。

Figure 7第 3.3 节中 Gaussian CondOT Probability Path 的图示:从训练好的 Flow Matching 模型模拟 ODE。数据分布是棋盘图案(右上角)。顶行:来自真实 Marginal Probability Path pt(x)p_t(x) 的直方图。底行:来自 Flow Matching 模型的样本直方图。可以看出,训练后顶行和底行匹配(直到训练误差)。该模型使用算法 3 进行训练。

让我们总结本节的结果。

小结 14(Flow Matching)

Flow Matching 训练包括学习 Marginal Vector Field uttarget\uref_t 为了构建它,我们选择一个满足 p0(z)=pinitp_0(\cdot|z)=\pinitp1(z)=δzp_1(\cdot|z)=\delta_{z}Conditional Probability Path pt(xz)p_t(x|\dap)。接下来,我们找到一个Conditional Vector Field uttarget(xz)\uref_t(x|z),使得其对应的流 ψttarget(xz)\psiref_t(x|z) 满足

X0pinitXt=ψttarget(X0z)pt(z),X_0\sim \pinit \quad \Rightarrow \quad X_t = \psiref_t(X_0|z) \sim p_t(\cdot|z),

或者等价地,uttarget\uref_t 满足连续性方程。那么由下式定义的Marginal Vector Field

uttarget(x)=uttarget(xz)pt(xz)pdata(z)pt(x)dz,(32)\uref_t(x) = \int \uref_t(x|z)\frac{p_t(x|z)\pdata(z)}{p_t(x)}\dd z, \tag{32}

遵循 Marginal Probability Path,即

X0pinit,dXt=uttarget(Xt)dtXtpt(0t1).(33)\begin{aligned}X_0&\sim\pinit,\quad \dd X_t =\uref_t(X_t)\dd t\Rightarrow X_t\sim p_t\quad (0\leq t\leq 1).\end{aligned}\tag{33}

特别地,对于这个 ODE,X1pdataX_1\sim \pdata,因此 uttarget\uref_t 按预期“将噪声转换为数据”。为了学习它,我们最小化 conditional flow matching 损失

LCFM(θ)=EtUnif,zpdata,xpt(z)[utθ(x)uttarget(xz)2].(34)\Lcond(\theta) = \mathbb{E}_{t\sim \Unif, z\sim \pdata, x\sim p_t(\cdot|\dap)}[\|u_t^\theta(x) - \uref_t(x|\dap)\|^2]. \tag{34}

最广泛使用的例子是Gaussian Probability Path。对于这种情况,公式变为:

pt(xz)=N(x;αtz,βt2Id)(35)\begin{aligned}p_t(x|z) =& \mathcal{N}(x;\alpha_t z,\beta_t^2 I_d)\end{aligned}\tag{35}
utflow(xz)=(α˙tβ˙tβtαt)z+β˙tβtx(36)\begin{aligned}\uflow_t(x|z)=&\left(\dot{\alpha}_t-\frac{\dot{\beta}_t}{\beta_t}\alpha_t\right)z+\frac{\dot{\beta}_t}{\beta_t}x\end{aligned}\tag{36}
LCFM(θ)=EtUnif,zpdata,ϵN(0,Id)[utθ(αtz+βtϵ)(α˙tz+β˙tϵ)2](37)\begin{aligned}\Lcond(\theta) =&\mathbb{E}_{t\sim\Unif,z\sim \pdata, \epsilon\sim \mathcal{N}(0,I_d)}[\|u_t^\theta(\alpha_tz+\beta_t\epsilon)-(\dot{\alpha}_tz+\dot{\beta}_t\epsilon)\|^2]\end{aligned}\tag{37}

对于noise schedules αt,βtR\alpha_t,\beta_t\in\mathbb{R},即我们选择的连续可微、单调的函数,使得 α0=β1=0\alpha_0=\beta_1=0 α1=β0=1\alpha_1=\beta_0=1(例如 αt=t,βt=1t\alpha_t=t,\beta_t=1-t)。

Score Function 与 Score Matching

在上一节中,我们展示了如何使用 Flow Matching 训练 Flow Model。在本节中,我们讨论 Diffusion Model,并演示如何使用 Score Matching 训练它们。

Conditional Score Function 与 Marginal Score Function

Figure 8Score Function logq(x)\nabla\log q(x) 的示意图,绘制为一般概率分布 q(x)q(x)(左)的黑色行(右)。

到目前为止,我们研究的核心对象是 Vector Field ut(x)u_t(x)。Diffusion Model [43, 43] 采取了不同的视角,聚焦于Score Function。因此,在本节中,我们将用 Score Function 的语言重新表述我们在此学到的内容——提供一个新颖的视角。设 q(x)q(x) 为任意概率分布。那么 qqScore Function定义为 logq(x)\nabla\log q(x),即 qq 关于 xx 的对数似然的梯度。得分具有直观的含义:logq(x)\nabla\log q(x) 是相对于对数似然的最陡上升方向。如图 8 所示。

让我们回到 Conditional Probability Path pt(xz)p_t(x|z) 和 Marginal Probability Path pt(x)p_t(x) 的设置,如第 3 节所述。那么我们可以等价地定义Conditional Score Functionlogpt(xz)\nabla\log p_t(x|z)Marginal Score Functionlogpt(x)\nabla\log p_t(x)。类似于式(18),边缘得分可以通过 Conditional Score Function logpt(xz)\nabla \log p_t(x|z) 表示为

logpt(x)=logpt(xz)pt(xz)pdata(z)pt(x)dz.(38)\nabla\log p_t(x) =\int \nabla \log p_t(x|z)\frac{ p_t(x|z)\pdata(z)}{p_t(x)}\dd z. \tag{38}

因此,条件得分与边缘得分之间的关系类似于 Conditional and Marginal Vector Fields 之间的关系。注意,我们可以通过以下方式证明式(38):

logpt(x)=pt(x)pt(x)=pt(xz)pdata(z)dzpt(x)=pt(xz)pdata(z)dzpt(x)=logpt(xz)pt(xz)pdata(z)pt(x)dz,(39)\nabla\log p_t(x) = \frac{\nabla p_t(x)}{p_t(x)} =\frac{\nabla \int p_t(x|z)\pdata(z)\dd z}{p_t(x)} =\frac{\int \nabla p_t(x|z)\pdata(z)\dd z}{p_t(x)} =\int \nabla \log p_t(x|z)\frac{ p_t(x|z)\pdata(z)}{p_t(x)}\dd z, \tag{39}

其中我们使用了规则 ylogy=1/y\partial_{y} \log y=1/y 并结合链式法则两次。

例 15(Gaussian Probability Path 的 Score Function。)

对于 Gaussian 路径 pt(xz)=N(x;αtz,βt2Id)p_t(x|z)=\mathcal{N}(x;\alpha_t z,\beta_t^2 I_d),我们可以利用 Gaussian 概率密度的形式(参见 式(97))得到

logpt(xz)=logN(x;αtz,βt2Id)=xαtzβt2.(40)\nabla \log p_t(x|z) = \nabla\log \mathcal{N}(x;\alpha_t z,\beta_t^2 I_d) = -\frac{x-\alpha_t z}{\beta_t^2}. \tag{40}

注意,Gaussian Probability Path 的 Score Function 是 xx 和 z 的线性函数。Conditional Vector Field ut(xz)u_t(x|z) 也是如此(参见式(20))。因此,可以在两者之间进行转换,如下一个命题所示。

命题 1(Gaussian Probability Path 的转换公式)

对于 Gaussian Probability Path pt(xz)=N(αtz,βt2Id)p_t(x|z)=\mathcal{N}(\alpha_t z,\beta_t^2 I_d),条件(分别地,边缘)Vector Field 与条件(分别地,边缘)得分通过以下恒等式相关联

uttarget(xz)=atlogpt(xz)+btx,at=(βt2α˙tαtβ˙tβt),bt=α˙tαt(41)\begin{aligned}\uref_t(x|z)=&a_t\nabla\log p_t(x|z)+b_tx,\quad a_t=\left(\beta_t^2\frac{\dot{\alpha}_t}{\alpha_t}-\dot{\beta}_t\beta_t\right),\quad b_t=\frac{\dot{\alpha}_t}{\alpha_t}\end{aligned}\tag{41}
uttarget(x)=atlogpt(x)+btx.(42)\begin{aligned}\uref_t(x)=&a_t\nabla\log p_t(x)+b_tx.\end{aligned}\tag{42}

特别地,我们注意到条件(相应地,边缘)Vector Field 可以从条件(相应地,边缘)得分中恢复,反之亦然。

证明

对于 Conditional Vector Field 和条件得分,我们可以推导出:

Dt(xz)=z,Dt(x)=zpt(xz)pdata(z)pt(x)dz=(i)1α˙tβtαtβ˙t(βtuttarget(xt)β˙txt).(43)D_t(x|z)=z, \quad D_t(x) =\int z\,\frac{ p_t(x|z)\pdata(z)}{p_t(x)}\dd z \overset{(i)}{=} \frac{1}{\dot{\alpha}_t\beta_t-\alpha_t\dot{\beta}_t}(\beta_t \uref_t(x_t)-\dot{\beta}_tx_t). \tag{43}

使用 SDE 采样

Figure 9定理 17 的图示。使用 SDE 模拟 Probability Path。这重复了图 6 中的图,但使用 式(44) 进行 SDE 采样。数据分布 pdata\pdata 在蓝色背景中。Gaussian pinit\pinit 在红色背景中。顶行:条件路径。底行:Marginal Probability Path。可以看出,SDE 将样本从 pinit\pinit 传输到 δz\delta_{z}(对于条件路径)以及到 pdata\pdata(对于边缘路径)。
X0pinit,dXt=uttarget(Xt)dt+σt22logpt(Xt)dt+σtdWt(44)\begin{aligned}X_0 &\sim \pinit,\qquad & \dd X_t &=\textcolor{RoyalBlue}{\uref_t(X_t)\dd t} + \textcolor{ForestGreen}{\frac{\sigma_t^2}{2}\nabla\log p_t(X_t)\dd t + \sigma_t\dd W_t}\end{aligned}\tag{44}
Xtpt(0t1).(45)\begin{aligned}\Rightarrow\quad X_t &\sim p_t\quad (0\le t\le 1). &&\end{aligned}\tag{45}
X0pinit,dXt=[(at+σt22)logpt(Xt)+btXt]dt+σtdWt(46)\begin{aligned}X_0\sim&\,\pinit,\quad \dd X_t =\left[\left(a_t+\frac{\sigma_t^2}{2}\right)\nabla\log p_t(X_t)+b_tX_t\right]\dd t +\sigma_t\dd W_t\end{aligned}\tag{46}
Xtpt(0t1)(47)\begin{aligned}\Rightarrow X_t\sim& \,p_t\quad (0\leq t\leq 1)\end{aligned}\tag{47}
Δwt(x)=i=1d2xi2wt(x)=div(wt)(x),(48)\begin{aligned}\Delta w_t(x)=&\sum\limits_{i=1}^{d}\frac{\partial^2}{\partial x_i^2}w_t(x)=\divv(\nabla w_t)(x),\end{aligned}\tag{48}
uttarget(xz)=(α˙tβ˙tβtαt)z+β˙tβtx=(i)(βt2α˙tαtβ˙tβt)(αtzxβt2)+α˙tαtx=(βt2α˙tαtβ˙tβt)logpt(xz)+α˙tαtx\begin{aligned}\uref_t(x|z)=&\left(\dot{\alpha}_t-\frac{\dot{\beta}_t}{\beta_t}\alpha_t\right)z+\frac{\dot{\beta}_t}{\beta_t}x \overset{(i)}{=}\left(\beta_t^2\frac{\dot{\alpha}_t}{\alpha_t}-\dot{\beta}_t\beta_t\right)\left(\frac{\alpha_tz-x}{\beta_t^2}\right)+\frac{\dot{\alpha}_t}{\alpha_t}x=\left(\beta_t^2\frac{\dot{\alpha}_t}{\alpha_t}-\dot{\beta}_t\beta_t\right)\nabla\log p_t(x|z)+\frac{\dot{\alpha}_t}{\alpha_t}x\end{aligned}
tpt(x)=div(ptut)(x)+σt22Δpt(x) for all xRd,0t1,(49)\partial_t p_t(x) = -\divv (p_t u_t)(x)+\frac{\sigma_t^2}{2}\Delta p_t (x)\quad \text{ for all }x\in\R^d, 0\leq t\leq 1, \tag{49}
dXt=σt22logp(Xt)dt+σtdWt,(50)\dd X_t = \frac{\sigma_t^2}{2}\nabla\log p(X_t)\dd t + \sigma_t dW_t , \tag{50}
Figure 10顶行:在由 式(50) 给出的 Langevin 动力学下演化的粒子,其中 p(x)p(x) 取为具有 5 个模态的 Gaussian 混合。底行:顶行中相同样本的核密度估计。可以看出,样本的分布收敛到平衡分布 pp(蓝色背景颜色)。

其中在 (i)(i) 中我们只是做了一些代数运算。通过取积分,相同的恒等式对于边缘流 Vector Field 和 Marginal Score Function 也成立:

utarget(x)=uttarget(xz)pt(xz)pdata(z)pt(x)dz=[atlogpt(xz)+btx]pt(xz)pdata(z)pt(x)dz=(i)atlogpt(x)+btx\begin{aligned}\uref(x)= \int \uref_t(x|z)\frac{ p_t(x|z)\pdata(z)}{p_t(x)}\dd z =&\int\left[a_t\nabla\log p_t(x|z)+b_tx\right] \frac{ p_t(x|z)\pdata(z)}{p_t(x)}\dd z\\ \overset{(i)}{=}&a_t\nabla\log p_t(x)+b_tx\end{aligned}

Score Matching

logpt(xz)=xαtzβt2.(51)\nabla\log p_t(x|z) = -\frac{x-\alpha_t z}{\beta_t^2}. \tag{51}

其中在 (i)(i) 中我们使用了式(38)以及后验密度积分为 11 的事实。

第 4.1 节 是引人注目的,因为它说明一旦我们学习了 uttarget\uref_t,我们也学习了 Score Function logpt(x)\nabla\log p_t(x),反之亦然。因此,许多 Diffusion Model 通过神经网络学习 Score Function logpt(x)\nabla\log p_t(x)。我们将在 第 4.3 节 中讨论这一点。

备注 16(分数的重新参数化)

式(41)中 Gaussian Probability Path 的重参数化公式之所以可能,是因为两边(Conditional Vector Field 和条件得分)都是 xxzz线性函数。一旦我们边缘化(Marginal Vector Field 和边缘得分),两边都只是后验均值 Ezx[z]\EE_{z|x}\left[z\right] 的线性重参数化。因此,任何能够恢复 Ezx[z]\EE_{z|x}\left[z\right] 的量都可以反过来用于恢复无 Conditional Vector Field 和得分。此外,从数值/训练稳定性的角度来看,这样做甚至可能更可取。一个常见的选择是后验均值本身,通常称为去噪器。形式上,我们定义条件和边缘去噪器

这里,(i)(i) 遵循与第 4.1 节类似的推导。去噪器有一个非常直观的解释:它是给定噪声数据 xx 时干净数据 zz 的期望值。〔脚注〕 人们通常称这样的模型为Denoising Diffusion Model,因为学习 DtD_t 和学习 uttarget\uref_t 在理论上是等价的。

Algorithm 4
算法 4 · Gaussian Probability Path 的 Score Matching 训练过程
输入: 样本数据集 zpdataz\sim\pdata,得分网络 stθs_t^\theta 或噪声预测器 ϵtθ\epsilon_t^\theta

- 循环: 遍历每个数据 mini-batch
- 从数据集中采样一个样本 zz
- 采样随机时间 tUnif[0,1]t\sim\mathrm{Unif}[0,1]
- 采样噪声 ϵN(0,Id)\epsilon\sim\mathcal{N}(0,I_d)
- 设置 xt=αtz+βtϵx_t=\alpha_t z+\beta_t\epsilon(一般情形:xtpt(z)x_t\sim p_t(\cdot\mid z)
- 计算损失:
L(θ)=stθ(xt)+ϵβt2(一般情形: stθ(xt)logpt(xtz)2或者L(θ)=ϵtθ(xt)ϵ2\begin{aligned} \mathcal{L}(\theta) &=\left\lVert s_t^\theta(x_t)+\frac{\epsilon}{\beta_t}\right\rVert^2 &&\text{(一般情形: }\lVert s_t^\theta(x_t)-\nabla\log p_t(x_t\mid z)\rVert^2\text{)}\\ \text{或者}\quad \mathcal{L}(\theta) &=\lVert\epsilon_t^\theta(x_t)-\epsilon\rVert^2 \end{aligned}
- 对 L(θ)\mathcal{L}(\theta) 做梯度下降,更新模型参数 θ\theta
- 结束循环

到目前为止,我们已经展示了如何通过 Marginal Vector Field uttarget\uref_t 构造一个遵循期望 Probability Path ptp_t 的 ODE 的轨迹 XtX_t。但这种方法仅限于 Flow Model。Diffusion Model 呢?使用 Score Function,现在让我们将此结果扩展到 SDE。

定理 17(SDE 扩展技巧)

如前定义 Conditional Vector Field 和 Marginal Vector Field uttarget(xz)\uref_t(x|z)uttarget(x)\uref_t(x)。然后,对于任何扩散系数 σt0\sigma_t\geq 0,我们可以通过向原始 ODE 的动力学中添加随机动力学来构造一个 SDE,如下所示:

=[uttarget(Xt)+σt22logpt(Xt)]dt+σtdWt\begin{aligned}&& &=\Big[\textcolor{RoyalBlue}{\uref_t(X_t)}+\textcolor{ForestGreen}{\frac{\sigma_t^2}{2}\nabla\log p_t(X_t)}\Big]\dd t + \textcolor{ForestGreen}{\sigma_t}\dd W_t\end{aligned}

特别地,对于此 SDE,X1pdataX_1\sim \pdata。我们注意到随机动力学与Langevin dynamics密切相关,可以看作是在保持边缘分布 ptp_t 的同时注入噪声。我们在 第 4.2 节 中简要讨论 Langevin dynamics。

我们在图 9 中展示了第 4.2 节中描述的动力学。可以看到,轨迹现在是锯齿状的,说明了 SDE 演化的随机性质。然而,正如第 4.2 节所确立的,边缘分布 ptp_t 保持不变。请注意,上述结果令人惊讶,因为我们可以选择任何扩散系数 σt0\sigma_{t}\geq 0,即使在训练网络之后也可以。理论上,第 4.2 节对任何 σt\sigma_t 的选择都成立。然而,在实践中,我们同时遭受训练误差(神经网络不能完美逼近 Marginal Vector Field 和评分)和模拟误差(例如,对于 σt0\sigma_{t}\gg0,我们需要在算法 2 中采取过小的步长)。在实践中,对于固定的训练模型,存在一个最优的 σt0\sigma_{t}\geq 0,可以通过经验确定 [23, 1, 28]。〔脚注〕

对于 Gaussian Probability Path,我们通过学习 Marginal Vector Field 免费获得了 Score Function。

例 18(GaussianSDE 扩展技巧)

根据第 4.1 节,对于 Gaussian Probability Path,我们可以仅使用 Score Function 来表达第 4.2 节中的 SDE:

其中 at,bta_t,b_t 的定义如第 4.1 节所述。

在本节的其余部分,我们将通过Fokker-Planck 方程证明第 4.2 节,该方程将连续性方程从 ODE 扩展到 SDE。为此,让我们首先定义拉普拉斯算子 Δ\Delta,通过

对于标量场 wt:RdRw_t:\R^d\to\R

定理 19(Fokker-Planck 方程)

ptp_t 为一个 Probability Path,并考虑 SDE

X0pinit,dXt=ut(Xt)dt+σtdWt.X_0\sim \pinit, \quad \dd X_t = u_t(X_t)\dd t + \sigma_t\dd W_t.

那么,当且仅当 Fokker-Planck 方程成立时,XtX_t 对所有 0t10\leq t\leq 1 具有分布 ptp_t

Fokker-Planck 方程的自包含证明可在 第 B 节 中找到。注意,当 σt=0\sigma_t=0 时,从 Fokker-Planck 方程可恢复出第 3.2 节。额外的拉普拉斯项 Δpt\Delta p_t 起初可能难以理解。熟悉物理的读者会注意到,同样的项也出现在热方程中(实际上热方程是 Fokker-Planck 方程的一个特例)。热量在介质中扩散。我们也添加一个扩散过程(不是物理上的,而是数学上的),因此我们添加了这个额外的拉普拉斯项。
现在让我们使用 Fokker-Planck 方程来帮助证明第 4.2 节。

证明(第 4.2 节的证明)

根据第 4.2 节,我们需要证明式(44)中定义的 SDE 满足 ptp_t 的 Fokker-Planck 方程。我们可以通过直接计算来做到这一点:

tpt(x)=(i)div(ptuttarget)(x)=(ii)div(ptuttarget)(x)σt22Δpt(x)+σt22Δpt(x)=(iii)div(ptuttarget)(x)div(σt22pt)(x)+σt22Δpt(x)=(iv)div(ptuttarget)(x)div(pt[σt22logpt])(x)+σt22Δpt(x)=(v)div(pt[uttarget+σt22logpt])(x)+σt22Δpt(x),\begin{aligned}\partial_t p_t(x) \overset{(i)}{=}& - \divv(p_t\uref_t)(x)\\ \overset{(ii)}{=}& - \divv(p_t\uref_t)(x) -\frac{\sigma_t^2}{2}\Delta p_t(x)+\frac{\sigma_t^2}{2}\Delta p_t(x)\\ \overset{(iii)}{=}& - \divv(p_t\uref_t)(x) -\divv(\frac{\sigma_t^2}{2}\nabla p_t)(x)+\frac{\sigma_t^2}{2}\Delta p_t(x)\\ \overset{(iv)}{=}& - \divv(p_t\uref_t)(x) -\divv(p_t\left[\frac{\sigma_t^2}{2}\nabla \log p_t\right])(x)+\frac{\sigma_t^2}{2}\Delta p_t(x)\\ \overset{(v)}{=}&- \divv\left(p_t\left[\uref_t+\frac{\sigma_t^2}{2}\nabla \log p_t\right]\right)(x)+\frac{\sigma_t^2}{2}\Delta p_t(x),\end{aligned}

其中在 (i)(i) 中我们使用了第 3.2 节,在 (ii)(ii) 中我们加上并减去了相同的项,在 (iii)(iii) 中我们使用了拉普拉斯算子的定义(式(48)),在 (iv)(iv) 中我们使用了 logpt=ptpt\nabla\log p_t=\frac{\nabla p_t}{p_t},在 (v)(v) 中我们使用了散度算子的线性性。上述推导表明,式(44)中定义的 SDE 满足 ptp_t 的 Fokker-Planck 方程。根据第 4.2 节,这意味着 XtptX_t\sim p_t 对于 0t10\leq t\leq 1 成立,如所愿。

备注 20(可选:Langevin Dynamics)

上述构造有一个著名的特例,即 Probability Path 是恒定的,即对于固定分布 pp,有 pt=pp_t=p。在这种情况下,我们设置 uttarget=0\uref_t=0 并获得 SDE

这通常被称为 Langevin 动力学ptp_t 是恒定的事实意味着 tpt(x)=0\partial_tp_t(x)=0。从第 4.2 节立即得出,这些动力学满足第 4.2 节中静态路径 pt=pp_t=p 的 Fokker-Planck 方程。因此,我们可以得出结论,pp 是 Langevin 动力学的一个平稳分布:

X0pXtp(t0).X_0 \sim p\quad \Rightarrow \quad X_t \sim p\quad (t\geq 0).

与许多 Markov process 一样,这些动力学在相当一般的条件下收敛到平稳分布 pp。也就是说,如果我们改为取 X0ppX_0 \sim p' \neq p,使得 XtptX_t \sim p_t',那么在温和条件下 ptpp_t \to p。这一事实使得 Langevin 动力学极其有用,因此它成为例如 分子动力学 模拟以及贝叶斯统计和自然科学中许多其他 Markov chain 蒙特卡洛(MCMC)方法的基础。特别是,当 pp 是 Gaussian distribution 时,Ornstein-Uhlenbeck 过程作为 Langevin 动力学的特例被恢复,并作为 Diffusion Model 初始公式的基础。

备注 21(可选:GLASS Flows、ODE 随机演化)

SDE 采样的显著特性(与 ODE 相比)在于演化变得随机,即初始点 X0X_0 并不能完全决定 XtX_t 对于 t>0t>0。也许令人惊讶的是,通过一个简单的采样技巧,称为 GLASS Flows [20],也可以纯粹通过 ODE 获得相同的随机转移。这允许利用 SDE 的随机性(例如通过搜索算法),同时保持 ODE 的效率。

仍需展示如何学习 Marginal Score Function logpt(x)\nabla\log p_t(x)。当然,对于 Gaussian Probability Path,我们可以简单地通过第 4.1 节转换 uttarget(x)\uref_t(x)。然而,一般情况下呢?事实证明,我们也可以直接学习 Marginal Score Function。为了近似边际得分 logpt\nabla\log p_t,我们使用一个神经网络,称为 得分网络 stθ:Rd×[0,1]Rds_t^\theta:\mathbb{R}^d\times[0,1]\to\mathbb{R}^d。与之前相同,我们可以设计一个 Score Matching 损失和一个 denoising score matching 损失:

LSM(θ)=EtUnif,zpdata,xpt(z)[stθ(x)logpt(x)2]score matching lossLCSM(θ)=EtUnif,zpdata,xpt(z)[stθ(x)logpt(xz)2]conditional score matching loss\begin{aligned}\mathcal{L}_{\text{SM}}(\theta) &=\mathbb{E}_{t\sim\text{Unif},\,z\sim \pdata,\, x\sim p_t(\cdot|\dap)}\left[\left\|s_t^\theta(x) - \nabla\log p_t(x)\right\|^2\right] &&\blacktriangleright\,\,\text{score matching loss}\\ \mathcal{L}_{\text{CSM}}(\theta) &=\mathbb{E}_{t\sim\text{Unif},\,z\sim \pdata,\, x\sim p_t(\cdot|\dap)}\left[\left\|s_t^\theta(x) - \nabla\log p_t(x|\dap)\right\|^2\right] &&\blacktriangleright\,\,\text{conditional score matching loss}\end{aligned}

其中差异再次在于使用边际得分 logpt(x)\nabla\log p_t(x) 与使用条件得分 logpt(xz)\nabla\log p_t(x|z)
和之前一样,我们理想情况下希望最小化 Score Matching 损失,但无法做到,因为我们不知道 logpt(x)\nabla\log p_t(x)。但同样地,denoising score matching 损失是一个可行的替代方案:

定理 22

Score Matching 损失等于 denoising score matching 损失加上一个常数:

LSM(θ)=LCSM(θ)+C,\mathcal{L}_{\text{SM}}(\theta) = \mathcal{L}_{\text{CSM}}(\theta) + C,

其中 CC 独立于参数 θ\theta。因此,它们的梯度一致:

θLSM(θ)=θLCSM(θ).\nabla_\theta \mathcal{L}_{\text{SM}}(\theta) = \nabla_\theta \mathcal{L}_{\text{CSM}}(\theta).

特别地,对于最小化器 θ\theta^*,将有 stθ=logpts_t^{\theta^*}=\nabla\log p_t

证明

注意 logpt\nabla\log p_t 的公式(式(38))看起来与 uttarget\uref_t 的公式(式(18))相同。因此,证明与第 3.3 节的证明相同,只需将 uttarget\uref_t 替换为 logpt\nabla\log p_t

例 23(Denoising Diffusion Model:Gaussian Probability Path 的 Score Matching)

让我们为 pt(xz)=N(αtz,βt2Id)p_t(x|z)=\mathcal{N}(\alpha_tz,\beta_t^2I_d) 的情况实例化 denoising score matching 损失。正如我们在式(40)中推导的,条件得分 logpt(xz)\nabla\log p_t(x|z) 具有公式

将这一公式代入,条件 Score Matching 损失变为:

LCSM(θ)=EtUnif,zpdata,xpt(z)[stθ(x)+xαtzβt22]=(i)EtUnif,zpdata,ϵN(0,Id)[stθ(αtz+βtϵ)+ϵβt2]=EtUnif,zpdata,ϵN(0,Id)[1βt2βtstθ(αtz+βtϵ)+ϵ2]\begin{aligned}\mathcal{L}_{\text{CSM}}(\theta) &=\mathbb{E}_{t\sim\Unif,\,z\sim \pdata,\,x\sim p_t(\cdot|\dap)}\left[\left\|s_t^\theta(x)+\frac{x-\alpha_t z}{\beta_t^2}\right\|^2\right]\\&\overset{(i)}{=}\mathbb{E}_{t\sim\Unif,\,z\sim \pdata,\,\epsilon\sim \mathcal{N}(0,I_d)}\left[\left\|s_t^\theta(\alpha_tz+\beta_t\epsilon)+\frac{\epsilon}{\beta_t}\right\|^2\right]\\ &=\mathbb{E}_{t\sim\Unif,\,z\sim \pdata,\,\epsilon\sim \mathcal{N}(0,I_d)}\left[\frac{1}{\beta_t^2}\left\|\beta_ts_t^\theta(\alpha_tz+\beta_t\epsilon)+\epsilon\right\|^2\right]\end{aligned}

其中在 (i)(i) 中我们代入了式(28),并将 xx 替换为 αtz+βtϵ\alpha_tz+\beta_t\epsilon。注意,网络 stθs_t^\theta 本质上学习的是预测用于破坏数据样本 zz 的噪声。这解释了为什么上述训练损失被称为 denoising score matching。人们很快意识到,当 βt0\beta_t\approx 0 接近零时(即 denoising score matching 仅在添加足够噪声时才有效),上述损失在数值上不稳定。因此,在关于 Denoising Diffusion Model 的早期工作中(参见 Denoising Diffusion Probabilistic Models,[17]),提出了在损失中丢弃常数 1βt2\frac{1}{\beta_t^2},并通过以下方式将 stθs_t^\theta 重新参数化为 噪声预测器 网络 ϵtθ:Rd×[0,1]Rd\epsilon_t^\theta:\mathbb{R}^d\times[0,1]\to\mathbb{R}^d

βtstθ(x)=ϵtθ(x)LDDPM(θ)=EtUnif,zpdata,ϵN(0,Id)[ϵtθ(αtz+βtϵ)ϵ2]\begin{aligned}-\beta_t s_t^\theta(x) = \epsilon_t^\theta(x)\quad \Rightarrow \quad \mathcal{L}_{\text{DDPM}}(\theta) =&\mathbb{E}_{t\sim\Unif,z\sim \pdata, \epsilon\sim \mathcal{N}(0,I_d)}\left[\|\epsilon_t^\theta(\alpha_tz+\beta_t\epsilon)-\epsilon\|^2\right]\end{aligned}

如前所述,网络 ϵtθ\epsilon_t^\theta 本质上学习的是预测用于破坏数据样本 zz 的噪声。在 Algorithm 4 中,我们总结了训练过程。

让我们总结本节的结果:

小结 24(Score Function、Score Matching、随机采样)

pt(xz),pt(x)p_t(x|z),p_t(x) 为 Conditional Probability Path 和 Marginal Probability Path。Conditional Score Functionlogpt(xz)\nabla\log p_t(x|z) 给出,Marginal Score Functionlogpt(x)\nabla\log p_t(x) 给出。对于每个扩散系数 σt0\sigma_{t}\geq 0,以下 SDE 的轨迹遵循该 Probability Path:

X0pinit,dXt=[uttarget(Xt)+σt22logpt(Xt)]dt+σtdWt(52)\begin{aligned}X_0\sim&\pinit,\quad \dd X_t =\left[\uref_t(X_t)+\frac{\sigma_t^2}{2}\nabla\log p_t(X_t)\right]\dd t +\sigma_t\dd W_t\end{aligned}\tag{52}
Xtpt(0t1),(53)\begin{aligned}\Rightarrow X_t\sim& p_t\quad (0\leq t\leq 1),\end{aligned}\tag{53}

其中 uttarget(x)\uref_t(x) 是如前所述的 Marginal Vector Field(参见式(18))。

Score Matching

为了学习 Marginal Score Function logpt(x)\nabla\log p_t(x),我们可以使用 得分网络 stθs_t^\theta,并通过 denoising score matching 对其进行训练。

LCSM(θ)=Ezpdata,tUnif,xpt(z)[stθ(x)logpt(xz)2](denoising score matching loss)(54)\begin{aligned}\mathcal{L}_{\text{CSM}}(\theta) &= \mathbb{E}_{z\sim \pdata, \,t\sim \text{Unif},\,x\sim p_t(\cdot|\dap)}[\|s_t^\theta(x) - \nabla\log p_t(x|\dap)\|^2] \quad &(\text{denoising score matching loss})\end{aligned}\tag{54}

Gaussian Probability Path

对于最重要的情况——Gaussian Probability Path pt(xz)=N(x;αtz,βt2Id)p_t(x|z)=\mathcal{N}(x;\alpha_tz,\beta_t^2I_d),无需分别训练 stθs_t^\thetautθu_t^\theta,因为我们可以通过以下公式在它们之间进行转换:

utθ(x)=atstθ(x)+btx,at=(βt2α˙tαtβ˙tβt),bt=α˙tαt\begin{aligned}u_t^\theta(x)=&a_t s_t^\theta(x)+b_tx,\quad a_t=\left(\beta_t^2\frac{\dot{\alpha}_t}{\alpha_t}-\dot{\beta}_t\beta_t\right), b_t=\frac{\dot{\alpha}_t}{\alpha_t}\end{aligned}

训练完成后,我们可以模拟以下 SDE。

X0pinit,dXt=[(1+σt22at)utθ(Xt)σt2bt2atXt]dt+σtdWt(55)\begin{aligned}X_0\sim\pinit,\quad \dd X_t =&\left[\left(1+\frac{\sigma_{t}^2}{2a_t}\right)u_t^\theta(X_t)-\frac{\sigma_{t}^2b_t}{2a_t}X_t\right]\dd t +\sigma_t\dd W_t\end{aligned}\tag{55}
=[(at+σt22)stθ(Xt)+btXt]dt+σtdWt(56)\begin{aligned}=&\left[\left(a_t+\frac{\sigma_{t}^2}{2}\right)s_t^\theta(X_t)+b_tX_t\right]\dd t +\sigma_t\dd W_t\end{aligned}\tag{56}

对于任意扩散系数 σt0\sigma_{t}\geq 0,以获得近似样本 X1pdataX_1\sim \pdata。可以通过经验找到最优的 σt0\sigma_t\geq 0

Guidance:如何以 prompt 为条件

到目前为止,我们考虑的 Generative Model 都是无引导的,例如,一个图像模型只会生成某个图像。从数学上讲,这意味着我们的模型从无条件数据分布 pdata(z)\pdata(z) 中返回样本。然而,在大多数情况下,我们的目标不仅仅是生成任意对象,而是生成基于某些附加信息的对象。换句话说,我们想要引导模型生成特定类型的对象。例如,可以想象一个图像 Generative Model,它接收文本 prompt yy,然后生成与文本 prompt yy 匹配的图像 xx。如第 1 节所讨论的,这意味着我们想要从 pdata(zy)\pdata(z|y) 中采样,即基于 yy 的引导数据分布。我们将在本节讨论这一点。

备注 25(术语)

为了避免与使用“条件”一词指代基于 zpdataz \sim \pdata 的条件(Conditional Probability Path/Vector Field)在符号和术语上产生冲突,我们将使用引导一词来特指基于 yy(如文本 prompt)的条件。

Vanilla guidance

首先,我们讨论构建引导 Generative Model 的“标准”方法。简短的答案是:我们只需在训练和推理期间将输入 prompt yy 提供给网络,并以与之前相同的方式完成所有操作。我们在下面将其形式化。我们认为条件变量或 prompt yy 存在于空间 Y\mathcal{Y} 中。例如,当 yy 对应于文本 prompt 时,Y\mathcal{Y} 是所有文本的空间。当 yy 对应于某个离散类标签时,Y\mathcal{Y} 将是离散的。我们对 Y\mathcal{Y} 不施加任何约束。

我们定义一个引导 Diffusion Model,它由一个Guided Vector Field utθ(y)u_t^{\theta}(\cdot | y)(由某个神经网络参数化)和一个时间相关的扩散系数 σt\sigma_t 组成,共同由下式给出

Neural network:uθ:Rd×Y×[0,1]Rd,(x,y,t)utθ(xy)Fixed:σt:[0,1][0,),tσt\begin{aligned}\textbf{\sffamily Neural network:}&\, u^\theta: \mathbb{R}^d \times \mathcal{Y} \times [0,1] \to \mathbb{R}^d,\,\, (x,y,t) \mapsto u_t^{\theta}(x|y)\\ \textbf{\sffamily Fixed:}&\, \sigma_t: [0,1] \to [0,\infty),\,\, t \mapsto \sigma_t\end{aligned}

注意与第 2.2 节摘要的区别:我们额外使用输入 yYy\in \mathcal{Y} 来引导 utθu_t^\theta。对于任何这样的 yYy \in \mathcal{Y},样本可以按如下方式从这样的模型中生成:

Initialization:X0pinitInitialize with simple distribution (such as a Gaussian)Simulation:dXt=utθ(Xty)dt+σtdWtSimulate SDE from t=0 to t=1.Goal:X1pdata(y)Goal is for X1 to be distributed like pdata(y).\begin{aligned}\textbf{\sffamily Initialization:}\quad X_0&\sim\pinit \quad &&\blacktriangleright\,\,\text{Initialize with simple distribution (such as a Gaussian)}\\ \textbf{\sffamily Simulation:}\quad \dd X_t &= u_t^\theta(X_t|y)\dd t + \sigma_t\dd W_t\quad &&\blacktriangleright\,\,\text{Simulate SDE from }t=0\text{ to }t=1\text{.}\\ \textbf{\sffamily Goal:}\quad X_1 &\sim \pdata(\cdot | y) \quad &&\blacktriangleright\,\,\text{Goal is for }X_1\text{ to be distributed like }\pdata(\cdot|y)\text{.}\end{aligned}

σt=0\sigma_t = 0 时,我们说这样的模型是引导 Flow Model。在下文中,为了简洁,我们将自己限制在 Flow Matching 和 Flow Model,但所有内容同样适用于一般情况。

接下来,我们讨论:如何训练引导 Flow Model utθ(xy)u_t^\theta(x|y)?一个简单的技巧可能是固定我们对 yy 的选择,并将我们的数据分布设为 pdata(xy)p_{\text{data}}(x|y)。然后我们恢复了之前的无引导生成问题,并可以相应地使用 conditional flow matching 目标构建 Generative Model,即:

Ezpdata(y),xpt(z)utθ(xy)uttarget(xz)2.(57)\mathbb{E}_{z \sim p_{\text{data}}(\cdot|y), x \sim p_t(\cdot|z)} \lVert u_t^{\theta}(x|y) - \uref_t(x|z)\rVert^2. \tag{57}

注意,标签 yy 不影响 Conditional Probability Path pt(z)p_t(\cdot|z) 或 Conditional Vector Field uttarget(xz)\uref_t(x|z)(尽管原则上,我们可以使其依赖)。
展开对所有此类 yy 选择的期望,我们因此获得一个引导 conditional flow matching 目标

LCFMguided(θ)=E(z,y)pdata(z,y),tUnif[0,1],xpt(z)utθ(xy)uttarget(xz)2.(58)\mathcal{L}_{\text{CFM}}^{\text{guided}}(\theta) = \mathbb{E}_{(z,y) \sim p_{\text{data}}(z,y),\,t\sim \text{Unif}[0,1],\,x\sim p_t(\cdot|z)} \lVert u_t^{\theta}(x|y) - \uref_t(x|z)\rVert^2. \tag{58}

式(58)中的引导目标与式(26)中的无引导目标之间的主要区别之一是,这里我们采样 (z,y)pdata(z,y) \sim \pdata 而不仅仅是 zpdataz \sim \pdata。原因是我们的数据分布现在原则上是一个联合分布,例如,同时覆盖图像 zz 和文本提示 yy。在实践中,这意味着式(58)的 PyTorch 实现将涉及一个数据加载器,该加载器返回同时包含 zzyy 的批次。

Figure 11使用 prompt/类 y=y=“柯基犬”进行图像生成。左:使用普通引导生成的样本——图像与 prompt 匹配不佳。右:使用 Classifier Guidance 和 w=4w = 4 生成的样本。如图所示,classifier-free guidance 提高了对 prompt 的遵循程度。图取自 [18]。

Classifier-free guidance

理论上,普通的引导应当能够产生对pdata(y)\pdata(\cdot|y) 的忠实生成过程。然而,人们很快在实验中发现,使用该过程生成的图像样本并不能很好地符合期望的标签yy(见图 11)。这可能有多种原因:模型可能欠拟合(即我们并未真正学到真实的 Marginal Vector Field),或者我们的数据可能不完美(例如,来自万维网的文本-图像对包含许多错误)。因此,为了真正生成更符合 prompt 的样本,我们必须找到一种方法来人为地强化 prompt 变量yy。实现这一目标的主要技术被称为classifier-free guidance,它在最先进的 Diffusion Model 中被广泛使用,我们接下来将讨论它。

Classifier Guidance

Figure 12分类器与 classifier-free guidance 的示意图。Classifier Guidance 将 Guided Vector Fielduttarget(xy)\uref_t(x|y) 与分类器的梯度logpt(yx)\log p_t(y|x) 分解,并用 guidance scale w>1w>1 放大分类器。classifier-free guidance 放大两个 Vector Field 之间的差异,从而在不训练单独分类器模型的情况下达到相同的效果。

为简单起见,我们这里将重点放在 Gaussian Probability Path 的情形。回顾式(15),Gaussian Conditional Probability Path 由pt(z)=N(αtz,βt2Id)p_t(\cdot|\dap) = \mathcal{N}(\alpha_t \dap,\beta_t^2 I_d) 给出,其中 noise schedulesαt\alpha_tβt\beta_t 是连续可微的、单调的,并且满足α0=β1=0\alpha_0 = \beta_1 = 0α1=β0=1\alpha_1 = \beta_0 = 1。此外,回顾我们可以使用第 4.1 节将 Guided Vector Fielduttarget(xy)\uref_t(x|y) 重写为以下形式,使用引导 Score Functionlogpt(xy)\nabla\log p_t(x|y)

uttarget(xy)=atlogpt(xy)+btx,(59)\uref_t(x|y) = a_t\nabla \log p_t(x|y)+b_tx, \tag{59}

接下来,注意到pt(xy)p_t(x|y) 是一个条件密度。因此,我们可以使用 Bayes' rule 将引导得分重写为

pt(xy)=pt(x)pt(yx)pt(y)(60)\begin{aligned}p_t(x|y)=&\frac{p_t(x)p_t(y|x)}{p_t(y)}\end{aligned}\tag{60}
logpt(xy)=log(pt(x)pt(yx)pt(y))=logpt(x)+logpt(yx),(61)\begin{aligned}\nabla \log p_t(x|y) =& \nabla \log \left(\frac{p_t(x)p_t(y|x)}{p_t(y)}\right) = \nabla \log p_t(x) + \nabla \log p_t(y|x),\end{aligned}\tag{61}

其中我们使用了梯度\nabla 是关于变量xx 取的,因此logpt(y)=0\nabla \log p_t(y) = 0。因此我们可以重写

uttarget(xy)=btx+at(logpt(x)+logpt(yx))=uttarget(x)+atlogpt(yx).\uref_t(x|y) = b_tx + a_t(\nabla \log p_t(x) + \nabla \log p_t(y|x)) = \uref_t(x) + a_t \nabla \log p_t(y|x).

注意上述等式的形式:Guided Vector Fielduttarget(xy)\uref_t(x|y) 是 Unguided Vector Fielduttarget(x)\uref_t(x) 加上引导变量yy 的似然pt(yx)p_t(y|x) 的梯度的和。由于人们观察到他们的图像xx 不能很好地符合他们的 prompt yy,一个自然的想法是放大logpt(yx)\nabla \log p_t(y|x) 项的贡献,从而得到

u~t(xy)=uttarget(x)+watlogpt(yx),(classifier guidance)(62)\begin{aligned}\tilde{u}_t(x|y) = \uref_t(x) + w a_t \nabla \log p_t(y|x),\quad &(\text{classifier guidance})\end{aligned}\tag{62}

其中w>1w > 1 被称为guidance scale。我们如何学习项logpt(yx)\log p_t(y|x)?注意,这可以被视为一种对噪声数据的分类器(即,它给出给定xxyy 的对数似然)。因此,我们可以简单地通过监督学习来学习它。这导致了Classifier Guidance[11, 43](见图 12 的示意图)。Classifier Guidance 在很大程度上被 classifier-free guidance 所取代,因此我们在此不再进一步讨论。然而,它构成了 classifier-free guidance 的基础,我们将在接下来看到。最后,注意这是一个启发式方法:对于w1w \neq 1,有u~t(xy)uttarget(xy)\tilde{u}_t(x|y) \neq \uref_t(x|y),即因此不是“真正的”Guided Vector Field。

classifier-free guidance

虽然 Classifier Guidance 原则上可行,但它带来了一些困难:首先,我们需要在流/Diffusion Model 旁边训练一个分类器——因此我们有 2 个网络而不是 1 个。此外,如果yy 是高维的,例如文本 prompt 而不仅仅是一个类别,那么pt(yx)p_t(y|x) 可能很难学习,梯度logpt(yx)\nabla\log p_t(y|x) 也很难获得。出于这个原因,classifier-free guidance[18] 被引入。classifier-free guidance 在理论上产生与 Classifier Guidance 等效的效果,但无需训练单独的分类器。

为此,我们可以再次应用等式

logpt(xy)=logpt(x)+logpt(yx)\nabla \log p_t(x|y) = \nabla \log p_t(x) + \nabla \log p_t(y|x)

得到

u~t(xy)=uttarget(x)+watlogpt(yx)=uttarget(x)+wat(logpt(xy)logpt(x))=uttarget(x)(wbtx+watlogpt(x))+(wbtx+watlogpt(xy))=(1w)uttarget(x)+wuttarget(xy).\begin{aligned}\tilde{u}_t(x|y) &= \uref_t(x) + w a_t \nabla \log p_t(y|x)\\ &= \uref_t(x) + w a_t (\nabla \log p_t(x|y) - \nabla \log p_t(x))\\ &= \uref_t(x) - (w b_tx + w a_t \nabla \log p_t(x)) + (w b_t x + w a_t \nabla \log p_t(x|y))\\ &= (1-w) \uref_t(x) + w \uref_t(x|y).\end{aligned}

因此,我们可以将缩放后的 Guided Vector Field u~t(xy)\tilde{u}_t(x|y) 表示为 Unguided Vector Field uttarget(x)\uref_t(x) 与 Guided Vector Field uttarget(xy)\uref_t(x|y) 的线性组合。其思路可能是同时训练一个未引导的 uttarget(x)\uref_t(x)(例如使用式(26))和一个引导的 uttarget(xy)\uref_t(x|y)(例如使用式(58)),然后在推理时将它们组合以获得 u~t(xy)\tilde{u}_t(x|y)。"但是等等!",你可能会问,"那我们岂不是需要训练两个模型?!"。事实证明,我们可以在一个模型中训练两者:我们可以用一个新的、额外的 \varnothing 标签来扩充我们的标签集,该标签表示不存在条件。然后我们可以将 uttarget(x)=uttarget(x)\uref_t(x)=\uref_t(x|\varnothing) 视为。这样,我们就不需要训练一个单独的模型来强化假设分类器的影响。这种在一个模型中训练条件模型和无条件模型(并随后强化条件)的方法被称为classifier-free guidance(CFG)[18](见图 12 的图示)。

备注 26(一般 Probability Path 的推导)

注意,该构造

u~t(xy)=(1w)uttarget(x)+wuttarget(xy),\tilde{u}_t(x|y) = (1-w) \uref_t(x) + w \uref_t(x|y),

对于任何选择 Probability Path 同样有效,而不仅仅是 Gaussian 路径。当 w=1w=1 时,很容易验证 u~t(xy)=uttarget(xy)\tilde{u}_t(x|y)=\uref_t(x|y)。我们使用 Gaussian 路径的推导只是为了说明该构造背后的直觉,特别是放大假设的"分类器" logpt(yx)\nabla \log p_t(y|x) 的贡献。

训练与 classifier-free guidance

我们现在必须修改式(58)中的引导 conditional flow matching 目标,以考虑 y=y = \varnothing 的可能性。挑战在于,当采样 (z,y)pdata(z,y) \sim \pdata 时,我们永远不会得到 y=y = \varnothing。因此,我们必须人为引入 y=y = \varnothing 的可能性。为此,我们将定义某个超参数 η\eta 作为我们丢弃原始标签 yy 并将其替换为 \varnothing 的概率。因此,我们得到了CFG conditional flow matching 训练目标

LCFMCFG(θ)=Eutθ(xy)uttarget(xz)2(63)\begin{aligned}\mathcal{L}_{\text{CFM}}^{\text{CFG}}(\theta) &= \,\,\mathbb{E}_{\square} \lVert u_t^{\theta}(x|y) - \uref_t(x|z)\rVert^2\end{aligned}\tag{63}
=(z,y)pdata(z,y),tUnif[0,1],xpt(z),replace y= with prob. η(64)\begin{aligned}\square &= (z,y) \sim p_{\text{data}}(z,y),\, t \sim \text{Unif}[0,1],\, x \sim p_t(\cdot|z),\text{replace }y=\varnothing\text{ with prob. }\eta\end{aligned}\tag{64}
Algorithm 5
算法 5 · Gaussian Probability Path pt(xz)=N(x;αtz,βt2Id)p_t(x\mid z)=\mathcal{N}(x;\alpha_tz,\beta_t^2I_d) 的 classifier-free guidance 训练
输入: 配对数据集 (z,y)pdata(z,y)\sim\pdata,neural-network Vector Field utθu_t^\theta

- 循环: 遍历每个数据 mini-batch
- 从数据集中采样一个样本对 (z,y)(z,y)
- 采样随机时间 tUnif[0,1]t\sim\mathrm{Unif}[0,1]
- 采样噪声 ϵN(0,Id)\epsilon\sim\mathcal{N}(0,I_d)
- 设置 x=αtz+βtϵx=\alpha_tz+\beta_t\epsilon
- 以概率 pp 丢弃标签:yy\leftarrow\varnothing
- 计算损失:
L(θ)=utθ(xy)(α˙tz+β˙tϵ)2\mathcal{L}(\theta) =\left\lVert u_t^\theta(x\mid y)-(\dot\alpha_tz+\dot\beta_t\epsilon)\right\rVert^2
- 对 L(θ)\mathcal{L}(\theta) 做梯度下降,更新模型参数 θ\theta
- 结束循环

我们在下面总结我们的发现。

小结 27(Flow Model 为 classifier-free guidance)

给定未引导的 Marginal Vector Field uttarget(x)\uref_t(x|\varnothing)、引导的 Marginal Vector Field uttarget(xy)\uref_t(x|y) 以及一个 guidance scale w>1w > 1,我们定义Classifier-Free GuidanceVector Field u~t(xy)\tilde{u}_t(x|y)

u~t(xy)=(1w)uttarget(x)+wuttarget(xy).(65)\tilde{u}_t(x|y) = (1-w) \uref_t(x|\varnothing) + w \uref_t(x|y). \tag{65}

通过使用同一个神经网络近似 uttarget(x)\uref_t(x|\varnothing)uttarget(xy)\uref_t(x|y),我们可以利用以下classifier-free guidance CFM(CFG-CFM)目标,由下式给出

LCFMCFG(θ)=Eutθ(xy)uttarget(xz)2(66)\begin{aligned}\mathcal{L}_{\text{CFM}}^{\text{CFG}}(\theta) &= \,\,\mathbb{E}_{\square} \lVert u_t^{\theta}(x|y) - \uref_t(x|z)\rVert^2\end{aligned}\tag{66}
=(z,y)pdata(z,y),tUnif[0,1],xpt(z),replace y= with prob. η(67)\begin{aligned}\square &= (z,y) \sim p_{\text{data}}(z,y),\, t \sim \text{Unif}[0,1],\, x \sim p_t(\cdot|z),\text{replace }y=\varnothing\text{ with prob. }\eta\end{aligned}\tag{67}

用通俗的话说,LCFMCFG\mathcal{L}_{\text{CFM}}^{\text{CFG}} 可能被近似为

(z,y)pdata(z,y)Sample (z,y) from data distribution.tUnif[0,1)Sample t uniformly on [0,1).xpt(xz)Sample x from the conditional probability path pt(xz).with prob.η,yReplace y with  with probability η.LCFMCFG(θ)^=utθ(xy)uttarget(xz)2Regress model against conditional vector field.\begin{aligned}(z,y) &\sim \pdata(z,y) \quad\quad\quad\quad && \blacktriangleright \quad \text{Sample }(z,y)\text{ from data distribution.}\\ t &\sim \text{Unif}[0,1) \quad\quad\quad\quad && \blacktriangleright \quad \text{Sample }t\text{ uniformly on }[0,1)\text{.}\\ x &\sim p_t(x|z) \quad\quad\quad\quad && \blacktriangleright \quad \text{Sample }x\text{ from the conditional probability path }p_t(x|z)\text{.}\\ \text{with prob.}&\,\eta,\, y \gets \varnothing \quad\quad\quad\quad && \blacktriangleright \quad \text{Replace }y\text{ with }\varnothing\text{ with probability }\eta\text{.}\\ \widehat{\mathcal{L}_{\text{CFM}}^{\text{CFG}}(\theta)} &= \lVert u_t^{\theta}(x|y) - \uref_t(x|z)\rVert^2 \quad\quad\quad\quad && \blacktriangleright \quad \text{Regress model against conditional vector field.}\end{aligned}

在推理时,对于固定的 yy 选择,我们可以通过以下方式采样

Initialization:X0pinit(x)Initialize with simple distribution (such as a Gaussian)Simulation:dXt=u~tθ(Xty)dtSimulate ODE from t=0 to t=1.Samples:X1Goal is for X1 to adhere to the guiding variable y.\begin{aligned}\textbf{\sffamily Initialization:}\quad X_0&\sim\pinit(x) \quad && \blacktriangleright\,\,\text{Initialize with simple distribution (such as a Gaussian)}\\ \textbf{\sffamily Simulation:}\quad \dd X_t &= \tilde{u}_t^\theta(X_t|y)\dd t \quad && \blacktriangleright\,\,\text{Simulate ODE from }t=0\text{ to }t=1\text{.}\\ \textbf{\sffamily Samples:}\quad X_1& \quad && \blacktriangleright\,\,\text{Goal is for }X_1\text{ to adhere to the guiding variable }y\text{.}\end{aligned}

注意,如果我们使用权重 w>1w>1X1X_1 的分布不再必然与 X1pdata(y)X_1 \sim \pdata(\cdot | y) 对齐。然而,经验上,这显示出与条件更好的对齐。因此,classifier-free guidance 是一个启发式方法,主要因其出色的实证结果而被证明合理。事实上,几乎所有你看到的 AI 生成的图像或视频都严重依赖 classifier-free guidance w4w\geq 4。在图 11 中,我们展示了在 128x128 ImageNet 上基于类别的 classifier-free guidance,如 [18] 中所述。类似地,在 Figure 13 中,我们可视化了在 MNIST 手写数字 dataset 上应用 classifier-free guidance 时各种引导尺度 ww 的影响。

Figure 13在 MNIST 手写数字 dataset 上以不同引导尺度应用 classifier-free guidance 的效果。左:guidance scale 设置为 w=1.0w = 1.0。中:guidance scale 设置为 w=2.0w = 2.0。右:guidance scale 设置为 w=4.0w = 4.0。你将在第三次实验室中自己生成类似的图像!
备注 28(Diffusion Model 指南)

将讨论从 flow 模型扩展到 Diffusion Model 是直接的。只需将 utθ(xy)u_t^\theta(x|y) 替换为 u~tθ(xy)\tilde{u}_t^\theta(x|y),并使用第 4 节中讨论的 SDE 进行采样。

构建大规模图像或视频生成器

在前面的章节中,我们学习了如何训练 Flow Matching 或 Diffusion Model 以从分布 pdata(xy)\pdata(x|y) 中采样。这个配方是通用的,可以应用于各种数据类型和应用。在本节中,我们深入探讨大规模图像和视频生成的特定案例,包括知名模型,如 FLUX 2.0, Stable Diffusion 3, Nano BananaVEO-3 或 Meta Movie Gen Video。最后,我们将在实验室中应用目前所学,从零开始构建我们自己的此类模型!本节大致安排如下:

  • 神经网络架构: 我们首先讨论原始条件输入,包括时间 tt 和引导变量 yrawy_{\text{raw}}(即离散类别标签或原始文本),如何被转换或embedding为模型 utθ(xy)u_t^\theta(x|y) 本身可消化的向量值形式。然后我们讨论 utθ(xy)u_t^\theta(x|y) 的流行架构选择,包括 U-NetDiffusion Transformer
  • latent space: 我们讨论Variational Autoencoder,它允许在较低维度的 latent space 中进行生成建模,从而实现超高分辨率图像生成。
  • 案例研究: 最后,我们将深入检查上述两个最先进的图像和视频模型 - Stable DiffusionMeta MovieGen - 让你了解大规模实践的方式。

神经网络架构

让我们首先将 attention 转向针对图像类模态(例如图像和视频)的 flow 和 Diffusion Model 的可扩展神经网络架构设计。具体来说,我们将探讨(引导的)Vector Field utθ(xy)u_t^\theta(x|y) 与参数 θ\theta 的任务在实践中如何实现。注意,神经网络必须有 3 个输入:向量 xRdx\in\R^d、条件变量 yYy\in\mathcal{Y} 和时间值 t[0,1]t\in [0,1],以及一个输出,向量 utθ(xy)Rdu_t^\theta(x|y)\in\mathbb{R}^d。对于低维分布(例如我们在前面章节中看到的玩具分布),将 utθ(xy)u_t^\theta(x|y) 参数化为多层感知器(MLP),也称为全连接神经网络,就足够了。也就是说,在这种简单设置中,通过 utθ(xy)u_t^\theta(x|y) 的前向传播将涉及将我们的输入 xxyytt 连接起来,并通过 MLP 传递。然而,对于复杂的高维分布,例如图像、视频和蛋白质上的分布,MLP 可能不够,通常使用特殊的、特定于应用的架构。在本小节的其余部分,我们将考虑图像(以及扩展的视频)的情况。首先,我们将考虑原始条件信息 - 时间 tt 和条件变量 yy - 如何被embedding为实际模型可消化的向量值形式。其次,我们将考虑此类模型的两种常见架构选择:U-Net [38, 17, 22, 11] 和 Diffusion Transformer(DiT)[12, 30, 28]。

条件变量的 embedding

时间 embedding

对于简单的玩具模型,将 tt 的原始值连接到输入就足以训练一个性能合理的网络。在实践中,标量时间通常使用傅里叶特征embedding 到更高维空间,使模型能够更忠实地捕捉高频时间依赖性 [46]。明确地,特征化由下式给出

TimeEmb(t)=2d[cos(2πw1t)cos(2πwd/2t)sin(2πw1t)sin(2πwd/2t)]T,(68)\begin{aligned}\text{TimeEmb}(t) = \sqrt{\frac{2}{d}}\begin{bmatrix} \cos(2\pi w_1 t) & \cdots & \cos(2\pi w_{d/2} t) & \sin(2\pi w_1 t) & \cdots & \sin(2\pi w_{d/2} t) \end{bmatrix}^T,\end{aligned}\tag{68}

其中频率 wiw_i 按以下方式设置

wi  =  wmin(wmaxwmin)i1d/21,i=1,,d/2.(69)w_i \;=\; w_{\min}\left(\frac{w_{\max}}{w_{\min}}\right)^{\frac{i-1}{d/2-1}}, \qquad i=1,\ldots,d/2. \tag{69}

这种 TimeEmb\text{TimeEmb} 的选择是标准选择,但确切形式并非严格必要。相反,上述只是获得维度 dd 的归一化 embedding 的便捷方式,即 TimeEmb(t)=1\|\text{TimeEmb}(t)\|=1(因为 sin2+cos2=1\sin^2+\cos^2=1)。

类别标签 embedding

yrawY{0,,N}y_{\text{raw}} \in \mathcal{Y} \triangleq \{0,\dots, N\} 只是一个类别标签时,通常最简单的方法是直接为 yrawy_{\text{raw}}N+1N+1 个可能值分别学习一个 embedding 向量,并将 yy 设置为该 embedding 向量。可以将这些 embedding 的参数视为包含在 utθ(xy)u_t^\theta(x|y) 的参数中,因此在训练过程中会学习这些参数。

embedding 文本输入

yrawy_{\text{raw}} 是文本 prompt 时,情况更为复杂,方法主要依赖于冻结的预训练模型。这类模型经过训练,可以将离散的文本输入 embedding 到捕获相关信息的连续向量中。其中一个模型被称为 CLIP(对比语言-图像预训练)。CLIP 通过训练学习一个图像和文本提示共享的 embedding 空间,其训练损失旨在鼓励图像 embedding 与其对应的提示接近,同时与其他图像和提示的 embedding 距离更远 [34]。因此,我们可以将 y=CLIP(yraw)RdCLIPy = \text{CLIP}(y_{\text{raw}}) \in \mathbb{R}^{d_{\text{CLIP}}} 设为冻结的预训练 CLIP 模型产生的 embedding。在某些情况下,将整个序列压缩为单个表示可能并不理想。此时,还可以考虑使用预训练的 Transformer 对 prompt 进行 embedding,以获得 embedding 序列。在条件化时,也常见将多个此类预训练 embedding 组合起来,以同时获得每个模型的好处 [14, 33]。就我们的目的而言,可以简单地假设在应用这样的模型后,prompt embedding 的形状为

PromptEmbed(yraw)RS×k\text{PromptEmbed}(y_{\text{raw}})\in \mathbb{R}^{S\times k}

Diffusion Transformer

在深入探讨这些架构的具体细节之前,让我们回顾一下引言中的内容:图像本质上是一个向量 xRCimage×H×Wx \in \mathbb{R}^{C_{\text{image}} \times H \times W}。这里 CimageC_{\text{image}} 表示 通道 的数量(RGB 图像通常有 Cinput=3C_{\text{input}} = 3 个颜色通道),HHWW 分别表示图像的 高度宽度(以像素为单位)。一个特别突出的架构类别是所谓的 扩散 Transformer(DiT)及其变体,它们使用 attention 机制来构建网络 [49, 30, 28]。扩散 Transformer 有不同的风格。我们在这里解释一种通用设计,但请注意,DiT 的具体实例可能因模型和应用而异。在本节的其余部分,我们将使用 dd 表示隐藏维度,LL 表示 Transformer 层的数量,hh 表示每层的头数。扩散 Transformer 基于 视觉 Transformer(ViT),其主要思想基本上是将图像划分为补丁,embedding 补丁以获得 token 序列,并通过标准的 attention 处理生成的 token [12]。最后应用一个去补丁化操作,以恢复正确形状的图像。初始的补丁化操作只是对图像张量 xRC×H×Wx\in \mathbb{R}^{C\times H\times W} 的重构:

Patchify(x)RN×C\text{Patchify}(x)\in \mathbb{R}^{N\times C'}
Figure 14左:Diffusion Transformer 架构的概述,取自 [30]。右:对比 CLIP 损失的示意图,其中学习了一个共享的图像-文本 embedding 空间,取自 [34]。

其中 C=CP2,N=(H/P)(W/P)C'=CP^2, N=(H/P)\cdot(W/P) 对于 PP 是 patch 大小。接下来,我们对输出应用线性变换,得到最终的 patch embedding

PatchEmb(x)=Patchify(x)WRN×d\text{PatchEmb}(x)=\text{Patchify}(x)W \in \mathbb{R}^{N\times d}

其中 WRC×dW\in \mathbb{R}^{C'\times d} 是一个可学习的权重矩阵。Diffusion Transformer 的输入包括时间 embedding、prompt embedding 以及补丁化的图像张量(见第 6.1.1 节):

t~=TimeEmb(t)Rdy~=PromptEmb(y)RS×dx~0=PatchEmb(x)RN×d\begin{aligned}\tilde{t} &=\text{TimeEmb}(t)\in \mathbb{R}^{d}\\ \tilde{y}&=\text{PromptEmb}(y)\in\mathbb{R}^{S \times d}\\ \tilde{x}_{0}&=\text{PatchEmb}(x)\in\mathbb{R}^{N\times d}\end{aligned}

注意,所有元素现在都具有所需的 Transformer 隐藏维度。Diffusion Transformer 然后通过 Transformer 层在 DiT Block 中迭代更新 z~i\tilde{z}_i,对于 i=0,,L1i=0,\cdots,L-1(详见备注 29):

x~i+1=DiTBlock(x~i,t~,y~)RN×d(i=0,,L1).(70)\tilde{x}_{i+1}=\text{DiTBlock}(\tilde{x}_i,\tilde{t},\tilde{y})\in\mathbb{R}^{N\times d}\quad (i=0,\dots,L-1). \tag{70}

其中 NN 是层数。
最后,一个最终操作应用去补丁化操作,将 DiT 的输出映射回所需的输出形状:

u=Depatchify(x~NW~)RC×H×W,u=\text{Depatchify}(\tilde{x}_N\tilde{W})\in \mathbb{R}^{C\times H\times W},

其中 W~Rd×C\tilde{W}\in \mathbb{R}^{d\times C'}。最终张量 uu 作为模型的输出,即预测的速度 utθ(xy)u_t^\theta(x|y)

备注 29(DiT 块)

为完整起见,我们给出单个 DiT 层的简要数学描述。虽然我们试图包含足够的细节以帮助读者大致理解 DiT 模型家族,但我们提醒读者,这些选择旨在强调关键算法选择而非架构细节。现在,设xRN×dx\in\mathbb{R}^{N\times d} 表示当前的 patch 令牌序列(此处x=x~ix=\tilde x_i),并设yRS×dy\in\mathbb{R}^{S \times d} 表示 embedding 的引导变量(此处y=y~y=\tilde y)。那么,一个典型的 DiT 块通过以下方式更新xx:(i)对补丁进行自 attention,(ii)与 prompt 进行交叉 attention,以及(iii)通过自适应归一化(AdaLN)进行时间条件化。

缩放点积 attention

给定查询QRN×dhQ\in\mathbb{R}^{N\times d_h}、键KRM×dhK\in\mathbb{R}^{M\times d_h} 和值VRM×dhV\in\mathbb{R}^{M\times d_h}

Attn(Q,K,V)  =  softmax ⁣(QKdh)V    RN×dh,\mathrm{Attn}(Q,K,V) \;=\; \mathrm{softmax}\!\left(\frac{QK^\top}{\sqrt{d_h}}\right)V \;\in\;\mathbb{R}^{N\times d_h},

其中 softmax 按行应用。

多头 attention

hh 表示头数,dh=dhd_h=\frac{d}{h} 表示每头维度。对于每个头h{1,,nheads}h\in\{1,\dots,n_{\text{heads}}\},学习投影矩阵WQ(h),WK(h),WV(h)Rk×dhW_Q^{(h)},W_K^{(h)},W_V^{(h)}\in\mathbb{R}^{k\times d_h}。定义

headh(x,z)  =  Attn ⁣(xWQ(h),zWK(h),zWV(h)),\text{head}_h(x,z) \;=\; \mathrm{Attn}\!\big(xW_Q^{(h)},\,zW_K^{(h)},\,zW_V^{(h)}\big),

其中源序列zz 要么是

z=x(self-attention on patches),z=y(cross-attention to the prompt).z=x \quad \text{(self-attention on patches)},\qquad z=y \quad \text{(cross-attention to the prompt)}.

拼接各头并应用输出投影WORd×dW_O\in\mathbb{R}^{d\times d}

MultiHeadattention(x,z)  =  Concat(head1(x,z),,headh(x,z))WO    RN×d.\mathrm{MultiHeadattention}(x,z) \;=\; \mathrm{Concat}\big(\text{head}_1(x,z),\dots,\text{head}_{h}(x,z)\big)\,W_O \;\in\;\mathbb{R}^{N\times d}.

通过自适应归一化进行时间条件化

t~Rd\tilde t\in\mathbb{R}^d 为 timestep embedding。在 DiTs 中,标准选择是使用t~\tilde t 来生成逐通道的缩放/平移参数,以调制归一化激活[31]。具体来说,设g:RdR2dg:\mathbb{R}^d\to\mathbb{R}^{2d} 为 MLP,并设置

(γ,β)=g(t~),(\gamma,\beta) = g(\tilde t),

其中γ,βRd\gamma,\beta\in\mathbb{R}^{d}(或根据实现,针对不同子层如 attention 和 MLP 使用单独的(γ,β)(\gamma,\beta) 对)。给定 token 矩阵xRN×dx\in\mathbb{R}^{N\times d} 和归一化算子Norm()\mathrm{Norm}(\cdot)(例如 LayerNorm),定义调制归一化

AdaNormt~(x)  =  (1+γ)Norm(H)  +  β,\mathrm{AdaNorm}_{\tilde t}(x) \;=\; \bigl(1+\gamma\bigr)\odot \mathrm{Norm}(H) \;+\; \beta,

其中\odot 表示逐元素乘法,并在 token 维度上进行广播。

综合起来

组合操作,因此也就是 DiT Block,由下式给出。

xx+gself(t~)MultiHeadattention ⁣(AdaNormt~(x),AdaNormt~(x))xx+gcross(t~)MultiHeadattention ⁣(AdaNormt~(x),y)xx+gMLP(t~)MLP ⁣(AdaNormt~(x)),\begin{aligned}x &\gets x + g_\text{self}(\tilde{t})\odot \mathrm{MultiHeadattention}\!\big(\mathrm{AdaNorm}_{\tilde t}(x),\,\mathrm{AdaNorm}_{\tilde t}(x)\big)\\ x &\gets x + g_\text{cross}(\tilde{t})\mathrm{MultiHeadattention}\!\big(\mathrm{AdaNorm}_{\tilde t}(x),\,y\big)\\ x &\gets x + g_\text{MLP}(\tilde{t})\mathrm{MLP}\!\big(\mathrm{AdaNorm}_{\tilde t}(x)\big),\end{aligned}

其中 MLP 是逐位置的(position-wise)前馈网络,gg_{\cdots} 是可学习的门控参数。输出 xRN×dx\in \RR^{N\times d} 成为下一层的 patch-token 序列(在我们的记号中,x~i+1\tilde x_{i+1})。最后,我们注意到类别条件化的 DiT(例如实验室中实现的那种)通常更简单,并且避免使用交叉 attention 层,而倾向于使用基于时间和类别的 AdaNorm 条件化。

U-Net

U-Net 架构 [38] 是 DiT 架构的一种替代架构,并且是一种特定类型的卷积神经网络。它最初是为图像分割设计的,其关键特征是它的输入和输出都具有图像的形状(可能具有不同数量的通道)。这使得它非常适合参数化 Vector Field xutθ(xy)x\mapsto u_t^\theta(x|y),因为对于固定的 y,ty,t,其输入具有图像的形状,其输出也是如此。因此,U-Net 在早期关于 Diffusion Model 的大量文献中得到了广泛使用 [17, 22, 11]。一个 U-Net 由一系列编码器 Ei\mathcal{E}_i 和相应的解码器序列 Di\mathcal{D}_i 组成,中间还有一个 latent 处理块,我们将其称为中编码器(midcoder)。〔脚注〕 举例来说,让我们跟随一张图像 xtR3×256×256x_t \in \mathbb{R}^{3 \times 256 \times 256}(我们取 (Cinput,H,W)=(3,256,256)(C_{\text{input}}, H, W) = (3, 256, 256))在被 U-Net 处理时所经过的路径:

xtinputR3×256×256Input to the U-Net.xtlatent=E(xtinput)R512×32×32Pass through encoders to obtain latent.xtlatent=M(xtlatent)R512×32×32Pass latent through midcoder.xtoutput=D(xtlatent)R3×256×256Pass through decoders to obtain output.\begin{aligned}x^{\text{input}}_t &\in \mathbb{R}^{3 \times 256 \times 256} \quad && \blacktriangleright\,\,\text{Input to the U-Net.}\\ x^{\text{latent}}_t = \mathcal{E}(x^{\text{input}}_t) &\in \mathbb{R}^{512 \times 32 \times 32} \quad && \blacktriangleright\,\,\text{Pass through encoders to obtain latent.}\\ x^{\text{latent}}_t = \mathcal{M}(x^{\text{latent}}_t) &\in \mathbb{R}^{512 \times 32 \times 32} \quad && \blacktriangleright\,\,\text{Pass latent through midcoder.}\\ x^{\text{output}}_t = \mathcal{D}(x^{\text{latent}}_t) &\in \mathbb{R}^{3 \times 256 \times 256} \quad && \blacktriangleright\,\,\text{Pass through decoders to obtain output.}\end{aligned}

注意,当输入经过编码器时,其表示中的通道数量增加,而图像的高度和宽度减小。编码器和解码器通常都由一系列卷积层组成(中间有激活函数、池化操作等)。上面未显示两点:首先,输入 xtinputR3×256×256x^{\text{input}}_t\in \mathbb{R}^{3 \times 256 \times 256} 通常被馈送到一个初始的预编码块中,以在馈送到第一个编码器块之前增加通道数量。其次,编码器和解码器通常通过残差连接连接。完整的图示见 Figure 15。

在高层次上,大多数 U-Net 都涉及上述某种变体。然而,上述某些设计选择可能与实际中的各种实现有所不同。特别是,我们上面选择了一种纯卷积架构,而在编码器和解码器中通常也会包含 attention 层。U-Net 的名字来源于其编码器和解码器形成的“U”形(见图 15)。

Figure 15一个简化的 U-Net 架构(本课程 2025 版本的实验 03 中使用了这样的架构)。

在 latent space 中工作:(变分)Autoencoder

到目前为止,我们一直在数据空间 Rd\RR^d 中操作。然而,随着我们扩展到越来越高分辨率的图像,直接在这样的空间中建模的成本很快就会变得高得令人望而却步。例如,一个 1024×10241024 \times 1024 图像,具有三个 RGB 颜色通道,对应的总维度为 d=HW33106d=H\cdot W\cdot 3\approx3*10^6!请注意,对于视频,维度会进一步增加,因为一切都随帧数 TT 缩放。可以想象,在这样的空间上进行训练很快就会变得不可行。与图像分类不同,图像分类的低维输出允许缩小卷积堆栈,而我们基于流(flow)的建模方法要求我们的输出 utθ(x)Rdu_t^\theta(x)\in\mathbb{R}^d 与输入一样大。因此,重要的问题变成了:我们如何在合理的内存和计算预算内对高维图像进行建模?

Standard Autoencoder

这个问题的一个自然答案在于压缩:例如,图像的实际空间可能位于高维图像空间的某个低维流形附近。更具体地说,我们可以考虑一个编码器 μϕ:RdRk\mu_{\phi}: \mathbb{R}^d \to \mathbb{R}^k 和一个解码器 μθ:RkRd\mu_{\theta}: \mathbb{R}^k \to \mathbb{R}^d,它们分别将原始图像 xRdx \in \RR^d 映射到潜在表示(latents)zRkz\in\RR^k 以及从潜在表示映射回来。维度 kk 通常选择得比 dd 小得多。对于图像,例如,d=3×1024×1024d = 3 \times 1024 \times 1024,下采样以获得例如 k=3×102416×102416k = 3 \times \tfrac{1024}{16} \times \tfrac{1024}{16} 的情况并不少见。μϕ\mu_{\phi}μθ\mu_{\theta} 一起被称为Autoencoder(autoencoder)。理想情况下,μϕ\mu_{\phi}μθ\mu_{\theta} 的选择应能实现高重建质量,换句话说,使得 μθ(μϕ(x))\mu_{\theta}(\mu_{\phi}(x)) 在平均意义上类似于 xx。因此,Autoencoder 通常使用重建损失进行训练

LRecon(ϕ,θ)=Expdata[μθ(μϕ(x))x2].\begin{aligned}\mathcal{L}_{\text{Recon}}(\phi,\theta)=&\mathbb{E}_{x\sim \pdata}\left[\|\mu_\theta(\mu_\phi(x))-x\|^2\right].\end{aligned}

该损失衡量原始数据点 xx 与重建数据点 μθ(μϕ(x))\mu_\theta(\mu_\phi(x)) 之间的平方误差。

对生成建模的适用性

不幸的是,上述重建损失不足以训练一个“好的”Autoencoder。回想一下,我们的最终目标是在 latent space 中训练一个 Generative Model,并针对由z=μϕ(x),xpdataz = \mu_{\phi}(x), x \sim \pdata 给出的 latent 分布platent(z)p_{\text{latent}}(z)。然后,通过将我们的 latentGenerative Model 的输出传递给解码器μθ\mu_{\theta},来实现pdata(x)\pdata(x) 的 Generative Model。我们目前所表述的 Autoencoder 出现了一个微妙的问题,即我们对platent(z)p_{\text{latent}}(z) 几乎没有控制,因此基本上无法保证platent(z)p_{\text{latent}}(z) 足够良好,以便于训练这样的 Generative Model(即,良好、简单、类似 Gaussian)。虽然将我们的数据在 latent space 中变换可能已经压缩了它,但我们可能已将数据分布pdata\pdata 变换成一个非常难以学习的分布platent\platent。因此,问题是:我们如何确保 latent 分布platent\platent 仍然表现良好且易于学习?为了允许对 latent 分布进行更显式的正则化,我们现在将 Autoencoder 的概念重新构建为一个更一般的概率框架,从而引出 Variational Autoencoder 的概念。

Variational Autoencoder

Variational Autoencoder(VAE)是通过放宽编码器和解码器是确定性函数的约束,从我们的(确定性)Standard Autoencoder 公式中获得的。特别地,让我们考虑一个参数为ϕ\phi 的编码器qϕ(zx)q_\phi(z|x),以及一个参数为θ\theta 的解码器pθ(xz)p_\theta(x|z)。最常见的选择是取

qϕ(zx)=N(z;μϕ(x),diag(σϕ2(x))),pθ(xz)=N(x;μθ(z),σθ2(z)Id)(71)q_\phi(z|x)=\mathcal{N}(z;\mu_\phi(x),\diag(\sigma_\phi^2(x))),\quad p_\theta(x|z)=\mathcal{N}(x;\mu_\theta(z),\sigma_\theta^2(z)I_d) \tag{71}

其中μϕ(x)Rk\mu_\phi(x)\in \mathbb{R}^kσϕ2(x)R0k\sigma_\phi^2(x)\in\mathbb{R}_{\geq 0}^kμθ(z)Rd\mu_\theta(z)\in \mathbb{R}^dσθ2(z)R0\sigma_\theta^2(z)\in\mathbb{R}_{\geq 0} 被参数化为神经网络,diag\diag 表示对角矩阵。为了编码或解码一个变量,我们采样

zqϕ(x)(encode)xpθ(z)(decode)\begin{aligned}z &\sim q_\phi(\cdot|x) &\quad& (\text{encode}) \\ x &\sim p_\theta(\cdot|z) &\quad& (\text{decode})\end{aligned}

最后,我们注意到当σϕ(x)=0\sigma_\phi(x) = 0σθ(x)=0\sigma_\theta(x) = 0 始终成立时,我们恢复了一个 Standard Autoencoder。让我们检查一下重建损失是什么样的。一个自然的目标如下:

LVAE-Recon(ϕ,θ)=Expdata(x),zqϕ(x)[logpθ(xz)](72)\begin{aligned}\mathcal{L}_{\text{VAE-Recon}}(\phi,\theta)=&-\EE_{x \sim \pdata(x),z\sim q_\phi(\cdot|x)} \left[\log p_\theta(x|z)\right]\end{aligned}\tag{72}

注意两个变化:我们不再使用确定性编码,而是现在采样zqϕ(zx)z\sim q_\phi(z|x)。此外,我们现在取xx 在解码下的负对数似然,即损失有效地询问:如果我们编码并解码了原始数据点xx,它会有多可能——并且由于现在事情变得随机,我们考虑了所有可能的解码/编码。对于 Gaussian 情况,这个重建损失变为:

LVAE-Recon(ϕ,θ)=Expdata(x),zqϕ(zx)[12σθ2(z)xμθ(z)2+d2logσθ2(z)]+const(73)\begin{aligned}\mathcal{L}_{\text{VAE-Recon}}(\phi,\theta)=&\EE_{x \sim \pdata(x),z\sim q_\phi(z|x)} \left[\frac{1}{2\sigma^2_\theta(z)}\|x-\mu_\theta(z)\|^2+\frac{d}{2}\log \sigma_{\theta}^2(z)\right]+\text{const}\end{aligned}\tag{73}

其中我们使用了正态分布的密度(参见式(97))。因此,VAE 重建损失与标准 AE 重建损失没有太大不同,我们只需考虑所有可能的编码zqϕ(x)z\sim q_\phi(\cdot|x)。依赖于解码器方差的第二项控制了重建精度和预测不确定性之间的权衡。许多实现,包括实验室中的实现,将σϕ(x)\sigma_\phi(x)σθ(z)\sigma_\theta(z) 固定为学习到的标量常数(即分别独立于xxzz),从而避免了学习方差时的病态行为和数值稳定性问题。因此,在这种情况下,VAE 重建损失基本上变成了 Standard Autoencoder 重建损失,直到编码中的随机性和常数:

LVAE-Recon(ϕ,θ)=Expdata(x),zqϕ(zx)[12σθ2xμθ(z)2]+const(74)\begin{aligned}\mathcal{L}_{\text{VAE-Recon}}(\phi,\theta)=&\EE_{x \sim \pdata(x),z\sim q_\phi(z|x)} \left[\frac{1}{2\sigma_\theta^2}\|x-\mu_\theta(z)\|^2\right]+\text{const}\end{aligned}\tag{74}

现在让我们重新审视我们的目标:我们想要创建数据分布pdata(x)\pdata(x) 的编码,使得在将其映射到 latent space 之后,分布变得“良好”或易于学习。为此,现在让我们引入一个在潜在变量zz 上的先验分布pprior(z)\prior(z)。为了我们的目的,我们将取pprior=N(0,Ik)\prior=\mathcal{N}(0,I_k) 为各向同性 Gaussian。这个先验分布pprior\prior 的选择有效地代表了 latent 分布应该是什么样子的“理想”情况。正态分布将非常容易学习,因此将满足我们获得“可训练”的 latent 分布的目标。因此,主要思想是正则化我们的编码器,以确保编码后的数据分布尽可能接近pprior\prior,我们通过辅助损失来实现这一点

LVAE-Prior(ϕ)=Expdata(x)[DKL ⁣(qϕ(x)pprior)],(75)\begin{aligned}\mathcal{L}_{\text{VAE-Prior}}(\phi)=&\EE_{x \sim \pdata(x)} \left[\dkl{q_\phi(\cdot|x)}{\prior}\right],\end{aligned}\tag{75}

其中DKLD_{KL}Kullback-Leibler(KL)散度。KL 散度是衡量两个概率分布差异的基本方法。详细解释它超出了本工作的范围,但我们在第 6.2.2 节中给出简要背景作为读者的提醒。这里定义的损失LVAE-Prior\mathcal{L}_{\text{VAE-Prior}} 现在非常直观:我们希望对于任何数据点xx,编码分布看起来像 Gaussian distribution。如果我们对所有xx 都这样做,那么自然可以预期我们的 latent 分布也将看起来像 Gaussian distribution。

备注 30(KL-divergence 背景)

对于两个概率密度q,pq,pKullback-Leibler 散度(KL 散度)定义为

DKL ⁣(q(x)p(x))=q(x)logq(x)p(x)=EXq[logq(X)p(X)].\dkl{q(x)}{p(x)}=\int q(x) \log \frac{q(x)}{p(x)}=\mathbb{E}_{X\sim q}\left[\log \frac{q(X)}{p(X)}\right].

KL 散度是分布之间不相似性的标准度量。特别是,KL 散度满足以下有用性质:

DKL ⁣(q(x)p(x))0,(76)\begin{aligned}\dkl{q(x)}{p(x)}&\geq 0,\end{aligned}\tag{76}
DKL ⁣(q(x)p(x))=0q=p.(77)\begin{aligned}\dkl{q(x)}{p(x)}&=0\quad \Leftrightarrow \quad q=p.\end{aligned}\tag{77}

即它总是非负的,并且当且仅当两个概率分布相同时为零。

为了定义 Variational Autoencoder 的损失函数,我们现在可以将重建损失和先验损失与一个参数权重 β0\beta\geq 0 结合起来,得到 VAE 训练目标,如下所示:

LVAE(ϕ,θ)=LVAE-Recon(ϕ,θ)+βLVAE-Prior(ϕ)(78)\begin{aligned}\mathcal{L}_{\text{VAE}}(\phi,\theta)&=\mathcal{L}_{\text{VAE-Recon}}(\phi,\theta)+\beta\mathcal{L}_{\text{VAE-Prior}}(\phi)\end{aligned}\tag{78}
=Expdata(x),zqϕ(zx)[logpθ(xz)]+βExpdata(x)[DKL(qϕ(x)pprior)](79)\begin{aligned}&=-\EE_{x \sim \pdata(x),z\sim q_\phi(z|x)} \left[\log p_\theta(x\mid z)\right]+\beta\EE_{x \sim \pdata(x)} \left[D_{KL}(q_\phi(\cdot|x)||\prior)\right]\end{aligned}\tag{79}

其中第一项确保 latent 变量可以被高效地解码回数据,第二项确保我们的 latent 分布接近 Gaussian distribution。参数 β\beta 控制每一项的强度。为了使这个损失更具体,让我们推导 Gaussian 情况下的 KL 散度:

例 31(各向同性 Gaussian distribution 的 KL 散度)

q(x)=N(x;μq,diag(σq2))q(x)=\mathcal N(x;\mu_q,\diag(\sigma_q^2))p(x)=N(x;μp,diag(σp2))p(x)=\mathcal N(x;\mu_p,\diag(\sigma_p^2)) 为具有对角协方差矩阵的 Gaussian distribution,其中 σq,σpR0d\sigma_q,\sigma_{p}\in\mathbb{R}_{\geq 0}^d,且 xRdx \in \RR^d。那么

DKL ⁣(qp)=12(K(σq2σp2)+μqμp2σp2),where K(α)=i=1dαilogαi1.(80)\dkl{q}{p} =\frac12\left( \mathcal{K}\left(\frac{\sigma_{q}^2}{\sigma_{p}^2}\right)+ \frac{\|\mu_q-\mu_p\|^2}{\sigma_p^2} \right),\quad \text{where }\mathcal{K}(\alpha)=\sum\limits_{i=1}^{d}\alpha_i-\log \alpha_i-1. \tag{80}

上述表达式很直观:如果均值和方差一致,那么 DKL ⁣(qp)=0\dkl{q}{p}=0。此外,它随着均值向量之间的平方误差 μqμp2\|\mu_{q}-\mu_{p}\|^2 的增加而增加。最后,函数 K(α)\mathcal{K}(\alpha)α=1\alpha=1 处有唯一最小值,因此当 σq=σp\sigma_{q}=\sigma_{p} 时,DKL ⁣(qp)\dkl{q}{p} 被最小化。

证明

我们对 d=1d=1 进行证明(对于 d>1d>1 的证明类似,只需对每个维度求和)。给定正态分布的密度,我们知道(参见 式(97)):

logq(x)=12log(2πσq2)12σq2xμq2,logp(x)=12log(2πσp2)12σp2xμp2\log q(x)= -\frac {1}{2}\log(2\pi\sigma_q^2)-\frac{1}{2\sigma_q^2}\|x-\mu_q\|^2, \qquad \log p(x)= -\frac 12\log(2\pi\sigma_p^2)-\frac{1}{2\sigma_p^2}\|x-\mu_p\|^2

那么

DKL(qp)=Exq[logq(x)logp(x)]=12logσp2σq2+12σp2Eq ⁣[xμp2]12σq2Eq ⁣[xμq2].(81)\begin{aligned}D_{\mathrm{KL}}(q\|p) &=\E_{x\sim q}\big[\log q(x)-\log p(x)\big]=\frac 12\log\frac{\sigma_p^2}{\sigma_q^2} +\frac{1}{2\sigma_p^2}\E_q\!\left[\|x-\mu_p\|^2\right] -\frac{1}{2\sigma_q^2}\E_q\!\left[\|x-\mu_q\|^2\right].\end{aligned}\tag{81}

对于 xN(μq,σq2I)x\sim \mathcal N(\mu_q,\sigma_q^2 I),我们有

Eq ⁣[xμq2]=tr(σq2I)=σq2.\E_q\!\left[\|x-\mu_q\|^2\right]=\mathrm{tr}(\sigma_q^2 I)=\sigma_q^2.

将此与 xμp=(xμq)+(μqμp)x-\mu_p=(x-\mu_q)+(\mu_q-\mu_p)Eq[xμq]=0\E_q[x-\mu_q]=0 的事实结合,我们得到

Eq ⁣[xμp2]=Eq ⁣[xμq2]+μqμp2=σq2+μqμp2.\E_q\!\left[\|x-\mu_p\|^2\right] =\E_q\!\left[\|x-\mu_q\|^2\right]+\|\mu_q-\mu_p\|^2 =\sigma_q^2+\|\mu_q-\mu_p\|^2.

将这些代入式(81)得到式(80)。

现在让我们假设编码器具有 Gaussian 形状。那么我们得到:

LVAE-Prior(ϕ)=Expdata(x)[DKL ⁣(qϕ(x)N(0,Ik))]=E[12K(σϕ2(x))+12μϕ(x)2](82)\begin{aligned}\mathcal{L}_{\text{VAE-Prior}}(\phi)=&\EE_{x \sim \pdata(x)} \left[\dkl{q_\phi(\cdot|x)}{\mathcal{N}(0,I_k)}\right]=\mathbb{E}\left[\frac{1}{2} \mathcal{K}\left(\sigma_{\phi}^2(x)\right)+ \frac{1}{2} \|\mu_\phi(x)\|^2\right]\end{aligned}\tag{82}

这个损失函数很直观:均值 μϕ(x)\mu_\phi(x) 因偏离零而受到惩罚,方差因偏离 11 而受到惩罚。作为 VAE 的总损失,我们得到

LVAE(ϕ,θ)=LVAE-Recon(ϕ,θ)+βLVAE-Prior(ϕ)=Expdata(x),zqϕ(zx)[12σθ2(z)xμθ(z)2recon. error+d2logσθ2(z)decoder confidence+β2K(σϕ2(x))make latent variance=1+β2μϕ(x)2make latent mean=0](83)\begin{aligned}&\mathcal{L}_{\text{VAE}}(\phi,\theta)\\ &=\mathcal{L}_{\text{VAE-Recon}}(\phi,\theta)+\beta\mathcal{L}_{\text{VAE-Prior}}(\phi)\\ &=\EE_{x \sim \pdata(x),z\sim q_\phi(z|x)} \left[\underbrace{\frac{1}{2\sigma^2_\theta(z)}\|x-\mu_\theta(z)\|^2}_{\text{recon. error}}+\underbrace{\frac{d}{2}\log \sigma_{\theta}^2(z) }_{\text{decoder confidence}}+\underbrace{\frac{\beta}{2} \mathcal{K}\left(\sigma_{\phi}^2(x)\right)}_{\text{make latent variance}=1}+ \underbrace{\frac{\beta}{2} \|\mu_\phi(x)\|^2}_{\text{make latent mean}=0} \right]\end{aligned}\tag{83}

上述损失函数的四项非常直观:第一项只是重构误差。第二项描述了解码器的不确定性:较小的方差使解码器更“自信”,但也更强烈地惩罚重构误差。此外,我们希望使 latent 的方差为 11,均值为 00——以强制 latent 中的分布接近 Gaussian distribution。

训练一个 VAE

接下来需要讨论如何最小化 VAE 损失 LVAE(ϕ,θ)\mathcal{L}_{\text{VAE}}(\phi,\theta)。该损失的问题在于,到目前为止,我们对其取期望的分布(qϕ(zx)q_\phi(z|x))仍然依赖于参数 ϕ\phi。然而,我们可以应用所谓的重参数化技巧来重写它。具体来说,对于

qϕ(zx)=N(z;μϕ(x),σϕ2(x)Ik)q_\phi(z|x)=\mathcal{N}(z;\mu_\phi(x),\sigma_\phi^2(x)I_k)

我们可以通过以下方式获得样本

ϵN(0,Ik),z=μϕ(x)+σϕ(x)ϵzqϕ(x)\epsilon\sim\mathcal{N}(0,I_k),\quad z=\mu_\phi(x)+\sigma_\phi(x)\epsilon \quad \Rightarrow\quad z\sim q_\phi(\cdot|x)

注意,在这个等式中,唯一的噪声/随机性来源是 ϵ\epsilon,其分布独立于 ϕ\phi。因此,我们可以将损失重写为:

LVAE(ϕ,θ)=Expdata(x),ϵN(0,Ik)[12σθ2(z)xμθ(μϕ(x)+σϕ(x)ϵ)2+d2logσθ2(z)+β2K(σϕ2(x))+β2μϕ(x)2]\begin{aligned}\mathcal{L}_{\text{VAE}}(\phi,\theta)=&\EE_{x \sim \pdata(x),\epsilon\sim\mathcal{N}(0,I_k)} \left[\frac{1}{2\sigma^2_\theta(z)}\|x-\mu_\theta(\mu_\phi(x)+\sigma_\phi(x)\epsilon)\|^2+\frac{d}{2}\log \sigma_{\theta}^2(z) +\frac{\beta}{2} \mathcal{K}\left(\sigma_{\phi}^2(x)\right)+ \frac{\beta}{2} \|\mu_\phi(x)\|^2 \right]\end{aligned}

重参数化之后,随机性仅来自 ϵN(0,Ik)\epsilon\sim\mathcal{N}(0,I_k),其分布不依赖于 ϕ\phi。因此,我们可以使用标准的深度学习工具来最小化这个损失。为了进一步简化,我们可以再次将 σθ2(z)=σ2\sigma_{\theta}^2(z)=\sigma^2 设为常数,得到:

LVAE(ϕ,θ)=Expdata(x),ϵN(0,Ik)[12σ2xμθ(μϕ(x)+σϕ(x)ϵ)2+β2K(σϕ2(x))+β2μϕ(x)2]\begin{aligned}\mathcal{L}_{\text{VAE}}(\phi,\theta)=&\EE_{x \sim \pdata(x),\epsilon\sim\mathcal{N}(0,I_k)} \left[\frac{1}{2\sigma^2}\|x-\mu_\theta(\mu_\phi(x)+\sigma_\phi(x)\epsilon)\|^2 +\frac{\beta}{2} \mathcal{K}\left(\sigma_{\phi}^2(x)\right)+ \frac{\beta}{2} \|\mu_\phi(x)\|^2 \right]\end{aligned}

在 Algorithm 6 中,我们总结了 VAE 的训练过程。

实际注意事项

我们在这里构建的结构展示了 Autoencoder 设计的原则。当然,在实践中,人们可能会添加更多的损失项或其他约束。因此,我们最后补充一些关于 Autoencoder 的实际注意事项:

较大的 β\beta 强制潜在变量更接近先验,但可能损害重构,并可能触发后验坍缩(编码器忽略 xx 并输出 qϕ(zx)N(0,Ik)q_\phi(z|x)\approx \mathcal{N}(0,I_k))。
一种常见的稳定化方法是KL 预热:从 β=0\beta=0 开始,在前几个 epoch 中逐渐增加到目标值。然而,在所有现代 Autoencoder 中,β\beta 的值非常小,即 β<<1\beta<<1
学习 Gaussian 解码器方差 σθ2\sigma_\theta^2 在数值上可能很微妙,并且除非进行正则化,否则可能导致退化解。
为了稳定性,许多实现将 pθ(xz)=N(x;μθ(z),σ2Id)p_\theta(x|z)=\mathcal{N}(x;\mu_\theta(z),\sigma^2 I_d) 固定为常数 σ2\sigma^2,这使得重构项与均方误差成正比(忽略常数)。
对于图像,像素级 Gaussian 似然(均方误差)通常会产生过于平滑的重构。
在实践中,人们会添加感知损失(使用预训练网络的特征空间损失)来提高清晰度和语义保真度。
为了进一步提高视觉真实感,可以将 VAE 目标与对抗损失(VAE-GAN 风格)结合,使用判别器对解码样本进行判别。
这通常会锐化输出,但会引入额外的优化不稳定性和额外的超参数。

  1. 选择 β\beta(以及 KL 预热)。
  2. 解码器方差。
  3. 超越像素 MSE 的重构损失。
  4. 对抗性和混合目标。
备注 32(工作于 latent space)
Algorithm 6
算法 6 · β\beta-VAE 训练过程(固定方差 Gaussian 解码器 pθ(xz)=N(x;μθ(z),σ~2Id)p_\theta(x\mid z)=\mathcal{N}(x;\mu_\theta(z),\tilde\sigma^2I_d)
输入: 样本数据集 xpdatax\sim\pdata,编码器 (μϕ(x),logσϕ2(x))(\mu_\phi(x),\log\sigma_\phi^2(x)),解码器 μθ(z)\mu_\theta(z),latent 维度 kk,常数 β0\beta\geq0σ~2>0\tilde\sigma^2>0

- 循环: 遍历 mini-batch {xi}i=1B\{x_i\}_{i=1}^B
- 编码每个 xix_iμiμϕ(xi)\mu_i\gets\mu_\phi(x_i)logσi2logσϕ2(xi)\log\sigma_i^2\gets\log\sigma_\phi^2(x_i)
- 采样噪声 ϵiN(0,Ik)\epsilon_i\sim\mathcal{N}(0,I_k)
- 重参数化:ziμi+σiϵiz_i\gets\mu_i+\sigma_i\odot\epsilon_i,其中 σi=exp(12logσi2)\sigma_i=\exp(\tfrac12\log\sigma_i^2)
- 解码均值:x^iμθ(zi)\hat x_i\gets\mu_\theta(z_i)
- 重建损失:
Lrecon1Bi=1B12σ~2xix^i2\mathcal{L}_{\mathrm{recon}} \gets\frac{1}{B}\sum_{i=1}^B\frac{1}{2\tilde\sigma^2}\lVert x_i-\hat x_i\rVert^2
- 相对于先验 pprior(z)=N(0,Ik)\prior(z)=\mathcal{N}(0,I_k) 的 KL 损失:
LKL1Bi=1B12j=1k(μi,j2+σi,j2logσi,j21)\mathcal{L}_{\mathrm{KL}} \gets\frac{1}{B}\sum_{i=1}^B\frac12\sum_{j=1}^k \left(\mu_{i,j}^2+\sigma_{i,j}^2-\log\sigma_{i,j}^2-1\right)
- 总损失:LLrecon+βLKL\mathcal{L}\gets\mathcal{L}_{\mathrm{recon}}+\beta\mathcal{L}_{\mathrm{KL}}
- 更新 (ϕ,θ)grad_update(L)(\phi,\theta)\gets\mathrm{grad\_update}(\mathcal{L})
- 结束循环

要训练一个 latent Generative Model,我们只需遵循现有的训练流程,但直接在 latent space 中工作。在训练时,我们从 qϕ(zx)q_\phi(z|x) 中采样,并使用 xpdatax \sim \pdata;在推理时,我们从 latent 扩散或 Flow Model 中采样 zz,然后使用 x=μmean(z)x = \mu_\text{mean}(z) 进行解码(注意,我们取均值而不是随机样本,以避免噪声引起的伪影)。直观地说,一个训练良好的 Autoencoder 可以被认为过滤掉了高频或语义上无意义的细节,使 Generative Model 能够“专注于”重要的、感知相关的特征 [36]。在撰写本文档时,几乎所有最先进的图像和视频生成方法都遵循所谓的 latent 扩散 范式,即在 Autoencoder 的 latent space 内训练流或 Diffusion Model [36, 48]。然而,重要的是要注意:还需要在训练 Diffusion Model 之前训练 Autoencoder。关键的是,性能现在也取决于 Autoencoder 将图像压缩为 latent space 并恢复美观图像的能力。

我们在 第 D 节 中提供了关于 VAE 的额外讨论。

案例:Stable Diffusion 3 与 Meta Movie Gen

我们通过简要考察两个大规模 Generative Model 来结束本节:用于图像生成的 Stable Diffusion 3 和 Meta 的用于视频生成的 Movie Gen Video [14, 33]。正如你将看到的,这些模型使用了我们在本工作中描述的技术,并进行了额外的架构增强,以扩展规模并适应丰富结构的条件模态,例如基于文本的输入。

Stable Diffusion 3

Stable Diffusion 是一系列最先进的图像 Generative Model。这些模型是最早使用大规模 latent Diffusion Model 进行图像生成的模型之一。如果你还没有尝试过,我们强烈建议你在线测试一下(https://stability.ai/news/stable-diffusion-3)。

Stable Diffusion 3 使用我们在本工作中研究的相同 conditional flow matching 目标(见算法 4)。〔脚注〕 正如他们论文中所概述的,他们广泛测试了各种流和扩散替代方案,发现 Flow Matching 表现最佳。对于训练,它使用 classifier-free guidance 训练(带有丢弃类标签),如上所述。此外,Stable Diffusion 3 遵循第 6.1 节中概述的方法,在预训练 Autoencoder 的 latent space 内进行训练。训练一个好的 Autoencoder 是第一批 stable diffusion 论文的一大贡献。

为了增强文本条件化,Stable Diffusion 3 使用了 3 种不同类型的文本 embedding(包括 CLIP embedding 以及 Google 的 T5-XXL [35] 编码器的预训练实例产生的序列输出,类似于 [3, 39] 中采用的方法)。CLIP embedding 提供了输入文本的粗略、总体 embedding,而 T5 embedding 提供了更细粒度的上下文,允许模型关注条件文本的特定元素。为了适应这些序列上下文 embedding,作者提出扩展 Diffusion Transformer,使其不仅关注图像的 patch,还关注文本 embedding,从而将条件化能力从最初为 DiT 提出的基于类的方案扩展到序列上下文 embedding。这种修改后的 DiT 被称为 多模态 DiT(MM-DiT),如图 Figure 16 所示。他们最终的最大模型有 80 亿参数。对于采样,他们使用 5050 步(即他们必须评估网络 5050 次),使用 Euler 模拟方案和 2.02.0-5.05.0 之间的 classifier-free guidance 权重。

Meta Movie Gen Video

接下来,我们讨论 Meta 的视频生成器 Movie Gen Videohttps://ai.meta.com/research/movie-gen/)。由于数据不是图像而是 视频,数据 xx 位于空间 RT×C×H×W\mathbb{R}^{T \times C \times H \times W} 中,其中 TT 表示新的 时间 维度(即帧数)。正如我们将看到的,在这个视频设置中做出的许多设计选择可以被视为将现有技术(例如,Autoencoder、扩散 Transformer 等)从图像设置中调整以适应这个额外的时间维度。

Figure 16[14] 中提出的多模态 Diffusion Transformer(MM-DiT)的架构。图也取自 [14]。

Movie Gen Video 使用 conditional flow matching 目标和相同的直线调度器 αt=t,σt=1t\alpha_t=t,\sigma_{t}=1-t。与 Stable Diffusion 3 一样,Movie Gen Video 也在冻结的预训练 Autoencoder 的 latent space 中运行。请注意,用于减少内存消耗的 Autoencoder 对于视频比对于图像更为重要——这就是为什么目前大多数视频生成器在生成视频的长度上相当有限。

具体来说,作者提出通过引入 时间 Autoencoder(TAE)来处理增加的时间维度,该 Autoencoder 将原始视频 xtRT×3×H×Wx_t' \in \mathbb{R}^{T' \times 3 \times H \times W} 映射到 latent xtRT×C×H×Wx_t\in\mathbb{R}^{T \times C \times H \times W},其中 TT=HH=WW=8\tfrac{T'}{T} = \tfrac{H'}{H} = \tfrac{W'}{W} = 8 [33]。为了适应长视频,提出了一种时间分块过程,将视频切成片段,每个片段分别编码,然后将潜在表示拼接在一起 [33]。模型本身——即 utθ(xt)u_t^\theta(x_t)——由一个 DiT 类骨干网络给出,其中 xtx_t 沿时间和空间维度进行 patch 化。然后,图像 patch 通过一个 Transformer,该网络在图像 patch 之间使用自 attention,并与语言模型 embedding 进行交叉 attention,类似于 Stable Diffusion 3 采用的 MM-DiT。对于文本条件化,Movie Gen Video 使用三种类型的文本 embedding:UL2 embedding,用于细粒度的基于文本的推理 [47];ByT5 embedding,用于关注字符级细节(例如,明确要求出现特定文本的提示)[50];以及 MetaCLIP embedding,在共享的文本-图像 embedding 空间中训练 [24, 33]。他们最终的最大模型有 300 亿参数。对于更详细和广泛的处理,我们鼓励读者查阅 Movie Gen 技术报告本身 [33]。

Discrete Diffusion Model:用扩散构建语言模型

在前面的章节中,我们将 Flow Model 和 Diffusion Model 视为欧几里得空间 Rd\mathbb{R}^d 上的 Generative Model,它们允许我们生成由向量 zRdz\in \mathbb{R}^d 表示的数据点。然而,并非所有数据都自然地建模为欧几里得空间 Rd\R^d 中的一个点。许多数据类型,如文本或 DNA,更自然地被视为离散状态空间 SS 的元素。最重要的是,语言由一系列我们希望建模的离散 token 组成。我们如何将 Flow Model 和 Diffusion Model 应用于这些数据类型?事实证明,我们在前面章节中学到的原理也适用于这些数据类型。由此产生的模型在机器学习文献中被称为Discrete Diffusion Model [5, 16]。然而,重要的是要记住,在离散状态空间中不存在数学上的扩散过程(SDE 在离散状态空间中不存在)。我们使用Continuous-Time Markov Chain (CTMC) 来代替 ODE/SDE。接下来,我们将解释 CTMC 模型(见 第 7.1 节)以及如何学习它们(见 第 7.2 节),从而允许我们使用 Flow Model 和 Diffusion Model 的原理来构建大型语言模型(LLM)。

Continuous-Time Markov Chain (CTMC) 模型

在本节中,我们解释 Continuous-Time Markov Chain (CTMC)。你可以将 CTMC 视为 SDE 的离散类比,我们可以用它来构建生成离散状态的神经网络模型。此外,我们将介绍 CTMC 模型,即允许使用 CTMC 生成离散序列(如文本)的神经网络模型。

Figure 17状态空间为 S={S1,S2,S3}S=\{S_1,S_2,S_3\}(序列长度为 d=1d=1)的 CTMC 轨迹的示意图。图改编自 [5]。

让我们首先描述我们的状态空间 SS。设 V={v1,,vV}\mathcal{V}=\{v_{1},\cdots,v_{V}\} 为我们的词汇表。状态空间由 S=VdS=\mathcal{V}^d 给出,其中 dNd\in \mathbb{N} 是序列长度,VNV\in \mathbb{N} 是词汇表大小。对于语言,{v1,,vV}\{v_{1},\cdots,v_{V}\} 可以枚举我们的字母表或一组离散 token,SS 将表示长度为 dd 的序列(或句子)的集合。对于 DNA,{v1,,vV}\{v_{1},\cdots,v_{V}\} 可以是所有 4 种 DNA 碱基,SS 是所有长度为 dd 的 DNA 序列。

接下来,设 XtX_tSS 上的一个随机过程,即 SS 中的随机轨迹 X:[0,1]S,tXtX:[0,1]\to S, t\mapsto X_{t}。我们要求 XtX_t 是一个Markov process,即一个无记忆的过程。具体来说,这意味着以下条件成立

p(Xt+hXt,Xt1,,Xtk)prob. of future given present and past=p(Xt+hXt)prob. of future given present(for all 0<h,0t1<t2<<tk<t)\underbrace{p(X_{t+h}|X_{t}, X_{t_1},\cdots, X_{t_k})}_{\text{prob. of future given present and past}}=\underbrace{p(X_{t+h}|X_{t})}_{\text{prob. of future given present}}\quad (\text{for all }0<h, 0\leq t_1<t_2<\cdots <t_k<t)

换句话说,未来事件的概率只取决于现在——过去对未来不再有影响。注意,ODE/SDE——虽然不在离散状态空间上——也是 Markov process。这里,XtX_t 在离散空间上,因此被称为 Markov,具体来说是continuous-time Markov chain(CTMC)。量 pt+ht(Xt+hXt)p_{t+h|t}(X_{t+h}|X_{t})转移概率,它们与 Markov chain 的初始分布 X0p0X_0\sim p_0 一起完全决定了 CTMC。因此,当我们说 CTMC 时,你也可以只考虑转移概率 pt+ht(Xt+hXt)p_{t+h|t}(X_{t+h}|X_{t})

接下来,让我们推导离散设置中 Vector Field 的类比。由于我们处于离散设置中,我们只能在状态之间跳跃(或切换)——我们不能像指定 ODE 时那样再沿某个方向前进。因此,我们定义一个Rate Matrix Qt(yx)Q_t(y|x),它有效地总结了从状态 xSx\in S 跳跃(或切换)到状态 ySy\in S 的速率。形式上,Rate Matrix QtQ_t 由一个有界函数(在时间上连续)给出

Q:S×S×[0,1]R,(x,y,t)Qt(yx)(84)Q:S\times S\times[0,1] \to \mathbb{R},\quad (x,y,t)\mapsto Q_t(y|x) \tag{84}

其中 Qt(yx)Q_t(y|x) 描述了从 xx 切换到 yy 的速率,使得

(1) Outgoing rates are positives: Qt(yx)0whenever xy(85)\begin{aligned}\text{(1) Outgoing rates are positives: }Q_t(y|x)\geq& 0 \quad \text{whenever }x\neq y\end{aligned}\tag{85}
(2) Rate staying equals negative outgoing rate: Qt(xx)=yxQt(yx) for all x(86)\begin{aligned}\text{(2) Rate staying equals negative outgoing rate: }Q_t(x|x)=&-\sum\limits_{y\neq x}Q_t(y|x)\quad \text{ for all }x\end{aligned}\tag{86}

这两个条件是直观的:第一个条件说从 xx 切换到不同状态 yxy\neq x 的速率只能是非负的(不切换对应于 00——因此速率小于 00 没有意义)。第二个条件说停留在 xx 的速率 Qt(xx)Q_t(x|x) 应该与离开 xx 的速率抵消——这本质上是一个一致性条件,说明你必须要么停留在 xx,要么离开(没有第三种选择)。注意,这些条件特别意味着 Qt(xx)0Q_t(x|x)\leq 0。因此,Qt(yx)Q_t(y|x) 是一个矩阵,其对角元素全部非正,而非对角元素全部非负。

我们现在可以定义微分方程的类比,即 CTMC “遵循”Rate Matrix 的条件。基本思想是 XX 的分布或演化应该遵循 Rate Matrix QtQ_t。换句话说,我们要求转移概率满足

ddhpt+ht(Xt+h=yXt=x)h=0=Qt(yx)for all x,yS,0t(87)\frac{\dd}{\dd h}p_{t+h|t}(X_{t+h}=y|X_{t}=x)_{|h=0}=Q_t(y|x)\quad\text{for all }x,y\in S, 0\leq t \tag{87}

左边是从 xx 切换到 yy 的概率的无穷小变化率。我们强加条件,即这些概率应该按照 Rate Matrix 指定的方式变化。让我们简要检查一下要求这些条件是否合理,即我们简单地设置 Qt(yx)Q_t(y|x) 如式(87),它会是一个有效的 Rate Matrix 吗?对于 h=0h=0,从 xx 切换到 yxy\neq x 的概率为零(因为没有时间过去),即对所有 yxy\neq xptt(yx)=0p_{t|t}(y|x)=0。因此,我们知道导数必须是非负的,并且当 yxy\neq xQt(yx)0Q_t(y|x)\geq 0。这验证了式(85)中的第一个条件成立。此外,我们知道

yxQt(yx)=yxddhp(Xt+h=yXt=x)h=0=ddhyxp(Xt+h=yXt=x)h=0=ddh(1p(Xt+h=xXt=x))=Qt(xx)\begin{aligned}\sum\limits_{y\neq x} Q_t(y|x)=\sum\limits_{y\neq x}\frac{\dd}{\dd h}p(X_{t+h}=y|X_{t}=x)_{|h=0}=\frac{\dd}{\dd h}\sum\limits_{y\neq x}p(X_{t+h}=y|X_{t}=x)_{|h=0}=&\frac{\dd}{\dd h}(1-p(X_{t+h}=x|X_{t}=x))\\ =&-Q_t(x|x)\end{aligned}

其中我们使用了概率之和为 11。这证明了式(86)。这验证了每个 CTMC 至少有一个满足式(87)的 Rate Matrix。但如果我们反过来——如果我们指定 QtQ_t,是否存在相应的 CTMC,如果存在,它是否唯一?事实确实如此。

定理 33(CTMC 存在且唯一)

对于任何 Rate MatrixQtQ_t(在时间tt 上有界且连续),存在唯一的 Markov chainXtX_{t}(即一组唯一的转移概率pt+ht(yx)p_{t+h|t}(y|x)),使得式(87)成立。

对于感兴趣的读者,我们在附录 C 中提供了一个自包含的证明。该定理的关键要点是,对于机器学习的目的,我们可以陈述并构造一个 Rate MatrixQtQ_t(例如通过神经网络),并假设存在一个与QtQ_t 对应的唯一 Markov chain。

例 34(跳跃率相等的两种状态 CTMC)

S={a,b}S=\{a,b\},并考虑一个时间齐次的 CTMC (Xt)t0(X_t)_{t\ge 0},它以恒定速率λ>0\lambda>0 在两个状态之间切换:

Q=abaλλbλλ.\begin{aligned}Q= \begin{array}{c|cc} & a & b\\\hline a & -\lambda & \lambda\\ b & \lambda & -\lambda \end{array}.\end{aligned}

那么,在时间增量h0h\ge 0 上的转移概率在时间tt 上也是恒定的,并由下式给出

(p(Xt+h=aXt=a)p(Xt+h=aXt=b)p(Xt+h=bXt=a)p(Xt+h=bXt=b))=12(1+e2λh1e2λh1e2λh1+e2λh).\begin{pmatrix} p(X_{t+h}=a|X_t=a) & p(X_{t+h}=a|X_t=b)\\ p(X_{t+h}=b|X_t=a) & p(X_{t+h}=b|X_t=b) \end{pmatrix} = \frac12 \begin{pmatrix} 1+e^{-2\lambda h} & 1-e^{-2\lambda h}\\ 1-e^{-2\lambda h} & 1+e^{-2\lambda h} \end{pmatrix}.

可以手动检查式(87)成立,即这些转移概率确实是该 Rate Matrix 的正确转移概率。事实上,这些速率非常直观:链以瞬时速率λ\lambda 不断翻转。指数项e2λhe^{-2\lambda h} 捕捉了初始状态记忆的衰减。当无限时间过去,即对于hh\to\infty,有

P(h)(12121212),\begin{aligned}P(h)\to \begin{pmatrix}\frac12&\frac12\\\frac12&\frac12\end{pmatrix},\end{aligned}

因此,链忘记了它从哪里开始,并以概率1/21/2 处于aabb。切换速率λ>0\lambda>0 越高,收敛越快。

CTMC 的模拟。接下来,让我们思考如何模拟 CTMC 的轨迹。设h>0h>0 为步长,pinit\pinitSS 上的初始分布,例如pinit=UnifS\pinit=\text{Unif}_{S}SS 上的均匀分布。然后我们可以通过设置X0pinitX_0\sim \pinit 并设置以下内容来迭代模拟它

Xt+hpt+ht(Xt)\begin{aligned}X_{t+h}\sim &p_{t+h|t}(\cdot|X_t)\end{aligned}

现在,如果我们知道pt+ht(Xt)p_{t+h|t}(\cdot|X_t),这会有效。然而,对于除最简单的 CTMC 之外的所有 CTMC,我们通常不知道转移核的闭式形式,只能访问 Rate MatrixQtQ_t。尽管如此,根据式(87):

pt+ht(Xt+h=yXt=x)=ptt(Xt=yXt=x)+hQt(yx)+Rt(h)=1y=x+hQt(yx)+Rt(h)p_{t+h|t}(X_{t+h}=y|X_t=x) =p_{t|t}(X_{t}=y|X_t=x)+hQ_t(y|x)+R_{t}(h)=1_{y=x}+hQ_t(y|x)+R_t(h)

其中Rt(h)R_t(h) 是一个误差项,对于小的hh 我们可以忽略它。因此,对于小的hh,我们可以设置

pt+ht(Xt+h=yXt=x)1y=x+hQt(yx)=:p~t+ht(yx)p_{t+h|t}(X_{t+h}=y|X_t=x)\approx 1_{y=x}+hQ_t(y|x)=:\tilde{p}_{t+h|t}(y|x)

可以检查,由于我们对 Rate Matrix 施加的条件,p~t+ht(yx)\tilde{p}_{t+h|t}(y|x) 对于小的hh 确实是一个有效的概率分布。因此,我们可以通过以下方式近似采样下一个点

Xt+hp~t+ht(x)=(1y=x+hQt(yx))yS(88)X_{t+h}\sim \tilde{p}_{t+h|t}(\cdot|x)=(1_{y=x}+hQ_t(y|x))_{y\in S} \tag{88}

由于上述只是一个离散分布,我们可以通过标准方法轻松地从中采样。这是模拟 CTMC 的一种简单方法。

CTMC 模型

接下来,让我们定义如何在神经网络中参数化 CTMC。一个CTMC 模型(或Discrete Diffusion Model)由初始分布pinit\pinit(在SS 上)和一个具有参数θ\theta 的神经网络QtθQ_t^\theta 给出,使得对于每个输入xSx\in S,模型返回 Rate Matrix 的单列。

x{Qtθ(yx)}ySx\mapsto \{Q_t^\theta(y|x)\}_{y\in S}

我们希望模型返回整个列,因为模拟 CTMC(式(88))时需要它,即采样下一个状态。

上述模型的一个复杂之处在于空间SS 可能非常大。特别是,S=Vd|S|=V^d,其中VV 是我们的词汇表大小,dd 是序列长度。这种指数增长使得在内存中存储 Rate Matrix 的整个列基本上不可能——{Qtθ(yx)}yS\{Q_t^\theta(y|x)\}_{y\in S} 永远无法在计算机中表示。因此,我们必须约束模型。具体来说,几乎所有 CTMC 模型都是因子化的(见图 Figure 18),这实际上是一种稀疏性约束。具体来说,一个因子化 CTMC 模型由一个 CTMC 模型QtθQ_t^\theta 给出,使得对于所有y=(y1,,yd),x=(x1,,xd)S=Vdy=(y_1,\cdots, y_d),x=(x_1,\cdots,x_d)\in S=\mathcal{V}^d,它满足

Qtθ(yx)=0whenever yixi for more than one position iQ_t^\theta(y|x)=0\quad \text{whenever }y_i\neq x_i\text{ for more than one position }i

我们将所有在至多一个 token 上与xx 不同的yy 称为xx邻居N(x)N(x)。我们可以将这样的因子化 CTMC 模型写为

x{Qtθ(yx)}yN(x)=(Qtθ(v1,1x)Qtθ(vV,1x)Qtθ(v1,dx)Qtθ(vV,dx))\begin{aligned}x\mapsto \{Q_t^\theta(y|x)\}_{y\in N(x)} =& \begin{pmatrix} Q_t^\theta(v_1,1|x) & \cdots Q_t^\theta(v_{V},1|x)\\ \cdots\\ Q_t^\theta(v_1,d|x) & \cdots Q_t^\theta(v_{V},d|x)\\ \end{pmatrix}\end{aligned}

其中Qt(yx)=Qtθ(vi,jx)Q_t(y|x)=Q_t^\theta(v_i,j|x) 现在给出了从x=(x1,,xd)x=(x_1,\cdots,x_d)xx 的邻居的速率,该邻居是通过将第jj 个元素替换为viv_i 而获得的,即y=(x1,,xj1,vi,xj+1,,xd)y=(x_1,\cdots,x_{j-1},v_{i},x_{j+1},\cdots,x_{d})。每一行对应于每个位置i=1,,di=1,\cdots,d 的 Rate Matrix,即我们要求

Qtθ(v,ix)0 if vxi,Qt(xi,ix)=vxiQtθ(v,ix)Q_t^\theta(v,i|x)\geq 0\text{ if }v\neq x_i,\quad Q_t(x_i,i|x)=-\sum\limits_{v\neq x_i}Q_t^\theta(v,i|x)

我们可以轻松地在神经网络的输出上强制执行这些条件,例如,可以使用一个序列长度dd、输出维度VV 的 Transformer 模型。还要注意,因子化 Rate Matrix 使得输出形状为d×Vd\times V——这个大小随维度线性增长(而不是指数增长)。

模拟 CTMC 模型

为了从 CTMC 模型中采样,我们先采样 X0pinitX_{0}\sim \pinit,再按照式(88)迭代采样下一个状态。Algorithm 7 给出了完整过程。如其中所示,对于 factorized CTMC 模型,可以采用并行的逐 token Euler 近似:在一个很小的步长 h>0h>0 内,每个 token 独立更新。该近似与完整 CTMC 的 Euler 步在 hh 的一阶精度上相同,但允许多个 token 以 O(h2)O(h^2) 的概率同时更新。

Figure 18因子化 CTMC 模型的图示。因子化 CTMC 只有在起点和终点仅在一个维度上不同时(此处为d=2d=2)才具有非零速率(Qt(yx)0Q_t(y|x)\neq 0)。图取自[26]。
Algorithm 7
算法 7 · 从因子化 CTMC 模型采样
输入: 速率网络QtθQ_t^\theta(因子化),初始分布pinit\pinit,步数nn

- 设置t0t \gets 0,步长h1nh \gets \frac{1}{n}

- 抽取样本X0pinitX_0 \sim \pinit,其中X0=(X0(1),,X0(d))VdX_0=(X_0^{(1)},\dots,X_0^{(d)})\in\mathcal{V}^d

- 循环: i=1,,ni=1,\dots,n

- 计算因子化跳跃速率{qj(v)}j=1..d, vVQtθ(Xt)\{q_{j}(v)\}_{j=1..d,\ v\in\mathcal{V}} \gets Q_t^\theta(\cdot \mid X_t)

- 循环: j=1,,dj=1,\dots,d (并行)

- xXt(j)x \gets X_t^{(j)} 位置jj 处的当前 token

- 定义逐位置欧拉转移概率p~j,t(Xt(j)=x)\tilde p_{j,t}(\cdot \mid X_t^{(j)}=x),通过
p~j,t(vx)={hqj(v),vx,[4pt]1hvV{x}qj(v),v=x.\begin{aligned}\tilde p_{j,t}(v\mid x) = \begin{cases} h q_{j}(v), & v\neq x, [4pt] 1 - h\sum\limits_{v'\in \mathcal{V}\setminus\{x\}} q_{j}(v'), & v=x. \end{cases}\end{aligned}

- 采样Xt+h(j)Categorical({p~j,t(vx)}vV)X_{t+h}^{(j)} \sim \textsc{Categorical} \left(\{\tilde p_{j,t}(v\mid x)\}_{v\in\mathcal{V}}\right)

- 结束循环

- 设置tt+ht \gets t + h

- 结束循环

- 返回: X1X_1

训练 CTMC 模型

我们接下来讨论如何学习 CTMC 模型。其原理与 Flow Matching 相同:(1)我们构造一个在噪声和数据之间插值的 Probability Path。(2)我们推导出一个 Conditional Rate Matrix 和 Marginal Rate Matrix。(3)我们以无模拟的方式学习 Marginal Rate Matrix。我们现在逐步解释这一方法。

在本节中,数据分布 pdata\pdataSS 上的一个分布,由概率质量函数刻画。即 pdata:SR0,zpdata(z)\pdata:S\to\mathbb{R}_{\geq 0}, z\mapsto \pdata(z),其中 zSpdata(z)=1\sum_{z\in S}\pdata(z)=1。我们不知道 pdata\pdata,但在训练期间我们可以以数据集的形式访问样本 zpdataz\sim \pdata。例如,万维网上的所有文本。我们的目标是学习生成样本 zpdataz\sim \pdata。我们的目标是训练 CTMC 模型 QtθQ_t^\theta,使得

X0pinit,Xt CTMC of QtθX1pdataX_0\sim \pinit, \quad X_t\text{ CTMC of }Q_t^\theta\quad \Rightarrow \quad X_{1}\sim \pdata

所以你可能意识到,这与欧几里得情况 Rd\mathbb{R}^d(见第 2 节、第 3 节)没有区别,只是我们使用 CTMC 模型而不是流/Diffusion Model。

Conditional and Marginal Probability Paths

我们定义 δz(x)\delta_{z}(x) 为这样的函数:如果 xzx\neq z,则 δz(x)=0\delta_{z}(x)=0;如果 x=zx=z,则 δz(x)=1\delta_{z}(x)=1。一个(离散的)Conditional Probability Path由一组分布 pt(xz)p_t(x|z) 给出,其中 x,zSx,z\in S0t10\leq t\leq 1,使得

p0(z)=pinit,p1(z)=δzp_0(\cdot|z)=\pinit, \quad p_1(\cdot|z)=\delta_{z}

因此,与欧几里得情况类似,离散 Conditional Probability Path 在独立于 zz 的分布和将所有质量放在 zz 上的分布之间进行插值。然后,一个(离散的)Marginal Probability Path由下式给出

pt(x)=zSpt(xz)pdata(z)p_t(x)=\sum\limits_{z\in S}p_t(x|z)\pdata(z)

人们可以很容易地检查 Marginal Probability Path 在“噪声”和数据之间进行插值:

p0=pinit,p1=pdata(89)p_0=\pinit, \quad p_1=\pdata \tag{89}
例 35(因子混合路径(每个 token 的独立噪声))

S=VdS=\mathcal{V}^d,并设 pinit(x)=j=1dpinit(j)(xj)\pinit(x)=\prod_{j=1}^d \pinit^{(j)}(x_j) 为一个因子化的初始分布。
固定一个调度器 0κt10\le \kappa_t\le 1,使得 κ0=0,κ1=1\kappa_{0}=0,\kappa_{1}=1,其中 ddtκ˙t0\frac{\dd}{\dd t}\dot{\kappa}_t\geq 0。通过下式定义条件路径

pt(xz)=j=1d[(1κt)pinit(j)(xj)+κtδzj(xj)].\begin{aligned}p_t(x|z) &=\prod_{j=1}^d\Big[(1-\kappa_t)\,\pinit^{(j)}(x_j)+\kappa_t\,\delta_{z_j}(x_j)\Big].\end{aligned}

等价地,可以通过抽取独立同分布的掩码 mj=0,1m_j=0,1 和噪声 ξjpinit(j)\xi_j\sim \pinit^{(j)},然后设置 xpt(z)x\sim p_t(\cdot\mid z) 来采样。

mjBernoulli(κt),ξjpinit(j)xj=mjzj+(1mj)ξj,j=1,,dx=(x1,,xd)\begin{aligned}m_j&\sim\mathrm{Bernoulli}(\kappa_t), \quad \xi_j\sim \pinit^{(j)}\\ x_j &= m_j\, z_j + (1-m_j)\,\xi_j,\qquad j=1,\dots,d\\ x&=(x_1,\cdots,x_d)\end{aligned}

我们将上述称为因子化混合路径。上述过程有效地“破坏”了序列中每个位置的第 jj 个 token,破坏概率为 1κt1-\kappa_{t},即对于 t=0t=01κt=11-\kappa_{t}=1,所有信息都被破坏;对于 t=1t=11κt=01-\kappa_{t}=0,没有信息被破坏。请注意,这与 Gaussian Probability Path 第 3.1 节类似,因为信息以由调度器 κt\kappa_{t} 决定的速度被逐步破坏。然而,它也与 Gaussian Probability Path 不同,因为因子化混合路径移动/传输概率质量(因为我们在离散空间中,没有方向)——它只是淡出一个分布并淡入另一个分布。

Figure 19d=2d=2 的离散 Probability Path 的图示。顶行:在初始分布和狄拉克分布之间插值的 Conditional Probability Path。底行:在初始分布和数据分布(此处为棋盘图案)之间的插值。注意与图 5 的相似之处和不同之处:这里,Probability Path 被“传送”(我们降低初始分布的权重并提高终端分布的权重)。

Conditional Rate Matrix 与 Marginal Rate Matrix

作为下一步,我们现在将构造离散 Flow Matching 的训练目标。首先,我们构造一个 Conditional Rate Matrix——这是 Flow Matching 的 Conditional Vector Field 的类比。设 Qtz(yx)Q_t^z(y|x) 为每个数据点 zSz\in S 的 Rate Matrix。然后,如果满足以下条件,我们称之为Conditional Rate Matrix

X0pinit,Xt CTMC of QtzXtpt(z)X_0\sim \pinit,\quad X_t\text{ CTMC of }Q_t^z\quad \Rightarrow \quad X_{t}\sim p_t(\cdot|z)

换句话说,Conditional Rate Matrix 使得其 CTMC“遵循”Conditional Probability Path。Conditional Rate Matrix 作为构建 Marginal Rate Matrix 的基石,该 Marginal Rate Matrix 遵循 Marginal Probability Path:

定理 36(离散边缘化技巧)

由下式定义的Marginal Rate Matrix

Qt(yx)=zSQtz(yx)pt(xz)pdata(z)pt(x)=zSQtz(yx)p1t(zx)where p1t(zx):=pt(xz)pdata(z)pt(x)(90)Q_t(y|x)=\sum\limits_{z\in S}Q_t^z(y|x)\frac{p_t(x|z)\pdata(z)}{p_t(x)}=\sum\limits_{z\in S}Q_t^z(y|x)p_{1|t}(z|x)\quad \text{where }p_{1|t}(z|x):=\frac{p_t(x|z)\pdata(z)}{p_t(x)} \tag{90}

是一个有效的 Rate Matrix,并满足以下条件:

X0pinit,Xt CTMC of QtXtptX_0\sim \pinit,\quad X_t\text{ CTMC of }Q_t\quad \Rightarrow \quad X_{t}\sim p_t

特别地,X1pdataX_1\sim \pdata 由式(89),即 Marginal Rate Matrix 的 CTMC 将噪声转换为数据。

为了证明这一陈述,我们需要 CTMC 的一个基本方程,即所谓的Kolmogorov 前向方程

命题 2(柯尔莫哥洛夫远期方程)

ptp_tSS 上的一组分布,对于每个0t10\leq t\leq 1。进一步,设XtX_t 是一个具有矩阵QtQ_t 和初始分布p0p_0 的 CTMC。那么,当且仅当Kolmogorov 前向方程(KFE)成立时,对于所有0t10\leq t\leq 1XtptX_t\sim p_t 成立:

ddtpt(x)=ySQt(xy)pt(y)\frac{\dd}{\dd t}p_t(x)=\sum\limits_{y\in S}Q_t(x|y)p_t(y)
证明(Proof of KFE)

为了证明 KFE 的必要性,假设pt(x)p_t(x) 是 CTMC 的真实边际分布,即对于每个0t10\leq t\leq 1XtptX_t\sim p_t。然后我们可以计算:

ddtpt(x)=(i)ddhh=0pt+h(x)=(ii)ddhh=0ypt+ht(xy)pt(y)=(iii)yddhh=0pt+ht(xy)pt(y)=(iv)yQt(xy)pt(y)\begin{aligned}\frac{\dd}{\dd t}p_t(x)&\overset{(i)}{=}\frac{\dd}{\dd h}_{|h=0}p_{t+h}(x)\\ &\overset{(ii)}{=}\frac{\dd}{\dd h}_{|h=0}\sum\limits_{y}p_{t+h|t}(x|y)p_t(y)\\ &\overset{(iii)}{=}\sum\limits_{y}\frac{\dd}{\dd h}_{|h=0}p_{t+h|t}(x|y)p_t(y)\\ &\overset{(iv)}{=}\sum\limits_{y}Q_t(x|y)p_t(y)\end{aligned}

其中在(i)(i) 中我们简单地使用时间偏移,在(ii)(ii) 中我们使用转移概率的定义,在(iii)(iii) 中我们交换求和与导数,在(iv)(iv) 中我们使用 Rate Matrix 的定义(见式(87))。

接下来,为了证明 KFE 的充分性,我们可以将 KFE 重写为矩阵形式:

ddtpt=Qtpt\frac{\dd}{\dd t}p_t = Q_tp_t

在此方程中,我们将 pt=(pt(x))xSp_t=(p_t(x))_{x\in S} 视为向量,将 Qt=(Qt(yx))x,ySQ_t=(Q_t(y|x))_{x,y\in S} 视为矩阵。注意,上述是向量空间 RS\mathbb{R}^S 上的线性 ODE。其初始条件由定理中所述的 p0p_0 固定。因此,如果任何其他边际集合 qtq_t 满足此方程,则由 ODE 的唯一性(见第 2.1 节)可知,我们可以得出结论 qt=ptq_t=p_t。这表明 KFE 也是充分的。

证明(第 7.2.2 节的证明)

利用 KFE,只需证明定理中定义的 Marginal Rate Matrix(见式(90))满足 KFE:

ddtpt(x)=(i)ddtzSpt(xz)pdata(z)=(ii)zSddtpt(xz)pdata(z)=(iii)zS[ySQtz(xy)pt(yz)]pdata(z)=(iv)ySpt(y)[zSQtz(xy)pt(yz)pdata(z)pt(y)]=(v)ySpt(y)Qt(xy)\begin{aligned}\frac{\dd}{\dd t}p_t(x)\overset{(i)}{=}&\frac{\dd}{\dd t}\sum\limits_{z\in S}p_t(x|z)\pdata(z)\\ \overset{(ii)}{=}&\sum\limits_{z\in S}\frac{\dd}{\dd t}p_t(x|z)\pdata(z)\\ \overset{(iii)}{=}&\sum\limits_{z\in S}\left[\sum\limits_{y\in S}Q_t^z(x|y)p_t(y|z)\right]\pdata(z)\\ \overset{(iv)}{=}&\sum\limits_{y\in S}p_t(y)\left[\sum\limits_{z\in S}Q_t^z(x|y)\frac{p_t(y|z)\pdata(z)}{p_t(y)}\right]\\ \overset{(v)}{=}&\sum\limits_{y\in S}p_t(y)Q_t(x|y)\end{aligned}

其中 (i)(i) 由 Marginal Probability Path 的定义得出,在 (ii)(ii) 中我们交换了求和与导数,在 (iii)(iii) 中我们对 Conditional Rate Matrix 使用 KFE,在 (iv)(iv) 中我们乘以并除以 pt(y)p_t(y),在 (v)(v) 中我们使用 Marginal Rate Matrix Qt(yx)Q_t(y|x) 的定义。这表明 KFE 得到满足。该结论由第 7.2.2 节得出。

现在让我们为因子化混合路径推导一个 Conditional Rate Matrix 的具体示例。

例 37(分解混合路径的 Conditional Rate Matrix)

ddtκt=κ˙t\frac{\dd}{\dd t}\kappa_t=\dot{\kappa}_t。因子化混合路径具有一个因子化的 Conditional Rate Matrix,由下式给出

Qtz(yx)=(Qtz(vi,jxj))vi,jQtz(vi,jxj)=κ˙t1κt(δzj(vi)δxj(vi))=κ˙t1κt{0if xj=zj1 if vi=zj,xjzj0 if vizj,xjzj1 if vi=xj,xjzj\begin{aligned}Q_t^z(y|x)&=(Q_t^z(v_i,j|x_j))_{v_i,j}\\ Q_t^z(v_i,j|x_j)&=\frac{\dot{\kappa}_t}{1-\kappa_t}(\delta_{z_j}(v_i)-\delta_{x_j}(v_i))\\ =&\frac{\dot{\kappa}_t}{1-\kappa_t}\begin{cases} 0 & \text{if }x_j=z_j\\ 1 & \text{ if }v_i=z_j, x_j\neq z_j\\ 0 & \text{ if }v_i\neq z_j, x_j\neq z_j\\ -1 & \text{ if }v_i=x_j, x_j\neq z_j \end{cases}\end{aligned}

注意,这是一个非常简单的 Rate Matrix:它只允许跳转到 zjz^j——即,如果任何 token jj 被更新,它必须跳转到终端数据点 z=(z1,,zd)z=(z_1,\cdots,z_d) 的 token 值——并且它只在尚未到达该值时跳转到 zjz^j

证明

我们注意到,因子化混合路径完全分解为独立的组件,所建议的 Conditional Rate Matrix 也是如此。因此,我们可以不失一般性地假设 d=1d=1。所以我们只需按维度进行计算。然后,我们可以推导出:

ddtpt(xz)=(i)ddt[(1κt)pinit(x)+κtδz(x)]=(ii)κ˙tδz(x)κ˙tpinit(x)=(iii)κ˙t1κt(δz(x)[(1κt)pinit(x)+κtδz(x)])=(iv)κ˙t1κt(δz(x)pt(xz))=(v)κ˙t1κtδz(x)(1pt(xz))+κ˙t1κt(δz(x)1)pt(xz)=(vi)yxκ˙t1κtδz(x)pt(yz)+κ˙t1κt(δz(x)1)pt(xz)=(vii)yxQtz(xy)pt(yz)+Qtz(xx)pt(xz)=(viii)ySQtz(xy)pt(yz)\begin{aligned}\frac{\dd}{\dd t}p_t(x|z) \overset{(i)}{=}&\frac{\dd}{\dd t}\left[(1-\kappa_t)\pinit(x)+\kappa_t\delta_{z}(x)\right]\\ \overset{(ii)}{=}&\dot{\kappa}_t\delta_{z}(x)-\dot{\kappa}_t\pinit(x)\\ \overset{(iii)}{=}&\frac{\dot{\kappa}_t}{1-\kappa_t}(\delta_{z}(x)-[(1-\kappa_t)\pinit(x)+\kappa_t\delta_{z}(x)])\\ \overset{(iv)}{=}&\frac{\dot{\kappa}_t}{1-\kappa_t}(\delta_{z}(x)-p_t(x|z))\\ \overset{(v)}{=}&\frac{\dot{\kappa}_t}{1-\kappa_t}\delta_{z}(x)\left( 1-p_t(x|z) \right)+\frac{\dot{\kappa}_t}{1-\kappa_t}(\delta_{z}(x)-1)p_t(x|z)\\ \overset{(vi)}{=}&\sum\limits_{y\neq x}\frac{\dot{\kappa}_t}{1-\kappa_t}\delta_{z}(x)p_t(y|z) +\frac{\dot{\kappa}_t}{1-\kappa_t}(\delta_{z}(x)-1)p_t(x|z)\\ \overset{(vii)}{=}&\sum\limits_{y\neq x}Q_t^z(x|y)p_t(y|z)+Q_t^z(x|x)p_t(x|z)\\ \overset{(viii)}{=}&\sum\limits_{y\in S}Q_t^z(x|y)p_t(y|z)\end{aligned}

其中 (i)(i) 使用了 d=1d=1 的因子化混合路径的定义,(ii)(ii) 通过求导并设置 ddtκt=κ˙t\frac{\dd}{\dd t}\kappa_t=\dot{\kappa}_t 得到,(iii)(iii) 由简单代数得出,(iv)(iv) 由因子化混合路径的定义得出,(v)(v) 由简单代数得出,(vi)(vi)ySpt(yz)=1\sum_{y\in S}p_t(y|z)=1 这一事实的定义得出,(vii)(vii) 由 Rate Matrix 的定义得出,(viii)(viii) 由简单代数得出。上述表明 KFE 得到满足,因此结论成立。

学习 Marginal Rate Matrix

在本节中,我们推导训练 CTMC 模型的基本算法。根据第 7.2.2 节,训练 CTMC 模型 Qtθ(yx)Q_t^\theta(y|x) 可以通过学习 Marginal Rate Matrix 来实现。

在本节中,我们现在将自身限制于因子化混合路径(见第 7.2.1 节),因为这是目前大多数离散扩散/Flow Matching 模型所使用的路径。在这种情况下,Marginal Rate Matrix 具有非常直观的形式:

定理 38(分解混合路径的边缘化技巧)

因子化混合路径的 Marginal Rate Matrix 是因子化的,并且具有如下形式

Qt(vi,jx)=κ˙t1κt(p1t(zj=vix)δxj(vi))\begin{aligned}Q_t(v_i,j|x)&=\frac{\dot{\kappa}_t}{1-\kappa_t}(p_{1|t}(z_j=v_i|x)-\delta_{x_j}(v_i))\end{aligned}

其中 p1t(zj=vix)p_{1|t}(z_j=v_i|x) 是给定完整噪声序列 xx 时,第 jj 个位置(序列中的第 jj 个 token)等于 viv_i 的条件概率。

证明

Marginal Rate Matrix 由下式给出

Qt(yx)=zSQtz(yx)p1t(zx)(91)\begin{aligned}Q_t(y|x)&=\sum\limits_{z\in S}Q_t^z(y|x)p_{1|t}(z|x)\end{aligned}\tag{91}

现在,每当 yyxx 不是邻居(相差超过一个 token)时,对于每个 zz,都有 Qtz(yx)=0Q_t^z(y|x)=0。因此,在这种情况下也有 Qt(yx)=0Q_t(y|x)=0。这表明 Marginal Rate Matrix 也是因子化的。于是有

Qt(vi,jx)=zSQtz(vi,jx)p1t(zx)(92)\begin{aligned}Q_t(v_i,j|x)&=\sum\limits_{z\in S}Q_t^z(v_i,j|x)p_{1|t}(z|x)\end{aligned}\tag{92}
=(i)zSκ˙t1κt(δzj(vi)δxj(vi))p1t(zx)(93)\begin{aligned}&\overset{(i)}{=}\sum\limits_{z\in S}\frac{\dot{\kappa}_t}{1-\kappa_t}(\delta_{z_j}(v_i)-\delta_{x_j}(v_i))p_{1|t}(z|x)\end{aligned}\tag{93}
=(ii)κ˙t1κt(zSδzj(vi)p1t(zx)δxj(vi))(94)\begin{aligned}&\overset{(ii)}{=}\frac{\dot{\kappa}_t}{1-\kappa_t}\left(\sum\limits_{z\in S}\delta_{z_j}(v_i)p_{1|t}(z|x)-\delta_{x_j}(v_i)\right)\end{aligned}\tag{94}
=(iii)κ˙t1κt(p1t(zj=vix)δxj(vi))(95)\begin{aligned}&\overset{(iii)}{=}\frac{\dot{\kappa}_t}{1-\kappa_t}\left(p_{1|t}(z_j=v_i|x)-\delta_{x_j}(v_i)\right)\end{aligned}\tag{95}

其中 (i)(i) 由 Conditional Rate Matrix 的公式得出(见第 7.2.2 节),(ii)(ii)zSp1t(zx)=1\sum_{z\in S}p_{1|t}(z|x)=1 这一事实得出,(iii)(iii) 由边缘化得出。证明完毕。

前面的定理非常引人注目:Marginal Rate Matrix 实际上是概率 p1t(zj=vix)p_{1|t}(z_j=v_i|x) 的重新参数化。这实际上不过是为每个 token 位置 j=1,,dj=1,\dots,d 学习一个分类器。换句话说,我们可以简单地将一个去噪概率网络定义为

p1tθ:xnetwork input(p1tθ(zj=vix))j=1,,d,viVnetwork outputp_{1|t}^\theta:\underbrace{x}_{\text{network input}}\mapsto\underbrace{(p_{1|t}^\theta(z_j=v_i|x))_{j=1,\cdots,d, v_i\in \mathcal{V}}}_{\text{network output}}

注意,网络输出的形状为 d×Vd\times V。可以通过简单的 softmax 层获得每个 token 位置的概率。网络本身可以是标准的序列到序列网络,例如 Transformer 就可以(见第 6.1.2 节)。

由于这仅仅是每个位置 jj 的一个分类器,我们可以通过每个 j=1,,dj=1,\cdots,d 的交叉熵损失来训练这样的网络。这导致了离散 Flow Matching 损失,由下式给出

LDFM(θ)=Ezpdata,tUnif[0,1],xpt(z)[j=1dlogp1tθ(zjx)]\mathcal{L}_{\text{DFM}}(\theta)=\mathbb{E}_{z\sim \pdata, t\sim\text{Unif}_{[0,1]}, x\sim p_t(\cdot|z)}\left[\sum\limits_{j=1}^{d}-\log p_{1|t}^\theta(z_j|x)\right]

这非常了不起:要训练一个 Generative Model,我们只需要为每个位置 jj 训练一个分类器模型。与连续 Flow Matching 简化为简单回归(见第 3 节)的方式相同,离散 Flow Matching 和 Discrete Diffusion Model 简化为简单的分类训练。在 Algorithm 8 中,我们总结了训练算法。训练后,我们可以通过算法 7 进行采样。

例 39(掩蔽扩散语言模型)

上述方法的一个特例是掩码扩散语言模型(MDLMs)。MDLM 的思想是,我们可以通过引入一个新的 token [mask]\text{[mask]} 来扩展 token 词汇表 V={v1,,vV}\mathcal{V}=\{v_1,\cdots,v_{V}\},该 token 表示这个 token 缺失(或被掩码)。具体来说,我们设置 V={v1,,vV,[mask]}\mathcal{V}=\{v_1,\cdots,v_{V},\text{[mask]}\},初始点就是 [mask]d\text{[mask]}^d,即全掩码序列。形式上,这意味着在上述框架中设置 pinit=δ[mask]d\pinit=\delta_{\text{[mask]}^d}。采样过程如图 Figure 20 所示。

Algorithm 8
算法 8 · 训练因子化 CTMC 模型(离散扩散)
输入: 序列 zpdataz\sim \pdata 的 dataset,其中 z=(z1,,zd)Vdz=(z_1,\dots,z_d)\in\mathcal{V}^d

初始(噪声)token 边际 pinit(j)\pinit^{(j)}V\mathcal{V} 上;调度 κt[0,1]\kappa_t\in[0,1]
后验网络 fθf_\theta,返回每个位置的 logits,覆盖 V\mathcal{V};优化器 Opt

- 循环: 每次训练迭代

- 采样一个数据点 zpdataz \sim \pdata

- 采样时间 tUnif[0,1]t \sim \mathrm{Unif}[0,1] 并计算 κκt\kappa \gets \kappa_t

- 采样一个噪声状态 xpt(z)x \sim p_t(\cdot\mid z)(因子化混合路径):

- 循环: j=1,,dj=1,\dots,d (并行)

- 采样掩码 mjBernoulli(κ)m_j\sim\mathrm{Bernoulli}(\kappa)

- 采样噪声 token ξjpinit(j)\xi_j\sim \pinit^{(j)}

- 设置 xjmjzj+(1mj)ξjx_j \gets m_j z_j + (1-m_j) \xi_j

- 结束循环

- x(x1,,xd)x \gets (x_1,\dots,x_d)

- 通过网络的 logits 预测终端 token 后验:
j()  fθ(x,t)jp1tθ(vx)j=Softmax(j)(v)\ell_j(\cdot)\ \gets\ f_\theta(x,t)_j \Rightarrow p^\theta_{1|t}(v\mid x)_j = \mathrm{Softmax} \big(\ell_j\big)(v)

- 离散 Flow Matching 损失(token 级别的 zz 的负对数似然):
LDFM(θ)j=1d[logp1tθ(zjx)j]\mathcal{L}_{\text{DFM}}(\theta) \gets \sum_{j=1}^d \Big[-\log p^\theta_{1|t}(z_j\mid x)_j\Big]

- 更新参数:θOpt.step(θLDFM(θ))\theta \gets \textsc{Opt.step}\big(\nabla_\theta \mathcal{L}_{\text{DFM}}(\theta)\big)

- 结束循环
Figure 20掩码扩散语言模型轨迹的示意图。

这完成了训练和采样 CTMC 模型的完整流程,使我们能够生成文本等离散序列。当前最先进的 Discrete Diffusion Model [4] 使用了本工作中描述的配方,其神经网络(通常是 Transformer)在 Web 规模的数据上进行训练。

备注 40(发电机匹配)

你可能想知道为什么 flow/diffusion 模型的原理能够如此无缝地迁移到离散状态空间。事实证明,Flow Matching 的原理并非 flow 或 CTMC 所独有。相反,这些是构建基于Markov process的 Generative Model 的一般学习原理。这一思想引出了 Generator Matching 框架 [19],该框架将离散和连续的 flow 与 diffusion 模型扩展并统一为一个框架。Generator 是 Vector Field utu_t 和 Rate Matrix QtQ_t 的推广。Markov process 和 generator 可以为任何数据模态和状态空间构建。例如,你可以为光滑流形 [8, 10](如几何数据)、混合状态空间(如文本和图像的联合生成)[6] 以及其他 Markov process(如跳跃过程)[19, 7] 构建模型。

参考文献

书目信息按原课件保留。

  1. Michael S Albergo, Nicholas M Boffi, and Eric Vanden-Eijnden. “Stochastic interpolants: A unifying frame- work for flows and diffusions”. In: arXiv preprint arXiv:2303.08797 (2023).
  2. Brian DO Anderson. “Reverse-time diffusion equation models”. In: Stochastic Processes and their Applications 12.3 (1982), pp. 313–326.
  3. Yogesh Balaji et al. eDiff-I: Text-to-Image Diffusion Models with an Ensemble of Expert Denoisers. 2023. arXiv: 2211.01324 [cs.CV]. url: https://arxiv.org/abs/2211.01324.
  4. Tiwei Bie et al. “Llada2. 0: Scaling up diffusion language models to 100b”. In: arXiv preprint arXiv:2512.15745 (2025).
  5. Andrew Campbell et al. “A continuous time framework for discrete denoising models”. In: Advances in Neural Information Processing Systems 35 (2022), pp. 28266–28279.
  6. Andrew Campbell et al. “Generative flows on discrete state-spaces: Enabling multimodal flows with applica- tions to protein co-design”. In: arXiv preprint arXiv:2402.04997 (2024).
  7. Andrew Campbell et al. “Trans-dimensional generative modeling via jump diffusion models”. In: Advances in Neural Information Processing Systems 36 (2023), pp. 42217–42257.
  8. Ricky TQ Chen and Yaron Lipman. “Flow matching on general geometries”. In: arXiv preprint arXiv:2302.03660 (2023).
  9. Earl A Coddington, Norman Levinson, and T Teichmann. Theory of ordinary differential equations. 1956.
  10. Valentin De Bortoli et al. “Riemannian score-based generative modelling”. In: Advances in neural information processing systems 35 (2022), pp. 2406–2422.
  11. Prafulla Dhariwal and Alex Nichol. Diffusion Models Beat GANs on Image Synthesis. 2021. arXiv: 2105.05233 [cs.LG]. url: https://arxiv.org/abs/2105.05233.
  12. Alexey Dosovitskiy. “An image is worth 16x16 words: Transformers for image recognition at scale”. In: arXiv preprint arXiv:2010.11929 (2020).
  13. Alexey Dosovitskiy et al. An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale. 2021. arXiv: 2010.11929 [cs.CV]. url: https://arxiv.org/abs/2010.11929.
  14. Patrick Esser et al. Scaling Rectified Flow Transformers for High-Resolution Image Synthesis. 2024. arXiv: 2403.03206 [cs.CV]. url: https://arxiv.org/abs/2403.03206.
  15. Lawrence C Evans. Partial differential equations. Vol. 19. American Mathematical Society, 2022.
  16. Itai Gat et al. “Discrete flow matching”. In: Advances in Neural Information Processing Systems 37 (2024), pp. 133345–133385.
  17. Jonathan Ho, Ajay Jain, and Pieter Abbeel. “Denoising diffusion probabilistic models”. In: Advances in neural information processing systems 33 (2020), pp. 6840–6851.
  18. Jonathan Ho and Tim Salimans. Classifier-Free Diffusion Guidance. 2022. arXiv: 2207.12598 [cs.LG]. url: https://arxiv.org/abs/2207.12598.
  19. Peter Holderrieth et al. “Generator matching: Generative modeling with arbitrary markov processes”. In: arXiv preprint arXiv:2410.20587 (2024).
  20. Peter Holderrieth et al. “GLASS Flows: Transition Sampling for Alignment of Flow and Diffusion Models”. In: arXiv preprint arXiv:2509.25170 (2025).
  21. Arieh Iserles. A first course in the numerical analysis of differential equations. Cambridge university press, 2009.
  22. Alexia Jolicoeur-Martineau et al. “Adversarial score matching and improved sampling for image generation”. In: arXiv preprint arXiv:2009.05475 (2020).
  23. Tero Karras et al. “Elucidating the design space of diffusion-based generative models”. In: Advances in Neural Information Processing Systems 35 (2022), pp. 26565–26577.
  24. Samuel Lavoie et al. Modeling Caption Diversity in Contrastive Vision-Language Pretraining. 2024. arXiv: 2405.00740 [cs.CV]. url: https://arxiv.org/abs/2405.00740.
  25. Yaron Lipman et al. “Flow matching for generative modeling”. In: arXiv preprint arXiv:2210.02747 (2022).
  26. Yaron Lipman et al. “Flow Matching Guide and Code”. In: arXiv preprint arXiv:2412.06264 (2024).
  27. Xingchao Liu, Chengyue Gong, and Qiang Liu. “Flow straight and fast: Learning to generate and transfer data with rectified flow”. In: arXiv preprint arXiv:2209.03003 (2022).
  28. Nanye Ma et al. “Sit: Exploring flow and diffusion-based generative models with scalable interpolant trans- formers”. In: arXiv preprint arXiv:2401.08740 (2024).
  29. Xuerong Mao. Stochastic differential equations and applications. Elsevier, 2007.
  30. William Peebles and Saining Xie. Scalable Diffusion Models with Transformers. 2023. arXiv: 2212 . 09748 [cs.CV]. url: https://arxiv.org/abs/2212.09748.
  31. Ethan Perez et al. “Film: Visual reasoning with a general conditioning layer”. In: Proceedings of the AAAI conference on artificial intelligence. Vol. 32. 1. 2018.
  32. Lawrence Perko. Differential equations and dynamical systems. Vol. 7. Springer Science & Business Media, 2013.
  33. Adam Polyak et al. Movie Gen: A Cast of Media Foundation Models. 2024. arXiv: 2410.13720 [cs.CV]. url: https://arxiv.org/abs/2410.13720.
  34. Alec Radford et al. Learning Transferable Visual Models From Natural Language Supervision. 2021. arXiv: 2103.00020 [cs.CV]. url: https://arxiv.org/abs/2103.00020.
  35. Colin Raffel et al. Exploring the Limits of Transfer Learning with a Unified Text-to-Text Transformer. 2023. arXiv: 1910.10683 [cs.LG]. url: https://arxiv.org/abs/1910.10683.
  36. Robin Rombach et al. High-Resolution Image Synthesis with Latent Diffusion Models. 2022. arXiv: 2112.10752 [cs.CV]. url: https://arxiv.org/abs/2112.10752.
  37. Robin Rombach et al. “High-resolution image synthesis with latent diffusion models”. In: Proceedings of the IEEE/CVF conference on computer vision and pattern recognition. 2022, pp. 10684–10695.
  38. Olaf Ronneberger, Philipp Fischer, and Thomas Brox. “U-net: Convolutional networks for biomedical image segmentation”. In: Medical image computing and computer-assisted intervention–MICCAI 2015: 18th inter- national conference, Munich, Germany, October 5-9, 2015, proceedings, part III 18. Springer. 2015, pp. 234– 241.
  39. Chitwan Saharia et al. Photorealistic Text-to-Image Diffusion Models with Deep Language Understanding. 2022. arXiv: 2205.11487 [cs.CV]. url: https://arxiv.org/abs/2205.11487.
  40. Simo Särkkä and Arno Solin. Applied stochastic differential equations. Vol. 10. Cambridge University Press, 2019.
  41. Jascha Sohl-Dickstein et al. “Deep unsupervised learning using nonequilibrium thermodynamics”. In: Inter- national conference on machine learning. PMLR. 2015, pp. 2256–2265.
  42. Yang Song and Stefano Ermon. “Generative modeling by estimating gradients of the data distribution”. In: Advances in neural information processing systems 32 (2019).
  43. Yang Song et al. Score-Based Generative Modeling through Stochastic Differential Equations. 2021. arXiv: 2011.13456 [cs.LG]. url: https://arxiv.org/abs/2011.13456.
  44. Yang Song et al. “Score-Based Generative Modeling through Stochastic Differential Equations”. In: Interna- tional Conference on Learning Representations (ICLR). 2021.
  45. Yang Song et al. “Score-based generative modeling through stochastic differential equations”. In: arXiv preprint arXiv:2011.13456 (2020).
  46. Matthew Tancik et al. Fourier Features Let Networks Learn High Frequency Functions in Low Dimensional Domains. 2020. arXiv: 2006.10739 [cs.CV]. url: https://arxiv.org/abs/2006.10739.
  47. Yi Tay et al. UL2: Unifying Language Learning Paradigms. 2023. arXiv: 2205.05131 [cs.CL]. url: https: //arxiv.org/abs/2205.05131.
  48. Arash Vahdat, Karsten Kreis, and Jan Kautz. “Score-based generative modeling in latent space”. In: Advances in neural information processing systems 34 (2021), pp. 11287–11302.
  49. Ashish Vaswani et al. Attention Is All You Need. 2023. arXiv: 1706.03762 [cs.CL]. url: https://arxiv.org/ abs/1706.03762.
  50. Linting Xue et al. ByT5: Towards a token-free future with pre-trained byte-to-byte models. 2022. arXiv: 2105. 13626 [cs.CL]. url: https://arxiv.org/abs/2105.13626.
  51. Jingfeng Yao, Bin Yang, and Xinggang Wang. “Reconstruction vs. generation: Taming optimization dilemma in latent diffusion models”. In: Proceedings of the Computer Vision and Pattern Recognition Conference. 2025, pp. 15703–15712.

概率论回顾

我们简要概述概率论的基本概念。本节部分内容取自 [26]。

随机向量

考虑 dd 维欧几里得空间 x=(x1,,xd)Rdx=(x^1,\ldots,x^d)\in \Real^d 中的数据,带有标准欧几里得内积 x,y=i=1dxiyi\ip{x,y}=\sum_{i=1}^d x^i y^i 和范数 x=x,x\norm{x}=\sqrt{\ip{x,x}}

我们将考虑具有连续概率密度函数(PDF)的随机变量(RV)XRdX\in\R^d,定义为连续函数 pX:RdR0p_X:\Real^d\too \Real_{\geq 0},为事件 AA 提供概率

P(XA)=ApX(x)dx,(96)\sP(X\in A) = \int_A p_X(x) \dd x, \tag{96}

其中 pX(x)dx=1\int p_X(x)\dd x = 1

按照惯例,当在整个空间上积分时,我们省略积分区间(Rd\int \equiv \int_{\Real^d})。

为了保持符号简洁,我们将随机变量XtX_t 的概率密度函数pXtp_{X_t} 简称为ptp_t

我们将使用符号XpX \sim pXp(X)X \sim p(X) 来表示XX 服从pp 分布。

生成建模中一个常见的概率密度函数是dd 维各向同性 Gaussian distribution:

N(x;μ,σ2I)=(2πσ2)d2exp(xμ222σ2),(97)\gN(x;\mu,\sigma^2 I) = (2\pi\sigma^2)^{-\frac{d}{2}}\exp\left(-\frac{\norm{x-\mu}_2^2}{2\sigma^2}\right), \tag{97}

其中μRd\mu\in \Real^dσR>0\sigma \in \Real_{>0} 分别表示分布的均值和标准差。

随机变量的期望是在最小二乘意义上最接近XX 的常数向量:

E[X]=arg minzRdxz2pX(x)dx=xpX(x)dx.(98)\E\brac{X}=\argmin_{z\in\Real^d} \int \norm{x-z}^2 p_X(x)\dd x = \int x p_X(x)\dd x. \tag{98}

计算随机变量函数期望的一个有用工具是无意识统计学家定律

E[f(X)]=f(x)pX(x)dx.(99)\E \brac{f(X)} = \int f(x) p_X(x) \dd x. \tag{99}

必要时,我们将在期望下标中注明随机变量为EXf(X)\E_{X} f(X)

条件密度与条件期望

Figure 21联合概率密度函数pX,Yp_{X,Y}(以阴影显示)及其边缘分布pXp_XpYp_Y(以黑色线条显示)。图来自[26]

给定两个随机变量X,YRdX,Y\in \Real^d,它们的联合概率密度函数pX,Y(x,y)p_{X,Y}(x,y) 具有边缘分布

pX,Y(x,y)dy=pX(x) and pX,Y(x,y)dx=pY(y).(100)\int p_{X,Y}(x,y)\dd y = p_X(x) \text{ and } \int p_{X,Y}(x,y)\dd x = p_Y(y). \tag{100}

参见图 21,其中展示了R\Reald=1d=1)中两个随机变量的联合概率密度函数的示例。

条件概率密度函数pXYp_{X|Y} 描述了在事件Y=yY=y(其密度为pY(y)>0p_Y(y)>0)条件下,随机变量XX 的概率密度函数:

pXY(xy):=pX,Y(x,y)pY(y),(101)p_{X|Y}(x|y)\defe\frac{p_{X,Y}(x,y)}{p_Y(y)}, \tag{101}

类似地,对于条件概率密度函数pYXp_{Y|X}。Bayes' rule 用pXYp_{X|Y} 表达了条件概率密度函数pYXp_{Y|X}

pYX(yx)=pXY(xy)pY(y)pX(x),(102)p_{Y|X}(y|x) = \frac{p_{X|Y}(x|y)p_Y(y)}{p_X(x)}, \tag{102}

对于 pX(x)>0p_X(x)>0

条件期望 E[XY]\E\brac{X | Y} 是在最小二乘意义上对 XX 的最佳逼近函数 g(Y)g_\star(Y)

g:=arg ming:RdRdE[Xg(Y)2]=arg ming:RdRdxg(y)2pX,Y(x,y)dxdy\begin{aligned}g_\star &\defe \argmin_{g:\Real^d\too\Real^d}\E\brac{\norm{X-g(Y)}^2} = \argmin_{g:\Real^d\too\Real^d}\int \norm{x-g(y)}^2 p_{X,Y}(x,y)\dd x \dd y\end{aligned}
=arg ming:RdRd[xg(y)2pXY(xy)dx]pY(y)dy.(103)\begin{aligned}&= \argmin_{g:\Real^d\too\Real^d} \int\brac{\textcolor{black}{\int \norm{x-g(y)}^2p_{X|Y}(x|y)\dd x}} p_Y(y)\dd y.\end{aligned}\tag{103}

对于满足 pY(y)>0p_Y(y)>0yRdy\in \Real^d,条件期望函数因此为

E[XY=y]:=g(y)=xpXY(xy)dx,(104)\E\brac{X|Y=y} \defe g_\star(y) = \int x p_{X|Y}(x|y) \dd x, \tag{104}

其中第二个等式来自对式(103)中内括号关于 Y=yY=y 取最小化,类似于式(98)。

gg_\star 与随机变量 YY 复合,我们得到

E[XY]:=g(Y),(105)\E\brac{X|Y} \defe g_\star(Y), \tag{105}

这是 Rd\Real^d 中的一个随机变量。

令人困惑的是,E[XY=y]\E\brac{X|Y=y}E[XY]\E\brac{X|Y} 都常被称为条件期望,但它们是不同的对象。

特别地,E[XY=y]\E\brac{X|Y=y} 是一个函数 RdRd\Real^d\too\Real^d,而 E[XY]\E\brac{X|Y} 是一个取值于 Rd\Real^d 的随机变量。

为了区分这两个术语,我们的讨论将采用这里引入的记号。

塔性质是一个有用的性质,有助于简化涉及两个随机变量 XXYY 的条件期望的推导:

E[E[XY]]=E[X](106)\E\brac{\E\brac{X|Y}} = \E\brac{X} \tag{106}

因为 E[XY]\E\brac{X|Y} 是一个随机变量,它本身是随机变量 YY 的函数,外层期望计算的是 E[XY]\E\brac{X|Y} 的期望。

塔性质可以通过使用上述一些定义来验证:

E[E[XY]]=(xpXY(xy)dx)pY(y)dy=式(101)xpX,Y(x,y)dxdy=式(100)xpX(x)dx=E[X].\begin{aligned}\E\brac{\E\brac{X|Y}} &= \int \parr{\int x p_{X|Y}(x|y) \dd x} p_Y(y) \dd y \\ &\overset{\text{式(101)}}{=} \int \int x p_{X,Y}(x,y) \dd x\dd y \\ &\overset{\text{式(100)}}{=} \int x p_X(x)\dd x= \E \brac{X}.\end{aligned}

最后,考虑一个涉及两个随机变量f(X,Y)f(X, Y)YY 的有用性质,其中XXYY 是两个任意随机变量。

然后,通过使用无意识统计学家定律与式(104),我们得到恒等式

E[f(X,Y)Y=y]=f(x,y)pXY(xy)dx.(107)\E\brac{f(X,Y)|Y=y} = \int f(x,y) p_{X|Y}(x|y) \dd x. \tag{107}

Fokker–Planck 方程的证明

在本节中,我们给出 Fokker-Planck 方程的自包含证明,该方程将连续性方程作为特例(第 3.2 节)。我们强调,本节对于理解本文档的其余部分并非必需,且在数学上更为高级。如果你希望理解 Fokker-Planck 方程的来源,那么本节正是为你准备的。

定理 41(Fokker-Planck 方程)

ptp_t 为 Probability Path,其中p0=pinitp_0=\pinit,并考虑 SDE

X0pinit,dXt=ut(Xt)dt+σtdWt.X_0\sim \pinit, \quad \dd X_t = u_t(X_t)\dd t + \sigma_t\dd W_t.

那么,当且仅当 Fokker-Planck 方程成立时,XtX_t 对所有 0t10\leq t\leq 1 具有分布 ptp_t

tpt(x)=div(ptut)(x)+σt22Δpt(x) for all xRd,0t1,(108)\partial_t p_t(x) = -\divv (p_t u_t)(x)+\frac{\sigma_t^2}{2}\Delta p_t (x)\quad \text{ for all }x\in\R^d, 0\leq t\leq 1, \tag{108}

我们首先证明 Fokker-Planck 是必要条件,即如果XtptX_t\sim p_t,则 Fokker-Planck 方程成立。证明的技巧是使用测试函数ff,即函数f:RdRf:\R^d\to\R 是无限可微的(“光滑的”),并且仅在有限域内非零(紧支撑)。我们利用以下事实:对于任意可积函数g1,g2:RdRg_1,g_2:\mathbb{R}^d\to\mathbb{R},有

g1(x)=g2(x) for all xRdf(x)g1(x)dx=f(x)g2(x)dx for all test functions f(109)g_1(x) = g_2(x)\text{ for all }x\in\mathbb{R}^d\quad \Leftrightarrow \quad \int f(x) g_1(x)\dd x = \int f(x) g_2(x)\dd x\text{ for all test functions }f \tag{109}

换句话说,我们可以将逐点相等表示为积分相等。测试函数的有用之处在于它们是光滑的,即我们可以取梯度和高阶导数。特别地,我们可以对任意测试函数f1,f2f_1,f_2 使用分部积分

f1(x)xif2(x)dx=f2(x)xif1(x)dx(110)\int f_1(x) \frac{\partial}{\partial x_i}f_2(x)\dd x = - \int f_2(x) \frac{\partial}{\partial x_i}f_1(x)\dd x \tag{110}

f1,f2f_1,f_2 及其乘积f1f2f_1\cdot f_2 可积的条件下。通过将此与散度和拉普拉斯算子的定义(见式(22))结合,我们得到恒等式:

f1T(x)f2(x)dx=f1(x)div(f2)(x)dx(f1:RdR,f2:RdRd)(111)\begin{aligned}\int \nabla f_1^T(x) f_2(x)\dd x =& -\int f_1(x)\divv(f_2)(x)\dd x \quad (f_1:\R^d\to\R,f_2:\R^d\to\R^d)\end{aligned}\tag{111}
f1(x)Δf2(x)dx=f2(x)Δf1(x)dx(f1:RdR,f2:RdR)(112)\begin{aligned}\int f_1(x) \Delta f_2(x) \dd x =& \int f_2(x) \Delta f_1(x) \dd x\quad (f_1:\R^d\to\R,f_2:\R^d\to\R)\end{aligned}\tag{112}

现在让我们进行证明。我们使用式(6)中 SDE 轨迹的随机更新:

Xt+h=Xt+hut(Xt)+σt(Wt+hWt)+hRt(h)(113)\begin{aligned}X_{t+h} =& X_{t}+hu_t(X_t) + \sigma_t(W_{t+h}-W_{t})+hR_t(h)\end{aligned}\tag{113}
Xt+hut(Xt)+σt(Wt+hWt)(114)\begin{aligned}\approx &X_{t}+hu_t(X_t) + \sigma_t(W_{t+h}-W_{t})\end{aligned}\tag{114}

其中目前我们为了可读性忽略误差项Rt(h)R_t(h),因为我们最终会取h0h\to 0。然后我们可以进行如下计算:

f(Xt+h)f(Xt)=式(114)f(Xt+hut(Xt)+σt(Wt+hWt))f(Xt)=(i)f(Xt)T(hut(Xt)+σt(Wt+hWt)))+12(hut(Xt)+σt(Wt+hWt)))T2f(Xt)(hut(Xt)+σt(Wt+hWt)))=(ii)hf(Xt)Tut(Xt)+σtf(Xt)T(Wt+hWt)+12h2ut(Xt)T2f(Xt)ut(Xt)+hσtut(Xt)T2f(Xt)(Wt+hWt)++12σt2(Wt+hWt)T2f(Xt)(Wt+hWt)\begin{aligned}f(X_{t+h})-f(X_t)\overset{\text{式(114)}}{=}&f(X_{t}+hu_t(X_t) + \sigma_t(W_{t+h}-W_{t}))-f(X_t)\\ \overset{(i)}{=}&\nabla f(X_t)^T\left(hu_t(X_t) + \sigma_t(W_{t+h}-W_{t}))\right)\\&+\frac{1}{2}\left(hu_t(X_t) + \sigma_t(W_{t+h}-W_{t}))\right)^T\nabla^2 f(X_t)\left(hu_t(X_t) + \sigma_t(W_{t+h}-W_{t}))\right)\\ \overset{(ii)}{=}&h\nabla f(X_t)^Tu_t(X_t) + \sigma_t\nabla f(X_t)^T(W_{t+h}-W_{t})\\&+\frac{1}{2}h^2u_t(X_t)^T\nabla^2 f(X_t)u_t(X_t)+ h\sigma_tu_t(X_t)^T\nabla^2 f(X_t)(W_{t+h}-W_{t})+\\&+\frac{1}{2}\sigma_t^2(W_{t+h}-W_t)^T\nabla^2 f(X_t)(W_{t+h}-W_t)\end{aligned}

其中在(i) 中我们使用了ffXtX_t 附近的二阶泰勒近似,在(ii) 中我们使用了 Hessian 矩阵2f\nabla^2 f 是对称矩阵的事实。注意E[Wt+hWtXt]=0\mathbb{E}[W_{t+h}-W_t|X_t]=0Wt+hWtXtN(0,hId)W_{t+h}-W_{t}|X_t\sim\mathcal{N}(0,hI_d)。因此

E[f(Xt+h)f(Xt)Xt]=hf(Xt)Tut(Xt)+12h2ut(Xt)T2f(Xt)ut(Xt)+h2σt2EϵtN(0,Id)[ϵtT2f(Xt)ϵt]=(i)hf(Xt)Tut(Xt)+12h2ut(Xt)T2f(Xt)ut(Xt)+h2σt2trace(2f(Xt))=(ii)hf(Xt)Tut(Xt)+12h2ut(Xt)T2f(Xt)ut(Xt)+h2σt2Δf(Xt)\begin{aligned}&\mathbb{E}[f(X_{t+h})-f(X_t)|X_t]\\ =&h\nabla f(X_t)^Tu_t(X_t)+\frac{1}{2}h^2u_t(X_t)^T\nabla^2 f(X_t)u_t(X_t)+\frac{h}{2}\sigma_t^2\mathbb{E}_{\epsilon_t\sim\mathcal{N}(0,I_d)}[\epsilon_t^T\nabla^2 f(X_t)\epsilon_t]\\ \overset{(i)}{=}&h\nabla f(X_t)^Tu_t(X_t)+\frac{1}{2}h^2u_t(X_t)^T\nabla^2 f(X_t)u_t(X_t)+\frac{h}{2}\sigma_t^2\text{trace}(\nabla^2 f(X_t))\\ \overset{(ii)}{=}&h\nabla f(X_t)^Tu_t(X_t)+\frac{1}{2}h^2u_t(X_t)^T\nabla^2 f(X_t)u_t(X_t)+\frac{h}{2}\sigma_t^2\Delta f(X_t)\end{aligned}

其中在(i)(i) 中我们利用了EϵtN(0,Id)[ϵtTAϵt]=trace(A)\mathbb{E}_{\epsilon_t\sim\mathcal{N}(0,I_d)}[\epsilon_t^T A\epsilon_t]=\text{trace}(A) 这一事实,在(ii)(ii) 中我们使用了拉普拉斯算子和 Hessian 矩阵的定义。由此我们得到

tE[f(Xt)]=limh01hE[f(Xt+h)f(Xt)]=limh01hE[E[f(Xt+h)f(Xt)Xt]]=E[limh01h(hf(Xt)Tut(Xt)+12h2ut(Xt)T2f(Xt)ut(Xt)+h2σt2Δf(Xt))]=E[f(Xt)Tut(Xt)+12σt2Δf(Xt)]=(i)f(x)Tut(x)pt(x)dx+12σt2Δf(x)pt(x)dx=(ii)f(x)div(utpt)(x)dx+12σt2f(x)Δpt(x)dx=f(x)(div(utpt)(x)+12σt2Δpt(x))dx\begin{aligned}&\partial_t \mathbb{E}[f(X_t)]\\ =&\lim\limits_{h\to 0} \frac{1}{h}\mathbb{E}[f(X_{t+h})-f(X_t)]\\ =&\lim\limits_{h\to 0} \frac{1}{h}\mathbb{E}[\mathbb{E}[f(X_{t+h})-f(X_t)|X_t]]\\ =&\mathbb{E}[\lim\limits_{h\to 0}\frac{1}{h}\left( h\nabla f(X_t)^Tu_t(X_t)+\frac{1}{2}h^2u_t(X_t)^T\nabla^2 f(X_t)u_t(X_t)+\frac{h}{2}\sigma_t^2\Delta f(X_t) \right)]\\ =&\mathbb{E}[\nabla f(X_t)^Tu_t(X_t)+\frac{1}{2}\sigma_t^2\Delta f(X_t)]\\ \overset{(i)}{=}&\int \nabla f(x)^Tu_t(x)p_t(x)\dd x+\int \frac{1}{2}\sigma_t^2\Delta f(x)p_t(x)\dd x\\ \overset{(ii)}{=}&-\int f(x)\divv(u_t p_t)(x)\dd x+\int \frac{1}{2}\sigma_t^2 f(x)\Delta p_t(x)\dd x\\ =&\int f(x)\left(-\divv(u_t p_t)(x)+\frac{1}{2}\sigma_t^2\Delta p_t(x)\right)\dd x\end{aligned}

其中在(i) 中我们使用了假设ptp_t 作为XtX_t 的分布,在(ii) 中我们使用了式(111)和式(112)。注意,要使用这一点,我们需要乘积pt(x)ut(x)p_t(x)u_t(x) 的可积性,即满足

pt(x)ut(x)dx<\int p_t(x)\|u_t(x)\|\dd x <\infty

注意,在机器学习中这一条件几乎总是成立(由于数值精度限制,数据和函数有界)。因此,以下成立

tE[f(Xt)]=f(x)(div(ptut)(x)+σt22Δpt(x))dx(for all f and 0t1)(115)\begin{aligned}\partial_t\mathbb{E}[f(X_t)] =& \int f(x)\left(-\divv (p_t u_t)(x)+\frac{\sigma_t^2}{2}\Delta p_t (x)\right)\dd x\quad (\text{for all }f\text{ and }0\leq t\leq 1)\end{aligned}\tag{115}
(i)tf(x)pt(x)dx=f(x)(div(ptut)(x)+σt22Δpt(x))dx(for all f and 0t1)(116)\begin{aligned}\overset{(i)}{\Leftrightarrow}\quad \partial_t\int f(x) p_t(x)\dd x =& \int f(x)\left(-\divv (p_tu_t)(x)+\frac{\sigma_t^2}{2}\Delta p_t (x)\right)\dd x\quad (\text{for all }f\text{ and }0\leq t\leq 1)\end{aligned}\tag{116}
(ii)f(x)tpt(x)dx=f(x)(div(ptut)(x)+σt22Δpt(x))dx(for all f and 0t1)(117)\begin{aligned}\overset{(ii)}{\Leftrightarrow}\quad \int f(x)\partial_t p_t(x)\dd x =& \int f(x)\left(-\divv (p_t u_t)(x)+\frac{\sigma_t^2}{2}\Delta p_t (x)\right)\dd x\quad (\text{for all }f\text{ and }0\leq t\leq 1)\end{aligned}\tag{117}
(iii)tpt(x)=div(ptut)(x)+σt22Δpt(x)(for all xRd,0t1)(118)\begin{aligned}\overset{(iii)}{\Leftrightarrow}\quad \partial_t p_t(x) =& -\divv (p_t u_t)(x)+\frac{\sigma_t^2}{2}\Delta p_t(x)\quad (\text{for all }x\in\R^d, 0\leq t\leq 1)\end{aligned}\tag{118}

其中在(i) 中我们使用了假设XtptX_t\sim p_t,在(ii) 中我们交换了导数与积分,在(iii) 中我们使用了式(109)。这完成了 Fokker-Planck 方程是必要条件的证明。

最后,我们解释为什么它也是充分条件。Fokker-Planck 方程是一个偏微分方程(PDE)。更具体地说,它是一个所谓的抛物型偏微分方程。类似于第 2.1 节,这类微分方程在给定初始条件下有唯一解(参见 [15])。现在,如果式(108)对ptp_t 成立,我们刚刚在上面证明了它也必须对qtq_t 的真实分布XtX_t 成立(即XtqtX_t\sim q_t)——换句话说,ptp_tqtq_t 都是该抛物型 PDE 的解。进一步,我们知道初始条件是相同的,即p0=q0=pinitp_0=q_0=\pinit,这是通过构造插值 Probability Path 得到的。因此,根据微分方程解的唯一性,我们知道对所有0t10\leq t\leq 1pt=qtp_t=q_t——这意味着Xtqt=ptX_t\sim q_t=p_t,这正是我们想要证明的。

连续时间 Markov chain 的存在性与唯一性

我们在本节中证明第 7.1 节。

证明

唯一性:我们需要证明只有一个转移核ptt(Xt=yXt=x)p_{t'|t}(X_{t'}=y|X_{t}=x) 满足式(87)。作为第一步,我们意识到式(87)意味着

ddtptt(Xt=yXt=x)(119)\begin{aligned}&\frac{\dd}{\dd t'}p_{t'|t}(X_{t'}=y|X_{t}=x)\end{aligned}\tag{119}
=ddhpt+ht(Xt+h=yXt=x)h=0(120)\begin{aligned}=&\frac{\dd}{\dd h}p_{t'+h|t}(X_{t'+h}=y|X_{t}=x)_{|h=0}\end{aligned}\tag{120}
=ddh[zSpt+ht(Xt+h=yXt=z)ptt(Xt=zXt=x)]h=0(121)\begin{aligned}=&\frac{\dd}{\dd h}\left[\sum\limits_{z\in S}p_{t'+h|t'}(X_{t'+h}=y|X_{t'}=z)p_{t'|t}(X_{t'}=z|X_{t}=x)\right]_{|h=0}\end{aligned}\tag{121}
=zSQt(yz)ptt(Xt=zXt=x)(122)\begin{aligned}=&\sum\limits_{z\in S}Q_{t'}(y|z)p_{t'|t}(X_{t'}=z|X_{t}=x)\end{aligned}\tag{122}

对于固定的x,tx,t,可以将tptt(Xt=yXt=x)t'\mapsto p_{t'|t}(X_{t'}=y|X_t=x) 视为向量值函数,上述是该函数的线性 ODE(实际上是 Kolmogorov 前向方程,见第 7.2.2 节),且具有已知的初始条件,即ptt(Xt=yXt=x)=δy(x)p_{t|t}(X_{t}=y|X_{t}=x)=\delta_{y}(x)。如我们所知,每个线性 ODE 都有唯一解(见第 2.1 节),因此ptt(Xt=yXt=x)p_{t'|t}(X_{t'}=y|X_{t}=x) 也必须是唯一的。

存在性:反之,任何线性 ODE 都有解,即我们知道对于每个x,tx,t,存在一个ptt(Xt=yXt=x)p_{t'|t}(X_{t'}=y|X_{t}=x) 使得

ptt(Xt=yXt=x)=δy(x)(123)\begin{aligned}p_{t|t}(X_{t}=y|X_{t}=x)&=\delta_{y}(x)\end{aligned}\tag{123}
ddtptt(Xt=yXt=x)=zSQt(yz)ptt(Xt=zXt=x)(124)\begin{aligned}\frac{\dd}{\dd t'}p_{t'|t}(X_{t'}=y|X_{t}=x)&=\sum\limits_{z\in S}Q_{t'}(y|z)p_{t'|t}(X_{t'}=z|X_t=x)\end{aligned}\tag{124}

对于t=tt'=t,这特别意味着式(87)。还需要证明在这种情况下ptt(Xt=yXt=x)p_{t'|t}(X_{t'}=y|X_{t}=x) 是有效的转移核,即以下三个性质必须成立:

ySptt(Xt=yXt=x)=1(125)\begin{aligned}\sum\limits_{y\in S}p_{t'|t}(X_{t'}=y|X_{t}=x)=&1\end{aligned}\tag{125}
ptt(Xt=yXt=x)0(126)\begin{aligned}p_{t'|t}(X_{t'}=y|X_{t}=x)\geq& 0\end{aligned}\tag{126}
zSpt2t1(Xt2=yXt1=z)pt1t0(Xt1=zXt0=x)=pt2t0(yx)(127)\begin{aligned}\sum\limits_{z\in S}p_{t_{2}|t_{1}}(X_{t_2}=y|X_{t_1}=z)p_{t_1|t_0}(X_{t_1}=z|X_{t_0}=x)=&p_{t_2|t_0}(y|x)\end{aligned}\tag{127}

对于第一个性质,可以观察到它对t=tt'=t 成立,由式(123)以及

ddtySptt(Xt=yXt=x)(128)\begin{aligned}&\frac{\dd}{\dd t'}\sum\limits_{y\in S}p_{t'|t}(X_{t'}=y|X_{t}=x)\end{aligned}\tag{128}
=ySddtptt(Xt=yXt=x)(129)\begin{aligned}=&\sum\limits_{y\in S}\frac{\dd}{\dd t'}p_{t'|t}(X_{t'}=y|X_{t}=x)\end{aligned}\tag{129}
=zS[ySQt(yz)]ptt(Xt=zXt=x)(130)\begin{aligned}=&\sum\limits_{z\in S}\left[\sum\limits_{y\in S}Q_{t'}(y|z)\right]p_{t'|t}(X_{t'}=z|X_t=x)\end{aligned}\tag{130}
=0(131)\begin{aligned}=&0\end{aligned}\tag{131}

这里我们利用了 Rate Matrix 的列之和为 00 这一事实。为了证明第二个性质,注意到它在时刻 t=tt'=t 成立。进一步,每当 ptt(Xt=yXt=x)=0p_{t'|t}(X_{t'}=y|X_{t}=x)=0 时,必有

ddtptt(Xt=yXt=x)=zyQt(yz)0ptt(Xt=zXt=x)0\begin{aligned}\frac{\dd}{\dd t'}p_{t'|t}(X_{t'}=y|X_{t}=x)&=\sum\limits_{z\neq y}\underbrace{Q_{t'}(y|z)}_{\geq 0}p_{t'|t}(X_{t'}=z|X_t=x)\\ &\geq 0\end{aligned}

因此,每当 ptt(Xt=yXt=x)=0p_{t'|t}(X_{t'}=y|X_{t}=x)=0 时,它只能增加。因此,ptt(Xt=yXt=x)p_{t'|t}(X_{t'}=y|X_{t}=x) 永远不会为负。

为了证明第三个性质,定义 qt2t0(yx)q_{t_2|t_0}(y|x)

qt2t0(yx)=zSpt2t1(Xt2=yXt1=z)pt1t0(Xt1=zXt0=x)\begin{aligned}q_{t_2|t_0}(y|x)=&\sum\limits_{z\in S}p_{t_{2}|t_{1}}(X_{t_2}=y|X_{t_1}=z)p_{t_1|t_0}(X_{t_1}=z|X_{t_0}=x)\end{aligned}

那么我们知道

qt2=t1t0(yx)=zSδy(z)pt1t0(Xt1=zXt0=x)=pt1t0(Xt1=yXt0=x)\begin{aligned}q_{t_2=t_1|t_0}(y|x)=&\sum\limits_{z\in S}\delta_{y}(z)p_{t_1|t_0}(X_{t_1}=z|X_{t_0}=x)=p_{t_1|t_0}(X_{t_1}=y|X_{t_0}=x)\end{aligned}

以及

ddt2qt2t0(yx)=zSddt2pt2t1(Xt2=yXt1=z)pt1t0(Xt1=zXt0=x)=zSz~SQt2(yz~)pt2t1(Xt2=z~Xt1=z)pt1t0(Xt1=zXt0=x)=z~SQt2(yz~)[zSpt2t1(Xt2=z~Xt1=z)pt1t0(Xt1=zXt0=x)]=z~SQt2(yz~)qt2t0(z~x)\begin{aligned}\frac{\dd}{\dd t_2}q_{t_2|t_0}(y|x)=&\sum\limits_{z\in S}\frac{\dd}{\dd t_2}p_{t_{2}|t_{1}}(X_{t_2}=y|X_{t_1}=z)p_{t_1|t_0}(X_{t_1}=z|X_{t_0}=x)\\ =&\sum\limits_{z\in S}\sum\limits_{\tilde{z}\in S}Q_{t_2}(y|\tilde{z})p_{t_2|t_1}(X_{t_2}=\tilde{z}|X_{t_1}=z)p_{t_1|t_0}(X_{t_1}=z|X_{t_0}=x)\\ =&\sum\limits_{\tilde{z}\in S}Q_{t_2}(y|\tilde{z})\left[\sum\limits_{z\in S}p_{t_2|t_1}(X_{t_2}=\tilde{z}|X_{t_1}=z)p_{t_1|t_0}(X_{t_1}=z|X_{t_0}=x)\right]\\ =&\sum\limits_{\tilde{z}\in S}Q_{t_2}(y|\tilde{z})q_{t_2|t_0}(\tilde{z}|x)\end{aligned}

这表明 pt2t0(zx)p_{t_2|t_0}(z|x)qt2t0(zx)q_{t_2|t_0}(z|x) 满足相同的 ODE。因此,必有

zSpt2t1(Xt2=yXt1=z)pt1t0(Xt1=zXt0=x)=qt2t0(yx)=pt2t0(yx)\begin{aligned}\sum\limits_{z\in S}p_{t_{2}|t_{1}}(X_{t_2}=y|X_{t_1}=z)p_{t_1|t_0}(X_{t_1}=z|X_{t_0}=x)=&q_{t_2|t_0}(y|x)=p_{t_2|t_0}(y|x)\end{aligned}

这就证明了第三个性质。所以 ptt(yx)p_{t'|t}(y|x) 确实是满足式(87)的转移核。证明完毕。

VAE 的补充视角

在本节中,我们详细阐述正文中对 VAE 的处理,并给出式(83)中总 VAE 损失的变分推导。作为第一步,注意到编码器和解码器都给出了关于 xx 和 latent zz 的联合分布,即

qϕ(x,z)=pdata(x)qϕ(x)(encoder joint)pθ(x,z)=pθ(xz)pprior(z)(decoder joint)\begin{aligned}q_\phi(x,z) &= \pdata(x)q_\phi(\cdot |x) \quad& (\text{encoder joint}) \\ p_\theta(x,z) &= p_\theta(x|z)\prior(z) \quad& (\text{decoder joint})\end{aligned}

因此,我们可能将训练 VAE 概念化为学习 ϕ\phiθ\theta,使得编码器和解码器的联合分布相当相似。我们可以通过联合 latent 与数据分布的 KL 散度来实现这一点:

DKL ⁣(qϕ(x,z)pθ(x,z))=DKL ⁣(pdata(x)qϕ(zx)pθ(xz)pprior(z))=E[log(pdata(x)qϕ(zx)pθ(xz)pprior(z))]=E[logpdata(x)]+E[log(qϕ(zx)pprior(z))]E[logpθ(xz)]=xpdata(x)zqϕ(zx).(132)\begin{aligned}\dkl{q_\phi(x,z)}{p_\theta(x,z)} &= \dkl{\pdata(x)q_\phi(z\mid x)}{p_\theta(x\mid z)\prior(z)}\\ &= \EE_{\blacksquare} \left[\log \left(\frac{\pdata(x)q_\phi(z\mid x)}{p_\theta(x\mid z)\prior(z)}\right)\right]\\ &= \textcolor{BrickRed}{\EE_{\blacksquare} \left[\log \pdata(x)\right]} + \textcolor{RoyalBlue}{\EE_{\blacksquare} \left[\log \left(\frac{q_\phi(z\mid x)}{\prior(z)}\right)\right]} - \textcolor{ForestGreen}{\EE_{\blacksquare} \left[\log p_\theta(x\mid z)\right]}\\ \blacksquare &= x \sim \pdata(x)\, z\sim q_\phi(z|x).\end{aligned}\tag{132}

现在让我们依次检查剩余的三项。首先,我们发现

E[logpdata(x)]=Expdata(x)[logpdata(x)]=C,(133)\textcolor{BrickRed}{\EE_{\blacksquare} \left[\log \pdata(x)\right] = \textcolor{BrickRed}{\EE_{x\sim \pdata(x)} \left[\log \pdata(x)\right]} = C}, \tag{133}

对于某个独立于 ϕ\phiθ\theta 的常数 CC。接下来,我们发现

E[log(qϕ(zx)pprior(z))]=Expdata(x)[DKL ⁣(qϕ(zx)pprior(z))](134)\textcolor{RoyalBlue}{\EE_{\blacksquare} \left[\log \left(\frac{q_\phi(z\mid x)}{\prior(z)}\right)\right] = \mathbb{E}_{x \sim \pdata(x)}\left[\dkl{q_\phi(z\mid x)}{\prior(z)}\right]} \tag{134}

鼓励 qϕ(zx)q_\phi(z\mid x) 类似于先验 pprior(z)\prior(z)。最后,我们发现

Expdata(x)zqϕ(zx)[logpθ(xz)](135)-\textcolor{ForestGreen}{\EE_{x \sim \pdata(x)\, z\sim q_\phi(z|x)} \left[\log p_\theta(x\mid z)\right]} \tag{135}

对应平均负对数似然,因此用于最小化重建损失。忽略常数项,我们将先验惩罚项和重建项合并,得到 VAE 损失实际上就是联合数据和 latent space 上的 KL 散度:

LVAE(ϕ,θ)=Expdata(x)[DKL ⁣(qϕ(zx)pprior(z))]prior enforcement lossExpdata(x)zqϕ(zx)[logpθ(xz)]reconstruction loss(136)\begin{aligned}\mathcal{L}_{\text{VAE}}(\phi, \theta) &= \underbrace{\textcolor{RoyalBlue}{\mathbb{E}_{x \sim \pdata(x)}\left[\dkl{q_\phi(z\mid x)}{\prior(z)}\right]}}_{\text{prior enforcement loss}} - \underbrace{\textcolor{ForestGreen}{\EE_{x \sim \pdata(x)\, z\sim q_\phi(z|x)} \left[\log p_\theta(x\mid z)\right]}}_{\text{reconstruction loss}}\end{aligned}\tag{136}
=DKL ⁣(qϕ(x,z)pθ(x,z))+const(137)\begin{aligned}&=\dkl{q_\phi(x,z)}{p_\theta(x,z)}+\text{const}\end{aligned}\tag{137}

因此,我们可以将 VAE 解释为潜在变量和图像联合空间中的 KL 散度。

作为 Generative Model 的 VAE

我们现在解释如何将 VAE 视为 Generative Model。我们可以通过设置 zpprior=N(0,Ik)z\sim \prior=\mathcal{N}(0,I_k) 并从解码器中采样 xpθ(z)x\sim p_\theta(\cdot|z) 来生成样本。我们得到的分布由下式给出:

pθ(x)=zpθ(xz)pprior(z)dzp_\theta(x) = \int_z p_\theta(x|z) \prior(z)\dd z

我们现在想要证明 VAE 学习近似地从 pθp_\theta 中采样。为了证明这一点,我们需要以下结果:

命题 3(链式法则)

q(x,z),p(x,z)q(x,z), p(x,z) 是两个变量 xRl1,zRl2x\in\mathbb{R}^{l_1},z\in\mathbb{R}^{l_2} 上的分布。那么,以下成立:

DKL ⁣(q(z,x)p(z,x))=DKL ⁣(q(x)p(x))+Exq[DKL ⁣(q(zx)p(zx))].\dkl{q(z,x)}{p(z,x)}= \dkl{q(x)}{p(x)}+\mathbb{E}_{x\sim q}\left[\dkl{q(z|x)}{p(z|x)}\right].

特别地,由于第二项根据式(76)是非负的,我们得到数据处理不等式

DKL ⁣(q(x)p(x))DKL ⁣(q(z,x)p(z,x)).(138)\dkl{q(x)}{p(x)}\leq \dkl{q(z,x)}{p(z,x)}. \tag{138}
证明
DKL ⁣(q(z,x)p(z,x))=Eq[logq(z,x)p(z,x)]=E(x,z)q[logq(zx)p(zx)q(x)p(x)]=E(x,z)q[logq(zx)p(zx)]+Exq[logq(x)p(x)]=DKL ⁣(q(x)p(x))+Exq[DKL ⁣(q(zx)p(zx))]\begin{aligned}\dkl{q(z,x)}{p(z,x)}&=\mathbb{E}_{q}\left[ \log \frac{q(z,x)}{p(z,x)} \right]\\ &=\mathbb{E}_{(x,z) \sim q}\left[ \log \frac{q(z|x)}{p(z|x)}\frac{q(x)}{p(x)} \right]\\ &=\mathbb{E}_{(x,z) \sim q}\left[ \log \frac{q(z|x)}{p(z|x)}\right]+\mathbb{E}_{x\sim q} \left[\log \frac{q(x)}{p(x)} \right]\\ &=\dkl{q(x)}{p(x)}+\mathbb{E}_{x\sim q}\left[\dkl{q(z|x)}{p(z|x)}\right]\end{aligned}

其中我们反复应用了 KL 散度的定义。

根据第 D 节,我们现在可以证明

LVAE(ϕ,θ)=DKL ⁣(qϕ(x,z)pθ(x,z))+constDKL ⁣(pdata(x)pθ(x))+const(139)\mathcal{L}_{\text{VAE}}(\phi, \theta) = \dkl{q_\phi(x,z)}{p_\theta(x,z)} +\text{const}\ge \dkl{\pdata(x)}{p_\theta(x)} +\text{const} \tag{139}

其中我们使用了 qϕ(x,z)q_\phi(x,z)xx 边缘分布是 pdata\pdata 这一事实。换句话说,VAE 损失最小化了数据分布 pdata\pdata 与 VAE 生成的分布之间的 KL 散度的上界。因此,我们可以将 VAE 视为独立的 Generative Model。同样,我们可以证明

LVAE(ϕ,θ)=DKL ⁣(qϕ(x,z)pθ(x,z))+constDKL ⁣(qϕ(z)pprior(z))+const(140)\mathcal{L}_{\text{VAE}}(\phi, \theta) = \dkl{q_\phi(x,z)}{p_\theta(x,z)} +\text{const}\ge \dkl{q_\phi(z)}{\prior(z)} +\text{const} \tag{140}

换句话说,VAE 目标最小化的是 latent 分布与先验之间的 KL 散度的上界。

为什么不就停在 VAE 呢?

根据上面的讨论,VAE 本身就可以作为 Generative Model 来实现,编码器仅仅是为了辅助训练一个互补的解码器,该解码器将 Gaussian distribution 变换为所需的数据分布。然后可以通过采样zppriorz \sim \prior 再采样xpθ(xz)x \sim p_\theta(x|z) 来获得样本。那么,为什么我们坚持要在学到的 latent space 中训练一个单独的 Generative Model 呢?答案与式(139)和式(140)左右两侧之间的所谓摊销差距有关,它恰好对应于信息处理不等式中的差距。当且仅当qϕ(zx)=pθ(zx)q_\phi(z|x)=p_\theta(z|x) 时,这个差距为零,此时编码器表示真实的后验。因此,虽然例如DKL ⁣(qϕ(x,z)pθ(x,z))\dkl{q_\phi(x,z)}{p_\theta(x,z)} 的最小化意味着DKL ⁣(qϕ(z)pprior(z))\dkl{q_\phi(z)}{\prior(z)} 的最小化(见式(140)),但前者的减少并不一定意味着后者的同等减少。因此,在训练结束时,同时存在DKL ⁣(qϕ(x,z)pθ(x,z))\dkl{q_\phi(x,z)}{p_\theta(x,z)} 和摊销差距

DKL ⁣(qϕ(x,z)pθ(x,z))DKL ⁣(qϕ(z)pprior(z))(141)\dkl{q_\phi(x,z)}{p_\theta(x,z)} - \dkl{q_\phi(z)}{\prior(z)} \tag{141}

并未完全最小化,因此qϕ(z)pprior(z)q_\phi(z) \neq \prior(z)。最后,注意在训练期间,解码器学习从qϕ(z)q_\phi(z) 重建,而不是从pprior(z)\prior(z) 重建,因此在推理时切换到从pprior(z)\prior(z) 重建将意味着偏离训练分布。然而在实践中,这种不匹配是一个特性而不是缺陷。实践表明,Flow Model 和 Diffusion Model 通常比用于实现 VAE 解码器的卷积堆栈更具能力,因此将部分生成复杂性转移给 latentGenerative Model 是有意义的。我们将在后面的讨论中回到这一思路。此外,超出这些笔记的范围,Diffusion Model 和 Flow Model 的变分形式将这些模型家族本身实现为 VAE。

证据下界

适当重新排列后,式(132)中的各项可以呈现各种互补的视角。其中之一是所谓的证据下界,我们如下提取。观察到对于固定的xx

Ezqϕ(zx)[log(qϕ(zx)pθ(xz)pprior(z))]=Ezqϕ(zx)[log(qϕ(zx)pθ(zx))]logpθ(x)=DKL ⁣(qϕ(zx)pθ(zx))logpθ(x)(142)\begin{aligned}\EE_{z\sim q_\phi(z|x)} \left[\log \left(\frac{q_\phi(z\mid x)}{p_\theta(x\mid z)\prior(z)}\right)\right] &= \EE_{z\sim q_\phi(z|x)} \left[\log \left(\frac{q_\phi(z\mid x)}{p_\theta(z\mid x)}\right)\right] - \log p_\theta(x)\\ &= \dkl{q_\phi(z\mid x)}{p_\theta(z\mid x)} - \log p_\theta(x)\end{aligned}\tag{142}

其中第一个等式由

pθ(zx)=pθ(xz)pprior(z)pθ(x).p_\theta(z\mid x) = \frac{p_\theta(x\mid z)\prior(z)}{p_\theta(x)}.

因此我们可以重新排列式(142)得到

Ezqϕ(zx)[log(pθ(xz)pprior(z)qϕ(zx))]+DKL ⁣(qϕ(zx)pθ(zx))=logpθ(x),(143)\EE_{z\sim q_\phi(z|x)} \left[\log \left(\frac{p_\theta(x\mid z)\prior(z)}{q_\phi(z\mid x)}\right)\right] + \dkl{q_\phi(z\mid x)}{p_\theta(z\mid x)} = \log p_\theta(x), \tag{143}

由此可得

Ezqϕ(zx)[log(pθ(xz)pprior(z)qϕ(zx))]ELBO(x;ϕ,θ)logpθ(x)evidence.(144)\underbrace{\EE_{z\sim q_\phi(z|x)} \left[\log \left(\frac{p_\theta(x\mid z)\prior(z)}{q_\phi(z\mid x)}\right)\right]}_{\triangleq \,\text{ELBO}(x;\phi, \theta)} \le \underbrace{\log p_\theta(x)}_{\text{evidence}}. \tag{144}

因此,左侧通常被称为证据下界,或 ELBO。我们现在可以通过 ELBO 重写式(136)中的LVAE\mathcal{L}_{\text{VAE}},即

LVAE=DKL ⁣(qϕ(x,z)pθ(x,z))+const=ExpdataEzqϕ(zx)[log(pdata(x)qϕ(zx)pθ(xz)pprior(z))]+const=Expdata[logpdata(x)ELBO(x;ϕ,θ)]+const=Expdata[ELBO(x;ϕ,θ)]H(pdata)+constconst=Expdata[ELBO(x;ϕ,θ)]+const(145)\begin{aligned}\mathcal{L}_{\text{VAE}} &= \dkl{q_\phi(x,z)}{p_\theta(x,z)}+\text{const}\\ &= \mathbb{E}_{x \sim \pdata} \mathbb{E}_{z \sim q_\phi(z|x)} \left[\log \left(\frac{\pdata(x)q_\phi(z\mid x)}{p_\theta(x\mid z)\prior(z)}\right)\right] + \text{const}\\ &= \mathbb{E}_{x \sim \pdata} \left[\log \pdata(x) - \text{ELBO}(x; \phi, \theta)\right] + \text{const}\\ &= -\mathbb{E}_{x \sim \pdata} \left[\text{ELBO}(x; \phi, \theta)\right] \underbrace{- H(\pdata) + \text{const}}_{\text{const}}\\ &= -\mathbb{E}_{x \sim \pdata} \left[\text{ELBO}(x; \phi, \theta)\right] + \text{const}\end{aligned}\tag{145}

因此,原始的 VAE 目标可以看作只是试图最大化期望的 ELBO。最后,让我们考虑在完美训练我们的 VAE 的极限情况下会发生什么。

备注 42(当qϕ(x,z)pθ(x,z)q_\phi(x, z) \approx p_\theta(x, z) 时会发生什么?)

首先,注意到用于训练我们的 latent Generative Model 的采样分布由边缘分布给出。

qϕ(z)=xqϕ(zx)pdata(x)dx.q_{\phi}(z) = \int_x q_\phi(z | x) \pdata(x)\, \d x.

如果 qϕ(x,z)=pθ(x,z)q_\phi(x, z) = p_\theta(x, z),那么特别地,

qϕ(z)=pθ(z)=pprior(z).q_{\phi}(z) = p_\theta(z) = \prior(z).

因此,qϕ(x,z)pθ(x,z)q_\phi(x, z) \approx p_\theta(x, z) 意味着对 latent 采样分布的正则化。其次,qϕ(x,z)pθ(x,z)q_\phi(x, z) \approx p_\theta(x, z) 意味着变分近似 pθ(xz)qϕ(xz)p_\theta(x \mid z) \approx q_\phi(x \mid z) 是好的,进而意味着低重建误差

备注 43(VAE 有什么变化?)

为什么我们不能简单地取 qϕ(x)=pθ(x)q_\phi(\cdot \mid x) = p_\theta(\cdot \mid x),从而保证 qϕ(x,z)=pθ(x,z)=0q_\phi(x, z) = p_\theta(x, z) = 0?原因在于,虽然我们知道似然 pθ(xz)p_\theta(x \mid z),但后验

pθ(zx)=pθ(xz)pprior(z)pθ(x)p_\theta(z \mid x) = \tfrac{p_\theta(x \mid z)\prior(z)}{p_\theta(x)}

通常是难以处理的,因为我们无法获得似然 pθ(x)p_\theta(x)。因此,VAE 中 变分 一词的存在是由于 qϕ(x)q_\phi(\cdot \mid x) 作为难以处理的后验 pθ(x)p_\theta(\cdot \mid x) 的替代品或变分近似

重建与生成

给定一个编码器 qϕ(zx)q_\phi(z | x)、解码器 pθ(xz)p_\theta(x | z) 以及训练用于从 qϕ(z)q_\phi(z) 采样的 latent Generative Model rψr_\psi,我们可以考虑以下两个 Generative Model:

rψ,θrecon(xout)=z,xinpθ(xoutz)qϕ(zxdata)pdata(xdata)dzdxin(reconstruction sampler)rψ,ϕgen(xout)=zgenpθ(xoutzgen)rψ(zgen)dzgen(generative sampler)\begin{aligned}r_{\psi, \theta}^{\text{recon}}(x_{\text{out}}) &= \int_{z, x_{\text{in}}} p_\theta(x_{\text{out}} \mid z)\, q_\phi(z \mid x_{\text{data}})\, \pdata(x_{\text{data}}) \d z \d x_{\text{in}} & (\text{reconstruction sampler}) \\ r_{\psi, \phi}^{\text{gen}}(x_{\text{out}}) &= \int_{z_{\text{gen}}} p_\theta(x_{\text{out}} | z_{\text{gen}})r_\psi(z_{\text{gen}})\, \d z_{\text{gen}} \quad& (\text{generative sampler})\end{aligned}

换句话说,重建 sampler 从 xdatapdatax_{\text{data}} \in \pdata 开始,编码到 zz,然后解码到 xoutx_{\text{out}};而生成 sampler 从 Generative Model 的 zgenrψz_{\text{gen}} \in r_{\psi} 开始,然后通过解码器。通过计算两个 sampler 分布相对于 pdata\pdata 的 Fréchet 初始距离,我们得到重建-FID(rFID)和生成-FID(gFID)。人们也可以考虑通过平均失真(重建的均方根误差)来衡量重建 sampler 的质量,尽管这样的指标对于生成 sampler 没有意义。事实证明,重建 sampler 的质量与生成 sampler 的质量之间存在自然的张力。低 rFID(高质量的重建 sampler)通常表明 latent 中的信息损失低,因此 latent 分布 qϕ(z)q_\phi(z) 在很大程度上类似于 pdata\pdata,并且学习 latent Generative Model 的任务可能更困难,从而提高 gFID。相反,高 rFID 通常表明信息损失高,并且 latent 分布 qϕ(z)q_\phi(z) 更容易学习,从而降低 gFID。这种现象在 Figure 22 中可视化。

分工

重建-生成 sampler 的权衡迫使我们考虑信息损失应如何在 Autoencoder 和 latent Generative Model rψr_\psi 之间分配。直观地说,rψr_{\psi} 通过某个学习到的 Vector Field utψ(zt)u_t^\psi(z_t) 将 standard Gaussian distribution 传输到 qϕ(z)ppriorq_\phi(z) \approx \prior,之后解码器 pθ(xz)p_\theta(x|z)qϕ(z)q_\phi(z) 传输到 pdata\pdata。现在让我们(不精确地)将速率定义为 latent 分布 qϕ(z)q_\phi(z)pprior(z)\prior(z) 匹配的程度,并由此定义生成任务被外包给 latent Generative Model 的程度。〔脚注〕这种分工可以通过绘制速率与失真之间的帕累托前沿来可视化,如 Figure 22 所示。特别是,当速率高时,失真低,反之亦然,这为前面关于重建与生成 sampler 质量的讨论提供了第二个视角。我们以下面的见解来总结我们的讨论。

直觉 44(分工)

图 Figure 22 中的关键洞见是,在帕累托前沿的“拐点”处存在最优的劳动分工,此时我们能够在不产生高失真的情况下获得低码率(高压缩率!)。换句话说,这样的点对应于一个压缩水平,它同时降低了训练底层 Generative Model 的难度,并保持了合理的重建质量。

Diffusion Model 文献导读

文献中有一整类围绕 Diffusion Model 和 Flow Matching 的模型。当你阅读这些论文时,你可能会发现与本节课内容不同的(但等价的)表述方式。这有时会让阅读这些论文变得有些困惑。因此,我们想简要概述各种框架及其差异,并将它们置于历史背景中。这不是理解本文档其余部分所必需的,而是为了在你阅读文献时提供支持。

离散时间与连续时间

Figure 22左图〔原文误作 Right;结合版面应为 Left〕:gFID 与 rFID 之间的权衡,图取自 [51]。这里,ff 表示下采样因子,dd 表示 latent 通道维度。右图:失真(重建质量)与码率的关系,取自 [17, 36]。这条曲线由 DDPM(其本身也是一种 VAE)生成。虽然失真和码率计算中的某些技术细节可能与本文给出的非精确定义有所不同,但整体直觉保持一致。

最初的 Denoising Diffusion Model 论文[41, 42, 17] 并未使用 SDE,而是构建了离散时间下的 Markov chain,即时间步为t=0,1,2,3,t=0,1,2,3,\dots。至今,你会在文献中找到许多使用这种离散时间公式的工作。虽然这种构造因其简单性而具有吸引力,但时间离散方法的缺点在于它迫使你在训练前选择时间离散化方案。此外,损失函数需要通过证据下界(ELBO)来近似——顾名思义,它只是我们实际想要最小化的损失的一个下界。后来,[43] 表明这些构造本质上是时间连续 SDE 的近似。此外,在连续时间情况下,ELBO 损失变得紧(即它不再是下界)(例如,注意第 3.3 节和第 4.3 节是等式而非下界——这在离散时间情况下会有所不同)。这使得 SDE 构造变得流行,因为它被认为在数学上“更干净”,并且可以在训练后通过 ODE/SDEsampler 控制模拟误差。然而,重要的是要注意,这两种模型使用相同的损失,并且并非根本不同。

“前向过程”与 Probability Path

第一波 Denoising Diffusion Model[41, 42, 17, 43] 并未使用Probability Path这一术语,而是通过所谓的前向过程构造了数据点zRdz\in\mathbb{R}^d 的加噪过程。这是一个 SDE,形式为

Xˉ0=z,dXˉt=utforw(Xˉt)dt+σtforwdWˉt(146)\bar{X}_0=z,\quad \dd \bar{X}_t = \uforw_t(\bar{X}_t)\dd t + \sigforw_t \dd \bar{W}_t \tag{146}

其思想是,在抽取数据点zpdataz\sim \pdata 后,模拟前向过程,从而破坏或“加噪”数据。前向过程被设计为使得当tt\to \infty 时,其分布收敛到 Gaussian distributionN(0,Id)\mathcal{N}(0,I_d)。换句话说,对于T0T\gg 0,有XˉTN(0,Id)\bar{X}_{T}\sim\mathcal{N}(0,I_d) 近似成立。注意,这本质上对应于一个 Probability Path:给定Xˉ0=z\bar{X}_0=zXˉt\bar{X}_t 的条件分布是一个 Conditional Probability Pathpˉt(z)\bar{p}_t(\cdot|z),而Xˉt\bar{X}_t 关于zpdataz\sim \pdata 边缘化的分布对应于 Marginal Probability Pathpˉt\bar{p}_t。〔脚注〕然而,请注意,使用这种构造,我们需要知道XtX0=zX_t|X_0=z 的闭式分布,以便训练我们的模型,避免模拟 SDE。这实质上将 Vector Fieldutforw\uforw_t 限制为那些我们知道分布XˉtXˉ0=z\bar{X}_t|\bar{X}_0=z 闭式形式的 Vector Field。因此,在整个 Diffusion Model 文献中,前向过程中的 Vector Field 总是仿射形式,即对于某个连续函数ata_t,有utforw(x)=atx\uforw_t(x)=a_t x。对于这种选择,我们可以使用已知的条件分布公式[40, 43, 23]:

XˉtXˉ0=zN(αtz,βt2I),αt=exp(0tardr),βt2=αt20t(σrforw)2αr2dr\bar{X}_t|\bar{X}_0=z\sim\mathcal{N}\left(\alpha_t z,\beta_t^2 I\right),\quad\alpha_t=\exp\parr{\int\limits_{0}^{t}a_r\dd r},\quad\beta_t^2=\alpha_t^2\int\limits_{0}^{t}\frac{(\sigforw_r)^2}{\alpha^2_r}dr

注意,这些只是 Gaussian Probability Path。因此,可以说前向过程是构造(Gaussian)Probability Path 的一种特定方式。 Probability Path 这一术语由 Flow Matching [25] 引入,旨在同时简化构造并使其更通用:首先,Diffusion Model 的“前向过程”从未被实际模拟(在训练期间仅从pˉt(z)\bar{p}_t(\cdot|z) 中采样)。其次,前向过程仅当tt\to\infty 时才收敛(即我们永远不会在有限时间内到达pinit\pinit)。因此,我们在本文档中选择使用 Probability Path。

时间反转与求解 Fokker-Planck 方程

Diffusion Model 的原始描述并没有通过 Fokker-Planck 方程(或连续性方程)来构造训练目标uttarget\uref_tlogpt\nabla\log p_t,而是通过前向过程的时间反转[2]。时间反转(Xt)0tT(X_t)_{0\leq t\leq T} 是一个 SDE,其轨迹上的分布在时间上反转,即

P[Xˉt1A1,,XˉtnAn]=P[XTt1A1,,XTtnAn](147)\mathbb{P}[\bar{X}_{t_1}\in A_1,\dots,\bar{X}_{t_n}\in A_n]=\mathbb{P}[X_{T-t_1}\in A_1,\dots,X_{T-t_n}\in A_n] \tag{147}
 for all 0t1,,tnT, and A1,,AnS(148)\text{ for all }0\leq t_1,\dots, t_n\leq T, \text{ and } A_1,\dots,A_n\subset S \tag{148}

如[2] 所示,可以通过 SDE 获得满足上述条件的时间反转:

dXt=[ut(Xt)+σt2logpt(Xt)]dt+σtdWt,ut(x)=uTtforw(x),σt=σˉTt\begin{aligned}\dd X_t =& \left[-u_t(X_t)+\sigma_t^2\nabla\log p_t(X_t)\right]\dd t+ \sigma_{t}\dd W_t,\quad u_t(x)=\uforw_{T-t}(x),\sigma_t=\bar{\sigma}_{T-t}\end{aligned}

作为ut(Xt)=atXtu_t(X_t)=a_tX_t,上述对应于我们在第 4.1 节中推导的训练目标的一个特定实例(这并非显而易见,因为使用了不同的时间约定。参见例如[26] 的推导)。然而,对于生成建模的目的,我们通常只使用 Markov process 的最终点X1X_1(例如,作为生成的图像),并丢弃较早的时间点。因此,Markov process 是“真正”的时间反转还是沿着 Probability Path 进行,对于许多应用来说并不重要。因此,使用时间反转并非必要,而且常常导致次优结果,例如概率流 ODE 通常更好[23, 28]。所有不同于时间反转的 Diffusion Model 采样方式都再次依赖于使用 Fokker-Planck 方程。我们希望这能说明为什么如今许多人直接通过 Fokker-Planck 方程构造训练目标——正如[25, 27, 1] 所开创并在本课程中所做的那样。

Flow Matching [25] 和 Stochastic Interpolants [1]

我们提出的框架与 Flow Matching 和随机插值(SIs)的框架最为密切相关。正如我们所了解的,Flow Matching 将自身限制于流。事实上,Flow Matching 的关键创新之一在于表明,不需要通过前向过程和 SDE 的构造,仅 Flow Model 就可以以可扩展的方式进行训练。由于这一限制,您应记住,从 Flow Matching 模型采样将是确定性的(只有初始X0pinitX_0\sim \pinit 是随机的)。随机插值既包括纯流,也包括我们在此使用的通过“Langevin dynamics”的 SDE 扩展(参见第 4.2 节)。随机插值得名于一个插值函数I(t,x,z)I(t,x,z),旨在在两个分布之间进行插值。在我们这里使用的术语中,这对应于构造 Conditional and Marginal Probability Paths 的一种不同但(主要)等价的方式。Flow Matching 和随机插值相对于 Diffusion Model 的优势在于其简单性和通用性:它们的训练框架非常简单,但同时允许您从任意分布pinit\pinit 到任意分布pdata\pdata——而 Denoising Diffusion Model 仅适用于 Gaussian 初始分布和 Gaussian Probability Path。这为生成建模开辟了新的可能性,我们将在本课程后面简要提及。

小结 45(另类扩散配方)

文献中流行的 Diffusion Model 替代公式通常涉及以下元素的某种组合:

  1. 离散时间:通常使用通过离散时间 Markov chain 对 SDE 的近似。
  2. 反转时间约定:通常使用反转时间约定,其中t=0t=0 对应于pdata\pdata(而在这里t=0t=0 对应于pinit\pinit)。
  3. 前向过程:前向过程(或加噪过程)是构造(Gaussian)Probability Path 的方法。
  4. 通过时间反转的训练目标:训练目标也可以通过 SDE 的时间反转来构造。这是此处呈现的构造的一个特定实例(使用反转时间约定)。