Schrödinger bridge based deep conditional generative learning¶
讲者: Jun Liu (Tsinghua University)
会场: Plenary Talk
报告题目: Conditional Generation via Diffusion, Flow, and Schrödinger Bridges
链接: arXiv
来源: JCSDS 2026 · 返回会议总览
一、领域脉络与小综述¶
这个方向是什么¶
本文研究的子方向是基于扩散过程的条件生成模型,其根本的统计问题是:给定观测到的协变量 \( z \in \mathbb{R}^{d_z} \),如何从条件分布 \( p(x|z) \) 中高效地生成样本 \( x \in \mathbb{R}^{d_x} \),尤其是在 \( d_x \) 和 \( d_z \) 都很大、传统密度估计方法失效的场景下。当前该方向的成熟度较高,已有多种基于 GAN、VAE、归一化流和扩散模型的方法,但核心瓶颈在于生成速度(需要长时扩散)和训练稳定性(GAN 的模式坍塌、迭代 SB 方法的计算负担)。
发展脉络(history)¶
- 奠基工作:生成对抗网络与变分自编码器
- Goodfellow et al. (2014):提出 GAN,通过生成器与判别器的对抗训练学习数据分布。但条件 GAN(如 Zhou et al. (2023) 的 GCDS)面临训练不稳定、模式坍塌等问题(作者引用 Karras et al. (2019) 指出其需要大量人工调参)。
-
Kingma (2013) / Sohn et al. (2015):条件 VAE(cVAE)通过变分下界学习条件分布,计算成本低但样本质量较差(作者原话:"tend to produce samples of lower visual quality")。
-
主要进展:扩散模型与得分匹配
- Sohl-Dickstein et al. (2015) / Song & Ermon (2020) / Ho et al. (2020) / Song et al. (2021):建立得分匹配与扩散概率模型(DDPM)框架,通过正向加噪、反向去噪生成样本。条件扩散模型(如 Dhariwal & Nichol (2021) 的 classifier guidance、Ho & Salimans (2022) 的 classifier-free guidance)在图像生成上达到 SOTA。
-
瓶颈:正向扩散需要运行足够长时间才能使分布接近高斯,导致生成速度慢(作者原话:"computationally expensive as one needs to run the forward noising diffusion long enough to converge to the reference distribution")。
-
当前 frontier:Schrödinger 桥(SB)加速
- Léonard (2013) / Chen et al. (2021) / Bortoli et al. (2021):将 SB 问题引入生成建模,通过熵正则化最优传输在有限时间内完成分布变换。
- Bortoli et al. (2021) 的 DSB:提出迭代比例拟合(IPF)算法近似 SB 解,但需要交替训练前向/后向网络,计算量大。
- Shi et al. (2022) 的 CDSB:将 DSB 扩展到条件生成,但每轮迭代需存储两个神经网络(作者原话:"computationally demanding because at each iteration a deep neural network needs to be trained which depends on the results obtained in the previous steps")。
-
Huang (2024):针对无条件生成,利用 Dirac Delta 初始分布得到 SB 闭式解,实现一步模拟、无需迭代。本文将其扩展到条件情形。
-
本文的位置:在 Huang (2024) 的无条件 SB 方法基础上,提出条件 Schrödinger 桥(SBCG),核心优势是一步模拟、无需迭代,且训练目标仅含一个未知函数 \( u_\theta(x,z,t) \),无需判别器或交替训练。
子线索聚类¶
-
线索 1:基于 GAN 的条件生成(Goodfellow et al., 2014; Zhou et al., 2023; Karras et al., 2019)
通过条件生成器与判别器对抗训练,样本质量高但训练不稳定、易模式坍塌。 -
线索 2:基于扩散/得分匹配的条件生成(Song et al., 2021; Dhariwal & Nichol, 2021; Ho & Salimans, 2022; Batzolis et al., 2021; Tashiro et al., 2021)
通过学习得分函数反向扩散,条件生成质量 SOTA,但需要长时扩散,生成速度慢。 -
线索 3:基于 Schrödinger 桥的条件生成(Bortoli et al., 2021; Shi et al., 2022; Liu et al., 2023; Peluchetti, 2024)
通过熵正则化最优传输在有限时间内完成分布变换,但现有方法(如 CDSB)需要迭代训练,计算成本高。 -
线索 4:其他条件生成方法(cVAE: Sohn et al., 2015; CNF: Winkler et al., 2023; CFF: Chang et al., 2024)
cVAE 和 CNF 计算成本低但样本质量有限;CFF 与本文类似,但初始分布为高斯而非固定点。
这个方向在追问的核心问题¶
- 如何在不牺牲样本质量的前提下加速条件生成?
- 主流方法:SB 加速(有限时间)、知识蒸馏(Luhman & Luhman, 2021)、非马尔可夫前向过程(Song et al., 2022)。
-
瓶颈:现有 SB 方法(CDSB)仍需迭代训练,计算负担大。
-
如何避免迭代训练,实现一步模拟的条件生成?
-
本文的切入点:利用 Dirac Delta 初始分布得到 SB 闭式解,将条件生成转化为一个回归问题(学习漂移项 \( u_\theta(x,z,t) \))。
-
如何在高维条件(如图像)下保持生成质量?
- 主流方法:神经网络参数化得分/漂移函数。
- 瓶颈:训练数据需求大,且条件变量 \( z \) 的维度可能很高(如图像修复中的部分像素)。
⚠️ 作者的 framing¶
- 作者把缺口 frame 成什么:作者认为现有条件 SB 方法(CDSB)的迭代训练是主要瓶颈,因此将自己的方法定位为"one-step simulation free algorithm which does not need to run iteration"(Section 5)。通过将初始分布设为 Dirac Delta,SB 问题有闭式解,从而将条件生成转化为一个简单的回归问题(最小化 (3.4) 或 (3.5))。
- 哪些竞争路线被他淡化或回避了:
- 作者提到 CFF(Chang et al., 2024)"shared similar spirit",但仅指出初始分布不同(高斯 vs 固定点),未深入比较两者在理论保证或实际性能上的差异。
- 作者未讨论 classifier-free guidance(Ho & Salimans, 2022)在 SB 框架下的对应物,仅在结论中提及"could be designed"作为未来工作。
- 作者未提及 SB 方法在条件生成中的理论保证(如收敛速度、样本质量界),而仅依赖数值实验。
- 什么明显该被引/该存在、却没出现在 intro 里:
- 没有引用关于条件 SB 问题的理论识别性的工作(如 SB 解是否唯一、条件信息如何影响最优传输路径)。
- 没有引用高维条件生成中的维数灾难分析(如条件得分函数的 Lipschitz 常数如何随 \( d_z \) 增长)。
- 没有引用SB 与最优传输的统计效率比较(如熵正则化参数的选择对样本质量的影响)。
张力¶
未见明显对立引用。所有被引工作基本沿着"GAN → 扩散模型 → SB 加速"的演进方向,彼此互补而非矛盾。
二、最核心、最简单的例子 / 数学问题¶
第一步:符号、模型、可观测数据交代清楚¶
- 符号:
- \( x \in \mathbb{R}^{d_x} \):要生成的响应变量(随机向量)。
- \( z \in \mathbb{R}^{d_z} \):条件变量(协变量),可以是离散标签或连续向量。
- \( (x,z) \sim \mu_{x,z} \):联合分布,可观测到 i.i.d. 样本 \( \{(x_i,z_i)\}_{i=1}^n \)。
- \( \mu_{x|z} \):目标条件分布,即要从中采样的分布(潜在、不可直接观测)。
- \( a \in \mathbb{R}^{d_x} \):初始固定点(如 \( a = 0 \)),SB 过程的起点。
- \( t \in [0,1] \):时间索引。
- \( x_t \):在时间 \( t \) 的随机过程值。
- \( u^\star(x,z,t) \):最优漂移项(SB 解中的额外漂移),是本文要估计的对象。
- \( u_\theta(x,z,t) \):用神经网络参数化的漂移项估计。
- \( \alpha(t), \beta(t) \):参考 SDE 中的时间函数,控制扩散速度。
-
\( \xi, \tau, \sigma^2_1, \sigma^2_2, \sigma^2 \):参考 SDE (2.5) 的辅助系数,见 (2.11)。
-
模型:
- 数据生成机制:假设存在一个联合分布 \( \mu_{x,z} \),研究者可从中获得 i.i.d. 样本。目标是从条件分布 \( \mu_{x|z} \) 中生成新样本。
- 统计模型:通过一个 Schrödinger 桥问题来建模从固定点 \( a \) 到目标条件分布 \( \mu_{x|z} \) 的随机过程。该过程由 SDE (3.3) 控制,其中漂移项 \( u^\star(x,z,t) \) 是待估函数。
-
已知量:参考 SDE 的漂移 \( b(x,t) \) 和扩散 \( \sigma(t) \) 是人为选择的(如 (2.4) 或 (2.5)),其过渡核 \( q(s,x_s,t,x_t) \) 有解析形式(高斯分布)。
-
可观测数据:
- 可观测:联合样本 \( \{(x_i,z_i)\}_{i=1}^n \)。
- 潜在/不可观测:条件分布 \( \mu_{x|z} \) 本身,以及 SB 过程中的漂移项 \( u^\star(x,z,t) \)。后者通过最小化 (3.4) 或 (3.5) 来间接学习,无需显式估计密度。
第二步:最小内核¶
最简特例:考虑一维情形(\( d_x = d_z = 1 \)),参考 SDE 取为标准布朗运动((2.4) 中 \( \alpha(t) = 1 \)),初始点 \( a = 0 \)。此时: - 参考 SDE:\( dx_t = dw_t \),过渡核 \( q(s,x_s,t,x_t) = \mathcal{N}(x_t; x_s, t-s) \)。 - 目标:给定 \( z \),从 \( \mu_{x|z} \) 中采样。
核心思路:SB 问题的解 \( u^\star(x,z,t) \) 是以下二次损失的最小化器:
为什么这个损失有效:可以证明,最优解 \( u^\star(x,z,t) = \frac{1}{1-t} \mathbb{E}[x_1 - x_t | x_t = x, z] \),即给定当前状态 \( x_t \) 和条件 \( z \) 下,从 \( x_t \) 到终点 \( x_1 \) 的期望位移除以剩余时间。因此,学习 \( u^\star \) 等价于学习条件期望 \( \mathbb{E}[x_1 | x_t = x, z] \),而后者完全决定了条件分布 \( \mu_{x|z} \)(因为 SB 过程是马尔可夫的)。
最简例子:假设 \( x = \tanh(z) + \epsilon \),\( \epsilon \sim \Gamma(1,0.3) \)(论文 Example 1)。此时: - 可观测数据:\( \{(x_i,z_i)\} \) 来自该模型。 - 训练:对每个 \( (x_i,z_i) \),采样 \( t_j \sim U[\epsilon,1-\epsilon] \),计算 \( x_{t_j}^{(i)} \sim \mathcal{N}(x_i, t_j(1-t_j)) \),然后最小化:
核心数学困难:如何保证 \( u_\theta \) 的估计误差在生成过程中不累积?论文未提供理论分析,仅依赖数值实验验证。
三、这篇论文做了什么¶
三句话¶
- 研究了什么问题:如何利用 Schrödinger 桥(SB)在有限时间内高效地从条件分布 \( p(x|z) \) 中生成样本,避免现有条件 SB 方法(如 CDSB)的迭代训练负担。
- 核心工具/方法:将初始分布设为 Dirac Delta,利用 SB 闭式解将条件生成转化为一个回归问题——用神经网络 \( u_\theta(x,z,t) \) 最小化一个二次损失((3.4) 或 (3.5)),然后通过欧拉-丸山法离散化 SDE 生成样本。
- 主要结论:数值实验表明,SBCG 在低维(2D 示例、条件均值/标准差估计)和高维(MNIST 图像生成、图像修复)任务中,生成的样本质量优于或持平于 GCDS、CKDE、NNKCDE、FlexCode 等方法,且训练更简单(无需迭代、无需判别器)。
关键设定与假设¶
- 设定:
- 可观测 i.i.d. 联合样本 \( \{(x_i,z_i)\}_{i=1}^n \)。
- 初始分布为 Dirac Delta \( \delta_a \)(固定点 \( a \),如 \( a=0 \))。
- 参考 SDE 为 (2.4)(布朗运动型)或 (2.5)(OU 过程型),其过渡核有解析高斯形式。
-
漂移项 \( u^\star(x,z,t) \) 用神经网络 \( u_\theta(x,z,t) \) 参数化。
-
假设:
- S1:参考 SDE 的过渡核 \( q(s,x_s,t,x_t) \) 是高斯分布(由 (2.4) 和 (2.5) 保证)。
- S2:\( \alpha(t) \) 和 \( \beta(t) \) 是 \( t \) 的非负函数,确保方差为正。
- S3:时间区间截断 \( [\epsilon, 1-\epsilon] \) 以避免 \( t=0,1 \) 处的数值不稳定(常见实践,如 Song & Ermon, 2020)。
-
S4:神经网络 \( u_\theta \) 有足够表达能力(未给出具体条件,如 Lipschitz 或 Sobolev 正则性)。
-
相比已有文献的放宽/强化:
- 放宽:相比 CDSB(Shi et al., 2022),无需迭代训练,只需训练一个网络。
- 强化:相比 CFF(Chang et al., 2024),初始分布为固定点而非高斯,且参考 SDE 更一般(CFF 是本文的特例,作者声称)。
- 未放宽:仍需要大量联合样本(实验中用 50,000 个训练点),且未提供理论保证(如收敛速度、样本质量界)。
主要结果¶
- Proposition 1(无条件 SB 的漂移项形式):对于参考 SDE (2.4) 和 (2.5),最优漂移 \( u^\star(x,t) \) 分别是 (2.9) 和 (2.10) 的唯一最小化器。证明通过条件期望性质完成(见附录 A)。
- Proposition 2(条件 SB 的漂移项形式):将 Proposition 1 推广到条件情形,\( u^\star(x,z,t) \) 是 (3.4) 和 (3.5) 的唯一最小化器。证明思路相同,仅将无条件期望替换为条件期望(见附录 B)。
- 数值实验:
- 低维(2D 示例):SBCG 生成的样本直方图与真实密度吻合(图 1、图 2)。
- 条件均值/标准差估计:在三个非线性模型(Example 4-6)上,SBCG 的 MSE 在 6 个比较项中 5 项最小(表 1),优于 GCDS、CKDE、NNKCDE、FlexCode。
- UCI 数据集(葡萄酒质量、鲍鱼):SBCG 构建的预测区间覆盖率接近名义水平(表 2,如 \( \alpha=0.1 \) 时覆盖率为 0.906-0.907)。
- MNIST 图像生成:给定标签生成图像(图 3)和图像修复(图 4),SBCG 能正确重建大部分数字,尤其在给定 3/4 图像时全部正确。
证明路线与技术技巧¶
- 整体路线(以 Proposition 1 为例):
- 将 SB 解 \( u^\star(x,t) \) 表达为条件期望形式:\( u^\star(x,t) = \frac{\alpha'(t)}{\alpha(1)-\alpha(t)} \mathbb{E}[x_1 - x_t | x_t = x] \)(由 (A.3) 得到)。
- 对任意 \( u(x,t) \),将损失 (2.9) 分解为三项:\( \| \text{目标} - u^\star \|^2 + \| u^\star - u \|^2 + 2 \langle \text{目标} - u^\star, u^\star - u \rangle \)。
- 利用条件期望性质证明交叉项为零:\( \mathbb{E}[\langle \text{目标} - u^\star, u^\star - u \rangle] = 0 \)。
-
因此损失在 \( u = u^\star \) 时最小,且唯一。
-
关键跳跃点:从 (A.1) 到 (A.3) 的推导——将积分形式转化为条件期望形式,依赖于高斯过渡核的解析形式。这一步是后续所有证明的基础。
-
技术技巧点名:
- 条件期望分解:将漂移项表达为条件期望,从而将优化问题转化为回归问题。
- 高斯过渡核的解析形式:利用 Särkkä & Solin (2019) 的公式计算 \( q(s,x_s,t,x_t) \),得到均值和方差的闭式表达式。
- 交叉项消去:通过条件期望性质证明交叉项为零,这是证明 \( u^\star \) 为唯一最小化器的关键。
真实例子与应用¶
- 模拟数据:
- 2D 示例(Section 4.1.1):三个非线性、非高斯模型(Example 1-3),以及四个复杂形状(checkerboard, moons, pinwheel, swissroll)。用 50,000 个训练点,4 层全连接网络(32-64-64-32),\( K=100 \) 步离散化。结果:直方图与真实密度吻合。
-
条件均值/标准差估计(Section 4.1.2):三个模型(Example 4-6),其中 Example 5 有异方差误差,Example 6 是混合高斯。用 50,000 个训练点,2000 个测试 \( z \),每个 \( z \) 生成 200 个 \( x \)。结果:SBCG 在 5/6 的 MSE 指标上最优。
-
UCI 数据集(Section 4.2):
- 葡萄酒质量:6497 个样本,11 个化学测量预测质量分数(0-10)。
- 鲍鱼:4177 个样本,物理测量预测环数(年龄)。
-
方法:90% 训练,10% 测试;对每个测试 \( z_i \) 生成 200 个 \( x \);构建预测区间并计算覆盖率。结果:覆盖率接近名义水平(如 90% 区间实际覆盖 90.6%-90.7%)。
-
MNIST 图像(Section 4.3):
- 标签生成(Section 4.3.1):60,000 张 28×28 图像,10 类标签。对每类单独训练(用子集数据),生成图像(图 3)。结果:生成图像与真实图像难以区分。
- 图像修复(Section 4.3.2):给定 1/4、1/2、3/4 的图像,重建缺失部分。用 10,000 个训练点,3 层 64 宽网络(SeLU 激活)。结果:给定 3/4 时全部正确重建;数字 0、2、4、6、8、9 在 1/4 时即可正确重建。
🔎 结论是否比证明窄¶
- 窄结论:Proposition 2 的证明仅针对两种参考 SDE((2.4) 和 (2.5)),但作者在 Section 2.1 声称"we consider two reference SDEs",未声称适用于更一般的 SDE。然而,在 Section 5 的讨论中,作者暗示框架可推广("our framework has more flexibility in selecting the reference SDE"),但未提供证明。
- 泛化 claim:作者在 Section 3.2.1 说"our framework does not need to rely on neural networks; there can be other ways to represent the function \( u_\theta \)",但全文仅用神经网络实现,未展示其他表示(如核方法)的有效性。
- 理论保证缺失:论文未提供任何关于估计误差 \( \| u_\theta - u^\star \| \) 的收敛速度、生成样本分布与真实条件分布之间的 KL 散度界、或离散化误差的定量分析。作者在 Section 5 仅说"validity and accuracy of our method are accessed using numerical experiments",未声称理论结果。
四、开放问题¶
- 理论保证:能否给出 \( u_\theta \) 的估计误差界(如 \( \| u_\theta - u^\star \|_{L^2(\mu)} \) 随样本量 \( n \) 和网络复杂度的收敛速度)?以及离散化误差(欧拉-丸山法的步长 \( h \) 对生成样本分布的影响)?
-
扎根点:Section 3.2.2 仅引用 Chen et al. (2023) 说"deviation can be controlled by the step size",但未给出具体界。
-
与 CFF 的详细比较:本文声称 CFF 是特例,但未在理论或实验上系统比较。CFF 的初始分布为高斯,本文为固定点——这对生成质量有何影响?
-
扎根点:Section 1 仅说"difference is that the diffusion process in our framework starts from a fixed point instead of a standard Gaussian",未深入分析。
-
更一般的参考 SDE:Proposition 2 的证明依赖于高斯过渡核。对于非高斯过渡核(如带跳的扩散),SB 解是否仍有类似形式?
-
扎根点:Section 2.1 仅考虑 (2.4) 和 (2.5),但 Section 5 暗示"more general class"。
-
条件 SB 的统计-计算权衡:本文方法避免了迭代训练,但代价是什么?是否在样本质量或多样性上有所牺牲?能否用低度多项式(low-degree)或 SQ 下界分析其计算复杂度?
- 扎根点:Section 5 提到"trade-off between sample quality and diversity"作为未来工作,但未量化。
Maintained by 陈星宇 · Homepage · Source on GitHub