Neural method learns two-way optimal transport maps from unpaired samples

CyclOT: Learning Quadratic Optimal Transport Maps via Synchronized Forward-Backward Interpolants

Machine Learning

Summary

The paper tackles how to learn efficient mappings between two different data sets without knowing which points match each other. The authors propose a neural network framework that simultaneously learns forward and reverse transformations, ensuring they are consistent with each other. This approach avoids needing predefined matchings or special mathematical forms, and they prove guarantees about the accuracy of the learned maps under certain conditions. They also test their method on several types of data, including images and biological measurements, to check how well it preserves structure and performs transformations.

What this means in practice

  • For medical image teams: Improve mapping between different medical image sets for tasks like disease progression analysis using learned forward and reverse transport maps.
  • For computational biology teams: Align single-cell perturbation data by learning bidirectional transformations to better understand cellular responses across conditions.

Authors

Shizhou Xu, Jiachen Liu, Shih-Hsin Wang, Stefan Broecker, Yuhao Huang, Bao Wang, Thomas Strohmer

Abstract

We study the recovery of forward and reverse quadratic optimal-transport maps from unpaired samples in high dimensions. We introduce a bidirectional neural framework in which the learned maps induce forward and backward displacement interpolants, while the training objective combines bidirectional quadratic action, discriminator-restricted Jensen-Shannon endpoint objectives, and two-sided cycle consistency. The construction requires neither precomputed sample pairings nor an explicit convex-potential parameterization. For absolutely continuous probability measures supported on a compact convex set, and under the stated generator-approximation, discriminator-richness, and minimizer-attainment conditions, we prove a population recovery theorem: for every prescribed accuracy, the sum of the corresponding \(L^2\) errors between any global minimizer and the forward and reverse quadratic Brenier maps is below that accuracy, provided the discriminator level is sufficiently large and the annealing action weight becomes sufficiently small. Moreover, the cycle loss is bounded above by \(λW_2^2(μ_0,μ_1)\). Complementary results quantify approximate invertibility and show that exact endpoint Jensen-Shannon divergence and cycle consistency control missing target mass and many-to-one map collapse, respectively. Experiments on Swiss roll, MNIST, CelebA, single-cell perturbation data, and chest X-ray images evaluate endpoint fidelity, transport cost, inverse consistency, and the geometry of the induced interpolations.