Table of Contents
Fetching ...

Multi-Marginal Schrödinger Bridge Matching

Byoungwoo Park, Juho Lee

TL;DR

This work tackles trajectory inference from multi-snapshot data by extending Schrödinger Bridge frameworks to the multi-marginal setting for $|\mathcal{T}|>2$ marginals. It proposes Multi-Marginal Schrödinger Bridge Matching (MSBM), which partitions the time horizon into local intervals and learns a shared drift/control that glues the local solutions into a globally consistent trajectory. The work develops multi-marginal reciprocal and Markov projection theory, proves convergence and uniqueness properties, and provides a practical training objective that scales via parallelism. Empirical results on synthetic data and real single-cell datasets show competitive or superior trajectory fidelity and notably faster training than competing methods, underscoring MSBM's applicability to developmental biology and systems medicine.

Abstract

Understanding the continuous evolution of populations from discrete temporal snapshots is a critical research challenge, particularly in fields like developmental biology and systems medicine where longitudinal tracking of individual entities is often impossible. Such trajectory inference is vital for unraveling the mechanisms of dynamic processes. While Schrödinger Bridge (SB) offer a potent framework, their traditional application to pairwise time points can be insufficient for systems defined by multiple intermediate snapshots. This paper introduces Multi-Marginal Schrödinger Bridge Matching (MSBM), a novel algorithm specifically designed for the multi-marginal SB problem. MSBM extends iterative Markovian fitting (IMF) to effectively handle multiple marginal constraints. This technique ensures robust enforcement of all intermediate marginals while preserving the continuity of the learned global dynamics across the entire trajectory. Empirical validations on synthetic data and real-world single-cell RNA sequencing datasets demonstrate the competitive or superior performance of MSBM in capturing complex trajectories and respecting intermediate distributions, all with notable computational efficiency.

Multi-Marginal Schrödinger Bridge Matching

TL;DR

This work tackles trajectory inference from multi-snapshot data by extending Schrödinger Bridge frameworks to the multi-marginal setting for marginals. It proposes Multi-Marginal Schrödinger Bridge Matching (MSBM), which partitions the time horizon into local intervals and learns a shared drift/control that glues the local solutions into a globally consistent trajectory. The work develops multi-marginal reciprocal and Markov projection theory, proves convergence and uniqueness properties, and provides a practical training objective that scales via parallelism. Empirical results on synthetic data and real single-cell datasets show competitive or superior trajectory fidelity and notably faster training than competing methods, underscoring MSBM's applicability to developmental biology and systems medicine.

Abstract

Understanding the continuous evolution of populations from discrete temporal snapshots is a critical research challenge, particularly in fields like developmental biology and systems medicine where longitudinal tracking of individual entities is often impossible. Such trajectory inference is vital for unraveling the mechanisms of dynamic processes. While Schrödinger Bridge (SB) offer a potent framework, their traditional application to pairwise time points can be insufficient for systems defined by multiple intermediate snapshots. This paper introduces Multi-Marginal Schrödinger Bridge Matching (MSBM), a novel algorithm specifically designed for the multi-marginal SB problem. MSBM extends iterative Markovian fitting (IMF) to effectively handle multiple marginal constraints. This technique ensures robust enforcement of all intermediate marginals while preserving the continuity of the learned global dynamics across the entire trajectory. Empirical validations on synthetic data and real-world single-cell RNA sequencing datasets demonstrate the competitive or superior performance of MSBM in capturing complex trajectories and respecting intermediate distributions, all with notable computational efficiency.
Paper Structure (31 sections, 12 theorems, 69 equations, 5 figures, 8 tables, 2 algorithms)

This paper contains 31 sections, 12 theorems, 69 equations, 5 figures, 8 tables, 2 algorithms.

Key Result

Proposition 0

For any $\mathbf{x}_{{\mathcal{T}}} := (\mathbf{x}_0, \mathbf{x}_{t_1}, \cdots, \mathbf{x}_{T}) \in \mathbb{R}^{d \times (k+1)}$ and $t \in [t_{i-1}, t_i)$, the marginal distribution of $\mathbb{Q}_{|{\mathcal{T}}}(\cdot | \mathbf{x}_{{\mathcal{T}}})$ at $t$ satisfies: Therefore, for any $\mathbb{P} \in {\mathcal{P}}_{[0, T]}$ the reciprocal projection ${\mathcal{R}}^{\texttt{mm}}(\mathbb{P}, {\m

Figures (5)

  • Figure 1: Training of MSBM
  • Figure 2: Comparison of generated population dynamics using MIOFlow, DMSB and MSBM on a 2-dim petal dataset. All trajectories are generated by simulating the dynamics from $\rho_{t_0}$.
  • Figure 3: Evaluation results of ${\mathcal{W}}_2$ and MMD.
  • Figure 3: Performance on the 100-dim PCA of EB dataset. MMD and SWD are computed between test $\rho^{\texttt{te}}_{t_i}$ and generated $\hat{\rho}_{t_i}$ by simulating the dynamics from test $\rho^{\texttt{te}}_{t_0}$. Best results are highlighted.
  • Figure 4: Comparison of generated population dynamics using DMSB and MSBM on a 100-dim PCA of EB dataset. The plot displays the first two principal components as the x and y axes, respectively.

Theorems & Definitions (19)

  • Proposition 0: Reciprocal Property
  • Proposition 0: Multi-Marginal Markovian Projection
  • Proposition 0: Uniqueness
  • Proposition 0: Convergence
  • Corollary 0: Multi-Marginal Schrödinger Bridge
  • Theorem A.1: Dynamics SB optimality pavon1991free
  • Theorem A.2: Multi-Marginal Schrödinger Potentials
  • proof
  • Remark A.3
  • Proposition 0: Reciprocal Property
  • ...and 9 more