cs.LGSep 2, 2025

Differentiable Expectation-Maximisation and Applications to Gaussian Mixture Model Optimal Transport

Authors: Samuel Boïté, Eloi Tanguy, Julie Delon, Agnès Desolneux, Rémi Flamary

Organizations: Université Paris Cité, CNRS, MAP5, F-75006 Paris, France · Centre Borelli, CNRS and ENS Paris-Saclay, F-91190 Gif-sur-Yvette, France · CMAP, CNRS, Ecole Polytechnique, Institut Polytechnique de Paris

Abstract

The Expectation-Maximisation (EM) algorithm is a central tool in statistics and machine learning, widely used for latent-variable models such as Gaussian Mixture Models (GMMs). Despite its ubiquity, EM is typically treated as a non-differentiable black box, preventing its integration into modern learning pipelines where end-to-end gradient propagation is essential. In this work, we present and compare several differentiation strategies for EM, from full automatic differentiation to approximate methods, assessing their accuracy and computational efficiency. As a key application, we leverage this differentiable EM in the computation of the Mixture Wasserstein distance MW2\mathrm{MW}_2 between GMMs, allowing MW2\mathrm{MW}_2 to be used as a differentiable loss in imaging and machine learning tasks. To complement our practical use of MW2\mathrm{MW}_2, we contribute a novel stability result which provides theoretical justification for the use of MW2\mathrm{MW}_2 with EM, and also introduce a novel unbalanced variant of MW2\mathrm{MW}_2. Numerical experiments on barycentre computation, colour and style transfer, image generation, and texture synthesis illustrate the versatility of the proposed approach in different settings.

Figures & tables

Appendix figures & tables21 assets

Supplementary material from the paper’s appendix.

Appendix

Explore similar work

Sep 25, 2026cs.LG

Averaged Mirror Descent and Dual Gradient Methods: Convergent Algorithms for Entropic Gromov-Wasserstein Problems

The Gromov-Wasserstein (GW) distance measures the discrepancy between metric measure (mm) spaces and identifies optimal alignments between them based solely on their intrinsic structure. Since it identifies isomorphic mm spaces, it provides a natural notion of distance for heterogeneous datasets which may admit isomorphic representations. In order to accelerate computation of GW distances, many practitioners employ entropic regularization to obtain an Entropic GW (EGW) problem. The most popular EGW solver is the Mirror Descent (MD) algorithm, which reduces EGW computations to an iterative process where an entropic optimal transport (EOT) problem is solved at each iteration. Despite its widespread use, the convergence of MD for this problem has only been established for restricted classes of costs. On the other hand, a recently proposed dual gradient method is available for general costs, but requires a choice of step size which depends on the regularization parameter. To address these two issues, we introduce Averaged Mirror Descent (AMD), which averages consecutive MD steps, and prove its convergence for arbitrary costs. Then, we establish that the dual gradient method with a fixed step size also converges for arbitrary costs at the cost of a more complicated iteration. In both cases, we also account for inexact iterations which are inescapable in practice. We compare the empirical performance of these methods across various settings and, in particular, show that AMD and the dual gradient method both converge on an example where classical MD fails.
May 6, 2026cs.LG

On the Wasserstein Gradient Flow Interpretation of Drifting Models

Recently, Deng et al. (2026) proposed Generative Modeling via Drifting (GMD), a novel framework for generative tasks. This note presents an analysis of GMD through the lens of Wasserstein Gradient Flows (WGF), i.e., the path of steepest descent for a functional in the space of probability measures, equipped with the geometry of optimal transport. Unlike previous WGF-based contributions, GMD can be thought of as directly targeting a fixed point of a specific WGF flow. We demonstrate three main results: first, that one algorithm proposed by Deng et al. (2026) corresponds to finding the limiting point of a WGF on the KL divergence, with Parzen smoothing on the densities. Second, that the algorithm actually implemented by Deng et al. (2026) corresponds to a different procedure, which bears some resemblance to the fixed point of a WGF on the Sinkhorn divergence, but lacks certain desirable properties of the latter. Third, the same same idea can be extended to the limiting point of other WGFs, including the Maximum Mean Discrepancy (MMD), the sliced Wasserstein distance, and GAN critic functions.
May 14, 2026cs.LG

Distance-Matrix Wasserstein Statistics for Scalable Gromov--Wasserstein Learning

Gromov--Wasserstein (GW) distances compare graphs, shapes, and point clouds through internal distances, without requiring a common coordinate system. This invariance is powerful, but discrete GW is a nonconvex quadratic optimal transport problem and is difficult to estimate at scale. We propose \emph{Distance-Matrix Wasserstein} (DMW), a hierarchy of Wasserstein statistics comparing laws of random finite distance matrices. Rather than optimizing a global point-level alignment, DMW samples nn points from each space, records their pairwise distances, and transports the resulting matrix laws. We prove that DMW is a relaxation and lower bound of GW, and establish a reverse approximation inequality: the GW--DMW gap is controlled by the Wasserstein error of approximating each original measure with nn samples. Hence population DMW converges to GW as sampled subspaces become dense. We further give finite-sample bounds, including intrinsic-dimensional rates that depend on the data manifold rather than the ambient matrix dimension (n2)\binom n2. For scalable computation, we introduce sliced and multi-scale DMW; for p=1p=1, the sliced multi-scale dissimilarity yields positive-definite exponential kernels. Experiments on synthetic metric spaces, scalability benchmarks, graph classification, and two-sample testing validate the theory and demonstrate an interpretable GW-style proxy for structural comparison.