RW-Flow: One-Step Generation on Compact Manifolds via Wasserstein Gradient Flows
Organizations: KAIST
Abstract
Manifold-valued data, and consequently the distributions they induce, are prevalent across many domains, ranging from the locations of geospatial events, such as earthquakes, to biomolecular torsion angles that encode information about three-dimensional structure. While diffusion and flow-based generative models have been successfully extended to compact manifolds, sampling typically requires tens or hundreds of sequential network evaluations. We introduce RW-Flow, a theoretically grounded framework for learning one-step generative models on compact manifolds via Wasserstein gradient flows. The main challenge is identifiability: driving the velocity field to zero should guarantee that the model distribution matches the target distribution. We establish a necessary and sufficient condition for identifiability on compact, connected Riemannian manifolds. We specifically show that, for a symmetric, Lipschitz-continuous cost function, the velocity field induced by the Sinkhorn divergence is identifiable if and only if the associated Gibbs kernel is nondegenerate. This characterization provides a general principle for designing identifiable costs on compact manifolds. It also reveals that the squared geodesic distance, the natural manifold analogue of the squared Euclidean distance, does not always guarantee identifiability. Across benchmarks involving geospatial events, protein side chain torsion angles, RNA backbone torsion angles, and general manifolds discretized as triangular meshes, RW-Flow outperforms existing one-step methods in nearly all settings under fair comparison conditions.
Figures & tables
| Cost | Gibbs kernel | Identifiable for | ||
| Squared Geodesic | a.e. on , | |||
| Spectral | all , any | |||
| Chordal | all , any | |||
| Geodesic | all on |
| Kernel | Parameters | |
| Matérn | ||
| Heat | ||
| Sub. Heat |
| Volcano | Earthquake | Flood | Fire | |||||||||
| Method | kMMD | COV | 1-NNA | kMMD | COV | 1-NNA | kMMD | COV | 1-NNA | kMMD | COV | 1-NNA |
| Held-out | ||||||||||||
| RFM | ||||||||||||
| GFM-L | ||||||||||||
| GFM-E | ||||||||||||
| GFM-S | ||||||||||||
| General | Glycine | Proline | Prepro | RNA | |||||||||||
| Method | kMMD | COV | 1-NNA | kMMD | COV | 1-NNA | kMMD | COV | 1-NNA | kMMD | COV | 1-NNA | kMMD | COV | 1-NNA |
| Held-out | |||||||||||||||
| RFM | |||||||||||||||
| GFM-L | |||||||||||||||
| GFM-E | |||||||||||||||
| GFM-S | |||||||||||||||
| Bunny ( ) | Bunny ( ) | Spot ( ) | Spot ( ) | |||||||||||||
| Method | kMMD | COV | 1-NNA | Time (s) | kMMD | COV | 1-NNA | Time (s) | kMMD | COV | 1-NNA | Time (s) | kMMD | COV | 1-NNA | Time (s) |
| Held-out | — | — | — | — | ||||||||||||
| RFM 1000 | ||||||||||||||||
| RW-Flow -H | ||||||||||||||||
Appendix figures & tables7 assets
Supplementary material from the paper’s appendix.
Appendix
| Symbol | Meaning |
|---|---|
| Manifold and geometry | |
| compact, connected, smooth Riemannian manifold with metric ; and denote the unit sphere and the flat torus of dimension | |
| ambient Euclidean space in which is embedded for the network parametrization | |
| tangent space of at ; and are the metric and its norm on | |
| , | exponential map at and its inverse on the domain where it is defined |
| geodesic distance on | |
| Energy Functional | Identifiability Condition | |
| Sinkhorn | nondegenerate, symmetric Lipschitz | |
| MMD | characteristic | |
| Smoothed KL (KGD) | characteristic, normalized |
| Data | |
| Datasets | : volcano, earthquake, fire, flood ( Mathieu & Nickel, 2020 ) ; : Top500 backbone torsion angles ( Lovell et al., 2003 ) by amino-acid type (General, Glycine, Proline, Pre-Pro); : RNA torsion angles ( Murray et al., 2003 ) ; all following the conventions of Chen & Lipman (2024) . |
| Coordinates | ambient unit vector on ; angles in on |
| Split | 80/10/10 train/validation/test, computed once with a split seed independent of the training seed and read by every implementation; validation for model selection, test once for the reported number |
| Network | |
| Architecture | MLP, ; no normalisation, residual connections or time embedding; PyTorch default initialisation |
| Width | 1024 on ; 512 on and |
| RNA | ||||
| Method | kMMD | MMD | COV | 1-NNA |
| Held-out | ||||
| RFM | ||||
| GFM-L | ||||
| GFM-E | ||||
| GFM-S | ||||