Source-linked AI summary

Aligned Diffusion Schrödinger Bridges

Vignesh Ram Somnath, Matteo Pariset, Ya-Ping Hsieh, Maria Rodriguez Martinez, Andreas Krause, Charlotte Bunne

arXiv:2302.11419v3cs.LGq-bio.QM

TL;DR

Existing diffusion Schrödinger bridge algorithms do not exploit paired observations, despite alignment being common in biological data. The paper introduces SBALIGN, combining Schrödinger bridge theory and Doob’s h-transform to solve aligned DSBs without IPF-like training. It reports sizeable improvements across synthetic and real experiments, including cellular differentiation and protein conformational-change tasks.

  • Problem

    Existing DSB frameworks fail to incorporate naturally paired initial and final observations, although such alignment occurs in biological phenomena.

  • Method

    SBALIGN combines classical Schrödinger bridge theory with Doob’s h-transform to derive an aligned-data loss that avoids IPF-like training and uses regularization.

  • Results

    The method improves over previous state-of-the-art methods across synthetic and real-world tasks, including cellular processes and protein conformational changes.

  • Takeaways & Limitations

    The experiments support treating data alignment as a relevant feature for diffusion Schrödinger bridge modeling.

  • Takeaways & Limitations

    SBALIGN may be limited when available information about pairings is insufficient, and the protein-docking application remains a proof of concept requiring integration with rigid-protein docking methods.

Abstract

from arXiv · show

Diffusion Schrödinger bridges (DSB) have recently emerged as a powerful framework for recovering stochastic dynamics via their marginal observations at different time points. Despite numerous successful applications, existing algorithms for solving DSBs have so far failed to utilize the structure of aligned data, which naturally arises in many biological phenomena. In this paper, we propose a novel algorithmic framework that, for the first time, solves DSBs while respecting the data alignment. Our approach hinges on a combination of two decades-old ideas: The classical Schrödinger bridge theory and Doob's $h$-transform. Compared to prior methods, our approach leads to a simpler training procedure with lower variance, which we further augment with principled regularization schemes. This ultimately leads to sizeable improvements across experiments on synthetic and real data, including the tasks of predicting conformational changes in proteins and temporal evolution of cellular differentiation processes.

1 INTRODUCTION

The paper addresses diffusion Schrödinger bridges’ failure to use naturally paired observations by introducing SBALIGN, a framework for aligned stochastic interpolation. It combines Schrödinger bridge theory with Doob’s h-transform to avoid IPF-like training, reduce variance, and improve results on synthetic and biological tasks.

  • Motivation: Aligned observations naturally pair initial and final states in biological processes, but existing DSB methods treat their dependence as unknown.This omission discards correspondence information and makes trajectory recovery harder than necessary.
  • Method: SBALIGN combines Schrödinger bridge theory with Doob’s h-transform to recover trajectories between paired states without iterative proportional fitting.Its learned drift represents an SB solution together with a pairing-related h-transform term.
  • Contributions: The framework formulates aligned interpolation in the DSB setting and introduces a reference-process construction that can extend to hybrid aligned/non-aligned bridges.The paper presents this as the first formulation of interpolation with aligned data in the DSB framework.
  • Method: The new loss avoids IPF-like procedures, while principled regularization schemes are designed to stabilize training and lower variance.The method is motivated by numerical instability in prior IPF-based approaches.
  • Experiments: The framework is evaluated on synthetic data, cellular developmental processes, and protein conformational changes during docking.The protein application models transitions between unbound and bound states, while aligned biological data also arises in cell differentiation studies.
  • Results: The method shows considerable improvement over prior methods across multiple metrics, supporting the importance of incorporating alignment.Related work compares against DSB methods designed for unaligned data and against an aligned-data framework with different pathwise objectives.

2 BACKGROUND

The paper formulates interpolation between paired endpoint observations as an aligned stochastic-process recovery problem, motivated by biological applications. Existing DSB objectives use only endpoint marginals, discarding pairing information and requiring IPF-based training.

  • 2 BACKGROUND: Aligned data pairs endpoint observations (x_i0, x_i1), enabling recovery of a stochastic trajectory between corresponding initial and final states.The paper highlights protein conformational changes, molecular dynamics, and other biological settings where correspondence between endpoints is naturally available.
  • 2 BACKGROUND: Protein docking illustrates the task: the initial and final molecular structures are paired, so interpolation should respect each molecule’s correspondence.Ignoring alignment would discard information about which initial structure maps to which final structure.
  • 2 BACKGROUND: Classical DSBs reconstruct processes from endpoint marginals, but their objective loses the joint pairing information between corresponding observations.The paper notes that the standard formulation uses only P̂0 and P̂1, not the paired samples (x_i0, x_i1).
  • 2 BACKGROUND: Because marginal interpolation requires forward-backward iterative proportional fitting, existing DSB training faces high variance and numerical or scalability issues.The proposed direction is to solve aligned interpolation without an IPF procedure.

3 ALIGNED DIFFUSION SCHRÖDINGER BRIDGES

The paper combines Schrödinger bridge theory with Doob’s h-transform to learn diffusion drifts from aligned endpoint pairs without IPF. It further uses learned aligned dynamics as data-informed references for classical DSBs, with regularization aimed at training stability and alignment fidelity.

  • 3 ALIGNED DIFFUSION SCHRÖDINGER BRIDGES: The framework derives an aligned DSB loss from Schrödinger bridge theory and Doob’s h-transform, avoiding IPF-like training procedures.The method represents the drift through an SB component and a pairing-related ∇log h term.
  • 3.1 LEARNING ALIGNED DIFFUSION SCHRÖDINGER BRIDGES: Aligned SB paths are constructed by sampling paired endpoints from the optimal coupling and connecting them with scaled Brownian bridges.The resulting process is a mixture of reference bridges weighted by the endpoint coupling, while its drift is learned through an SDE characterization.
  • 3.1 LEARNING ALIGNED DIFFUSION SCHRÖDINGER BRIDGES: Doob’s h-transform adds ∇log h_t to the drift, where h_t represents the conditional probability of the terminal endpoint given the current state.This conditional SDE supplies the bridge-compatible drift representation used for learning.
  • 3.1 LEARNING ALIGNED DIFFUSION SCHRÖDINGER BRIDGES: A second neural network and alternating minimization reduce training variance by avoiding unconditional-path sampling, while ℓ2 regularization stabilizes h-transform behavior near t → 1.The regularization also encourages drifts whose h-transforms diminish and whose trajectories respect endpoint alignment in expectation.
  • 3.2 PAIRED SCHRÖDINGER BRIDGES AS PRIOR PROCESSES: SBALIGN can use its learned aligned drift to define a data-informed reference process for a subsequent classical SB problem.This is intended to improve the reference coupling relative to a standard Brownian motion, while limited pairings constrain the accuracy of the aligned solution.
  • 3.2 PAIRED SCHRÖDINGER BRIDGES AS PRIOR PROCESSES: On synthetic Moon and T datasets, SBALIGN learns drifts that respect the true alignment, and its learned drift can improve another training method when used as a reference process.The Moon example contrasts alignment-respecting dynamics with classical SB trajectories that favor nearby points under Brownian motion.
  • 3.2 PAIRED SCHRÖDINGER BRIDGES AS PRIOR PROCESSES: Unlike prior reference-process strategies based on Gaussian-to-data training or Gaussian approximations, the method shapes the drift directly from alignments sampled from π⋆.The learned aligned drift is therefore tied to the original endpoint-interpolation problem rather than a related surrogate.

4 EXPERIMENTS

Across synthetic, cellular, and protein-docking experiments, SBALIGN uses aligned data to learn or improve stochastic trajectories and outperforms relevant baselines. Its evaluations also show strong performance with partial alignments and only 10 simulation steps, while protein docking remains a proof of concept.

  • 4 EXPERIMENTS: The experiments cover 2-dimensional synthetic interpolation, genetically barcoded cell differentiation, and protein structures filtered for substantial conformational change.The protein dataset contains 2,370 examples with provided Cα RMSD > 3.0Å, reduced to 1,591 after preprocessing and split into train, validation, and test sets.
  • 4.1 SYNTHETIC EXPERIMENTS: SBALIGN reproduces target alignments on the Moon and T synthetic datasets, whereas Brownian-prior FBSB tends to map points to nearby rather than aligned destinations.A learned SBALIGN drift can also serve as a reference drift that enables FBSB to recover desired alignments without hand-crafted time-varying drifts.
  • 4.2 CELL DIFFERENTIATION: SBALIGN outperforms FBSB on all reported cell-differentiation metrics, including distributional and alignment-based measures, while predicting overall differentiation trends.It struggles to isolate rare cell types, but using noisy barcode alignments to learn a prior and fine-tuning FBSB improves distributional scores.
  • 4.3 CONFORMATIONAL CHANGES IN PROTEIN DOCKING: SBALIGN outperforms EGNN by a large margin for protein conformational changes and predicts almost 70% of examples with RMSD < 5Å.The evaluation uses RMSD between predicted and true bound structures and fractions below 2.0, 5.0, and 10.0Å.
  • 4.3 CONFORMATIONAL CHANGES IN PROTEIN DOCKING: SBALIGN achieves impressive protein-conformation performance with just 10 simulation steps, leaving the speed–quality tradeoff for future work.The authors frame this as a response to the slow sampling speed often associated with diffusion models.
  • 4.3 CONFORMATIONAL CHANGES IN PROTEIN DOCKING: For protein docking, the study models conformational changes between apo and holo structures but does not provide a complete docking solution.Combining SBALIGN with newer rigid-protein docking methods is identified as future work.

5 CONCLUSION

The paper presents SBALIGN as a diffusion Schrödinger bridge framework for aligned-data interpolation. It combines Schrödinger bridge theory with Doob’s h-transform to avoid IPF-like training and reports improvements across synthetic and real-world tasks.

  • 5 CONCLUSION: SBALIGN combines Schrödinger bridge theory with Doob’s h-transform to solve diffusion Schrödinger bridges with aligned data.The framework derives novel loss functions and avoids the iterative proportional fitting procedure used by prior methods.
  • 5 CONCLUSION: The proposed losses are numerically stable because they do not rely on iterative proportional fitting.The paper verifies the framework on varied synthetic and real-world tasks and reports noticeable improvement over previous state-of-the-art methods.

(Supplementary Material)

The supplementary material identifies the authors and their institutional affiliations.

  • Supplementary Material: The paper is authored by Vignesh Ram Somnath, Matteo Pariset, Ya-Ping Hsieh, Maria Rodriguez Martinez, Andreas Krause, and Charlotte Bunne.The listed affiliations are ETH Zürich, IBM Research Zürich, and EPFL.

A ADDITIONAL RESULTS

The appendix motivates direct parameterization of Doob’s h-transform score because conditional probabilities are difficult to approximate and numerically unstable, especially near t≈0.

  • A.1 VARIANCE REDUCTION: Faithful approximation of the conditional probability requires good early-training paths, exponentially many trajectories, and close agreement between conditional and unconditional endpoint samples.These requirements are difficult to satisfy in high-dimensional spaces and before the drift has been learned.
  • A.1 VARIANCE REDUCTION: Directly parameterizing mϕ_t≈∇log h_t sidesteps these numerical difficulties while allowing the score magnitude to be controlled and regularized.
  • A.1 VARIANCE REDUCTION: P(X1 = x1|Xt = x) spans 11 orders of magnitude and is smallest near t≈0, making direct manipulation numerically challenging.The resulting score errors can be amplified across timesteps and lead trajectories astray.

B.1 SYNTHETIC DATASETS

The synthetic datasets test whether methods recover prescribed alignments rather than merely connecting marginal distributions by short or analytically convenient paths.

  • B.1 SYNTHETIC DATASETS: Figure 5 displays initial and final marginals for the moon and T datasets, with arrows marking selected alignments.
  • B.1 SYNTHETIC DATASETS: The moon dataset rotates one noisy semicircular distribution by 233° and requires methods to use alignment information rather than nearest-endpoint connections.Classic generative models choose shortest paths between nearby ends of the moons.
  • B.1 SYNTHETIC DATASETS: The T dataset exposes swapped pairings from classical Schrödinger bridges with a Brownian prior, while resisting simple symmetric or time-constant reference drifts.It therefore motivates general plug-and-play methods for approximate reference drifts.

B.2 CELL DIFFERENTIATION DATASETS

The biological datasets use lineage or structural correspondence to evaluate aligned stochastic trajectories for cell differentiation and protein conformational change.

  • B.2 CELL DIFFERENTIATION DATASETS: SBALIGN uses genetic barcodes to trace progenitor cells at t into descendants at t+1, then combines Doob’s h-transform with Brownian bridges to recover trajectories.Figure 6 presents this aligned cell-differentiation pipeline.
  • B.2 CELL DIFFERENTIATION DATASETS: Figure 7 compares ground-truth and SBALIGN cell-population marginals after projection onto the first two principal components.
  • B.2 CELL DIFFERENTIATION DATASETS: The cell dataset retains barcode-matched cells across days 2 and 4, excludes already differentiated day-2 cells, and produces 4702 non-overlapping pairs split 80%/10%/10%.The observations are reduced from 1622 gene features to 50 principal components.
  • B.3 CONFORMATIONAL CHANGES IN PROTEIN DOCKING: The protein task uses unbound and bound structures from 4330 proteins, retaining examples with provided Cα RMSD above 3Å before preprocessing.
  • B.3 CONFORMATIONAL CHANGES IN PROTEIN DOCKING: After Kabsch superposition and filtering, 1591 protein pairs are split into 1291/150/150 train, validation, and test examples, with Brownian bridges sampled between aligned states.
  • B.3 CONFORMATIONAL CHANGES IN PROTEIN DOCKING: Protein structures are represented using Cα coordinates, residue-level biochemical features, time embeddings, radial distance bases, and spherical-harmonic edge features.

C EXPERIMENTAL DETAILS

The experiments define distributional, matching, perturbation, and cell-fate metrics, with evaluation procedures tailored to aligned cell-differentiation trajectories and their predicted endpoint populations.

  • C.1.1 Cell Differentiation: Cell differentiation evaluation compares predicted and observed endpoint marginals using Wε, MMD, and ℓ2-based perturbation-signature distances.These metrics support comparison with FBSB, which operates only at the cell-distribution level.
  • C.1.1 Cell Differentiation: MMD is estimated unbiasedly with an RBF kernel and averaged across length scales 2, 1, 0.5, 0.1, 0.01, and 0.005.
  • C.1.1 Cell Differentiation: The perturbation-signature ℓ2 metric measures the distance between observed and predicted differences in feature means.
  • C.1.1 Cell Differentiation: RMSD evaluates aligned SBALIGN matchings by measuring predicted versus observed cell statuses at day 4.When squared, it represents the mean squared norm of status differences.
  • C.1.1 Cell Differentiation: Cell-type classification accuracy is computed by training an MLP on observed cells and evaluating predicted day-4 trajectories against ground-truth labels.The reported subset accuracy counts labels coinciding with the ground truth.

C.2.1 Cell Differentiation and Synthetic Datasets

The models use multilayer perceptrons to encode spatial and temporal information and predict drift magnitudes, with dropout and constant optimized diffusivity. Protein docking instead uses an SE(3)-equivariant graph neural network with residue-neighborhood features.

  • C.2.1 Cell Differentiation and Synthetic Datasets: Three modules encode spatial coordinates or drift, time, and their concatenation before predicting drift magnitudes along each dimension.The spatial encoder is a 3-layer MLP, the temporal encoder uses sinusoidal embeddings followed by a 2-layer MLP, and the final MLP has three layers.
  • C.2.1 Cell Differentiation and Synthetic Datasets: Dropout of 0.1 follows every non-final linear layer, while the diffusivity function is set to a constant optimized during training.
  • C.2.2 Conformational Changes in Protein Docking: Protein docking models use an SE(3)-equivariant tensor-product graph network over up to 40 residue neighbors within 40Å, with residue, edge, and distance features.
  • C.3 HYPERPARAMETERS: The hyperparameter section introduces the selected hyperparameters and training procedures that are detailed in subsequent subsections.

C.3.1 Synthetic Tasks

Synthetic and cell-differentiation tasks tune activation functions and diffusivity constants with Ray Tune, while protein docking uses specified optimization, validation, and model-selection procedures. The reported training setups differ in batch size, epochs, parameter count, and validation protocol.

  • C.3.1 Synthetic Tasks: Synthetic-task tuning selects among four activations and diffusivity constants {1, 2, 5, 10}; selu performs marginally better and g = 1 yields optimal results.
  • C.3.2 Cell Differentiation: Cell-differentiation tuning finds noticeable gains from silu and optimal results at diffusivity constant g = 1 across the tested values.
  • C.3.3 Conformational Changes in Protein Docking: Protein docking uses AdamW with learning rate 0.001, batch size 2, ten sampled timepoints per epoch, and regularization strength 1.0 for mϕ.Training typically stops after 200 epochs without validation improvement, with exponential moving averages used for validation inference.
  • C.3.3 Conformational Changes in Protein Docking: The protein model has 0.54M parameters and trains for 200 epochs, whereas the EGNN baseline has 0.76M parameters and trains for 1000 epochs.The best model is selected by validation-set mean RMSD after trajectory simulation and then used for test inference.
Loading 2302.11419v3…