Scalable Supervised Optimal Transport of Gaussian Mixture Models
Abstract
Optimal Transport (OT) is a principled framework for comparing probability distributions, but its effectiveness depends critically on the ground metric, the cost used to compare observations. In high-dimensional settings, fixed ground metrics like the Euclidean distance can make OT distances reflect task-irrelevant variation rather than low-dimensional class structure, limiting OT in tasks such as classification and clustering. Supervised OT addresses this by learning a parameterized ground metric such that the induced OT distance reflects class structure. However, existing approaches on point clouds scale poorly with the number of observations. We propose a scalable supervised OT framework for Gaussian Mixture Models (GMMs) that lifts ground metric learning from observations to mixture components, replacing large point-cloud OT problems with smaller component-level problems whose transport costs admit closed forms. We define a learnable Wasserstein-type distance between GMMs using the Generalized Bures-Wasserstein (GBW) distance between Gaussian components as the ground metric, parameterized by a rectangular linear map that projects Gaussian components into a low-dimensional latent space inducing an OT distance between GMMs. For additional scalability, we introduce a diagonal approximation of the GBW metric, reducing the covariance computation between components to linear in the latent dimension. We bound the resulting approximation error by the mutual information between latent Gaussian variables and the Frobenius norm of the learned map, yielding a principled regularization strategy that also enables interpretation of latent axes. Across synthetic and scRNA-seq benchmarks, our method improves or matches classification and clustering over fixed-metric and supervised OT baselines on point clouds while substantially reducing OT distance computation time.