Organizations: Carnegie Mellon University, Department of Statistics · Amazon Web Services · Amazon Web Services (AWS) LogAnalytics · Carnegie Mellon University, Machine Learning Department · Cornell University, Department of Electrical and Computer Engineering
Optimal transport (OT) and Gromov-Wasserstein (GW) alignment provide interpretable geometric frameworks for comparing, transforming, and aggregating heterogeneous datasets---tasks ubiquitous in data science and machine learning. Because these frameworks are computationally expensive, large-scale applications often rely on closed-form solutions for Gaussian distributions under quadratic cost. This work provides a comprehensive treatment of Gaussian, quadratic cost OT and inner product GW (IGW) alignment, closing several gaps in the literature to broaden applicability. First, we treat the open problem of IGW alignment between uncentered Gaussians on separable Hilbert spaces by giving an exact variational characterization through a quadratic optimization over isometries and co-isometries (orthogonal matrices in finite dimensions), for which we derive tight analytic upper and lower bounds. If at least one Gaussian measure is centered, the solution reduces to a fully closed-form expression, which we further extend to an analytic solution for the IGW barycenter between centered Gaussians. We also present a reduction of Gaussian multimarginal OT with pairwise quadratic costs to a tractable optimization problem and prove that every second-order stationary point of this problem is globally optimal. To demonstrate utility, we compare embedding distributions of already-trained language-model distillations and cluster synthetic users using covariance spectra of their text embeddings.
Figures & tables
Figure 1: (a) The OT problem compares measures over the same space by searching for the coupling π which minimizes expected transportation cost c(x,y) . Here we depict the quadratic-cost OT problem, with c(x,y)=∥x−y∥22 . (b) GW alignment compares measures over (possibly different) spaces by searching for the coupling π which matches the kernels kX and kY as well as possible. Here, we depict the inner product Gromov-Wasserstein (IGW) problem with kX(⋅,⋅)=kY(⋅,⋅)=⟨⋅,⋅⟩ , which seeks to find a coupling which is as close to unitary as possible (roughly, preserving pairwise angles).
(a) IGW :
(a) W2 :
(b) IGW :
(b) W2 :
t=0
t=0.33
t=0.67
t=1
Figure 2: Comparison of displacement interpolations between two Gaussian distributions (origin marked in red). (a) has the origin at the bottom and (b) has the origin at the top; the distributional orientation along the IGW interpolation depends on the origin, while the W2 interpolation is invariant to translating both measures. Particles follow straight segments in both interpolations. Top row of (a) and (b): IGW displacement interpolation between two Gaussians with OT map estimated using RGD. Bottom row of (a) and (b): 2-Wasserstein displacement interpolation with OT map from the Bures-Wasserstein formula in ( 4 ).
Figure 3: Contour plots of the ρ -weighted 2-Wasserstein barycenter and ρ -weighted IGW barycenter. The 2-Wasserstein barycenter does not preserve the covariance structure of the input measures, but the IGW barycenter naturally does.
Figure 4: Comparison of multimarginal OT between two sets of Gaussian distributions. Top row: multimarginal OT between three aligned Gaussians. Bottom row: multimarginal OT between three misaligned Gaussians.
Figure 5: Upper and lower bounds on the IGW distance between the bert - base - uncased model and its distillations on the (a) amazon_polarity and (b) ag_news datasets. The analytic upper bound is tightened by numerically optimizing over the Stiefel manifold. We see that some of the smaller distillations preserve the embedding distribution almost as well as the larger ones. The resulting upper and lower bounds are usually extremely tight in practice.
Figure 6: CKA between the bert - base - uncased model and its distillations on the (a) amazon_polarity and (b) ag_news datasets, compared to the upper bound on the IGW distance obtained by RGD. Pearson correlation is −0.504 ( p=0.014 ) for amazon_polarity and −0.198 ( p=0.365 ) for ag_news .
Figure 7: (a) 2-D MDS embedding of the IGW distance matrix between users, colored by the embedding model used for each user. (b) The same MDS embedding, colored by the ground truth interests of each user. (c) Euclidean k -means++ clustering of the zero-padded ordered covariance spectra of the heterogeneous users. The IGW distance is able to almost recover the original clusters, even only using heterogeneous covariance information.
Figure 8: Comparison of the Burer-Monteiro and interior-point approaches for solving multimarginal OT between p Gaussian measures on R3 : (a) computation times, (b) covariance-pairing objectives J=∑i<jtr(Cij) , with relative error ∣JBM−JSDP∣/∣JSDP∣ , and (c) number of optimization variables.
Appendix figures & tables8 assets
Supplementary material from the paper’s appendix.
Appendix
H=128
H=256
H=512
H=768
L=2
4.4
9.7
22.8
39.2
L=4
4.8
11.3
29.1
53.4
L=6
5.2
12.8
35.4
67.5
L=8
5.6
14.4
41.7
81.7
L=10
6.0
16.0
48.0
95.9
L=12
6.4
17.6
54.3
110.1
Appendix
Table 1: Number of parameters of the bert model distillations of Turc et al. [ 68 ] (in millions), as a function of the number of transformer layers L and hidden embedding dimension H . The teacher model bert - base - uncased has L=12 and H=768 , with 110.1 million parameters.
Figure 9: First two principal components of the embeddings produced by the bert - base - uncased model on the (a) amazon_polarity and (b) ag_news datasets (blue), along with contour lines from a Gaussian fit (red). The leading two-dimensional projections are visually consistent with an approximate Gaussian model.
Figure 10: (a) 2-D multidimensional scaling (MDS) embedding of the IGW distance matrix between users, colored by their ground truth interests. (b) Euclidean k -means++ clustering of the ordered covariance spectra, whose arithmetic means are the spectral IGW barycenters. The IGW distance is able to almost recover the original clusters, using only covariance information.
Representation
ARI
NMI
Mean only
1.000±0.000
1.000±0.000
Covariance spectrum only
1.000±0.000
1.000±0.000
Combined
1.000±0.000
1.000±0.000
Appendix
Table 2: Adjusted Rand index (ARI) and normalized mutual information (NMI), reported as mean and standard deviation over 10 repeated user samples, using the same embedding model, users, and clustering procedure for all representations.
Figure 11: Computation times at a common relative objective-gap target of 10−7 , including computation of the coupling matrices Ui . Curves show medians and shaded bands show ranges over the 15 runs in each setting. The specialized fixed-point method is faster at the matched accuracy in every tested setting.
Diagnostic over 240 runs
RGD
Fixed point
Maximum Riemannian gradient norm
9.29×10−4
1.57×10−3
Maximum relative objective gap
1.00×10−7
9.99×10−8
Maximum relative covariance-constraint error
4.08×10−15
3.77×10−14
Appendix
Table 3: Optimization accuracy at the same stopping tolerance. The relative objective gap is defined above; the covariance-constraint error is maxi(∥UiUi⊺−Σi∥F/∥Σi∥F) . All per-run results are provided with the code.
Figure 12: A visual depiction of the Monge mass transportation problem.
Figure 13: Relationships between various distances on spaces of objects and distributions.
The Applied Mathematics and Physics Department, Graduate School of Informatics, Kyoto University, Kyoto, Japan · Daniel Guggenheim School of Aerospace, Georgia Institute of Technology, Atlanta, USA