Differentiable Expectation-Maximisation and Applications to Gaussian Mixture Model Optimal Transport
Listen
Radio episode about this paper
Transcript
Introduction to the show: ident: AI Radio. Generated commentary on the latest Artificial Intelligence papers.
Tom: Today's paper: "Differentiable Expectation-Maximisation and Applications to Gaussian Mixture Model Optimal Transport".
Jane: This work presents methods to differentiate the Expectation-Maximisation (EM) algorithm, which is typically treated as a non-differentiable black box,
Tom: First, who's behind it and why it matters.
Title and authors: Tom: Let's talk about the title and who wrote this thing: "Differentiable Expectation-Maximisation and Applications to Gaussian Mixture Model Optimal Transport." It’s clear they are focusing on two major areas here, EM differentiation and the application to GMM Optimal Transport.
Jane: The authors are Samuel Boïté, Eloi Tanguy, Julie Delon, Agnès Desolneux, and Rémi Flamary. They come from a mix of institutions in Paris and Gif-sur-Yvette.
Lu: What's interesting is that they are specifically focusing on the Mixture Wasserstein distance between GMMs as the key application for this differentiable EM setup, which connects statistics with transport theory.
Meng: I wonder how much practical benefit we can expect from focusing on MW2 specifically when there are so many other ways to measure distribution similarity?
Lalam: The focus on MW2 suggests they are aiming for a loss function that is not only mathematically sound but also directly useful for image processing and machine learning tasks.
The paper's summary: Tom: So, what’s the core summary of this paper? It explains that they present several ways to compute the gradient of EM with respect to data, ranging from full automatic differentiation down to approximate methods like One-Step Gradient Approximation.
Jane: They are showing how these different strategies perform in terms of accuracy and computational cost for calculating gradients when you're dealing with an initialisation and a set number of iterations.
Lu: They compare the complexity of these methods, noting that Full Automatic Differentiation is quite costly, while the Approximate Implicit Gradient has a specific complexity involving the inverse of a matrix.
Meng: The comparison between the complexities—O(nK2d4) for AI versus O(Kd2 + nd) for OS—is something I’ll pay close attention to when we think about deployment on real hardware.
Lalam: It sounds like they are providing a practical toolkit so researchers aren't stuck with just one way to get gradients from EM, which is very helpful for flexibility in research.
The paper's improvements: Tom: Moving into the improvements section, the authors highlight a few things they did better than previous approaches. They specifically detail the advantages of using Approximate Implicit Gradient over the One-Step gradient approximation regarding median Mean Squared Error.
Jane: They also provide some stability results for MW2; they show that if the GMM parameters are close due to EM convergence, then the MW2 costs will also be close, which justifies using it as a loss function.
Lu: The paper introduces a novel unbalanced variant called UMW2, which penalizes marginal conditions instead of enforcing them in the balanced formulation, suggesting it might be more stable and easier to optimize with minibatch sampling.
Meng: That unbalanced variant sounds promising for practical use because stability with respect to minibatch sampling is a major hurdle when we try to implement these kinds of metrics in production systems.
Lalam: The introduction of UMW2 seems like a smart move by the authors, trying to address the stability issues they found in the standard balanced formulation by making it more amenable to optimization.
Conclusion: Tom: So, wrapping up this discussion on "Differentiable Expectation-Maximisation and Applications to Gaussian Mixture Model Optimal Transport," the main point is that they've successfully shown how to make EM differentiable enough for loss calculation using MW2.
Jane: They’ve mapped out different differentiation strategies and provided theoretical backing for why MW2 works well with the EM process when parameters are close, which is a big step forward.
Lu: The authors show that this differentiability allows us to apply these GMMs in concrete tasks like Barycentre Flow in 2D and Colour Transfer, which shows real versatility.
Meng: For practical application, the focus on providing an unbalanced variant suggests they’re steering this toward more robust training scenarios where we deal with noisy data or sampling techniques.
Lalam: Overall, this work helps shift GMM fitting from a static black box into a trainable part of generative modeling pipelines, which is really powerful for advancing our AI capabilities.
Université Paris Cité · Centre Borelli, CNRS and ENS Paris-Saclay Centre Borelli, CNRS and École Polytechnique, Institut Polytechnique de Paris
cs.LG, math.PR, stat.ML
Submitted: 2025-09-02
Updated: 2026-09-30
License: http://arxiv.org/licenses/nonexclusive-distrib/1.0/
Importance score: 75/100
The gist: This work presents methods to differentiate the Expectation-Maximisation (EM) algorithm, which is typically treated as a non-differentiable black box, and applies this differentiability to Gaussian
Key concepts
- Mixture Wasserstein distance (MW2)
- This is a specialized version of the Wasserstein distance used to compare two Gaussian Mixture Models. It measures how different the components of the two models are by solving a discrete transport problem between them, offering a way to use MW2 as a differentiable loss function.
- Approximate Implicit Gradient (AI)
- This strategy computes the gradient of EM without needing to solve complex systems directly. It uses an approximation involving an inverse matrix, which is theoretically exact under certain conditions but requires solving a large linear system for computation.
- One-Step Gradient (OS)
- This is a numerically simpler way to approximate the EM gradient by only calculating the gradient for the final iteration. While computationally cheap, it introduces approximation errors and relies on assumptions that are not always true for EM convergence.
- Unbalanced GMM-OT (UMW2)
- A novel variant of MW2 that penalizes marginal conditions instead of enforcing them. This version is proposed as a potentially more stable alternative to the balanced formulation, especially when using minibatch sampling in optimization.
Terminology
Summary
This work presents methods to differentiate the Expectation-Maximisation (EM) algorithm, which is typically treated as a non-differentiable black box, and applies this differentiability to Gaussian Mixture Models (GMMs) through the Mixture Wasserstein distance (MW2). This allows MW2 to be used as a differentiable loss function in imaging and machine learning tasks. The paper explores various differentiation strategies—from full automatic differentiation to approximate methods—and provides theoretical justifications for using MW2 with EM, contributing novel stability results and unbalanced variants of the distance.
Differentiation Strategies for EM
The paper investigates several approaches to compute the gradient of the EM algorithm with respect to data, given an initialisation and number of iterations. These methods include:
-
Full Automatic Differentiation (AD): Considered a
natural baseline
for computing the exact gradient up to numerical precision, though it can be costly. -
Approximate Implicit Gradient (AI): Approximates the gradient using the chain rule derived from Proposition 2.1, yielding an expression involving the inverse of a matrix:
/I − ∂F/∂θ(θT, X)−1∂F/∂X(θT, X) (Eq. 10). This method is theoretically exact when θT = θ∗ but requires solving a large linear system. The complexity is noted as O(nK2d4). The AI gradient is superior to the One-Step gradient (OS) in terms of median MSE, though it suffers from high variance. It outperforms OS when the spectral norm of the Jacobian is close to 0, and performs substantially better than OS for larger T. 3. Approximate Implicit Gradient (AI) complexity: O(nK2d4). The One-Step gradient (OS) complexity: O(Kd2 + nd). The Full Automatic Differentiation (AD) complexity: O(T(nKd2 + Kd3)).
- One-Step Gradient Approximation (OS): Approximates the gradient by only computing the last step, neglecting the dependence of penultimate iterations on X. It is numerically inexpensive but suffers from approximation error and requires a contraction assumption that is not verified for EM.
Gaussian Mixture Model Optimal Transport (GMM-OT)
The paper focuses on leveraging differentiable EM in the computation of MW2 between GMMs, which compares GMMs by matching their components using a small-scale discrete OT problem. Key aspects include:
/Mixture Wasserstein distance (MW2): A variant of the Wasserstein distance restricted to couplings being GMMs on Rd × Rd, defined as MW22(µ0, µ1). If the parameters of the GMMs are known, computing Eq. (14) involves evaluating K0 × K1 Wasserstein distances between Gaussians and solving a discrete transport problem of size K0 × K1. The distance has a closed-form expression for Gaussian measures: W22(µ, µ˜) = ∥m − m˜∥22 + tr(Σ + Σ˜ − 2(Σ1/2Σ'1/2)).
**/Stability of MW2: A stability result shows that if the GMM parameters are sufficiently close (thanks to EM convergence), then the MW2 costs will also be close, providing theoretical justification for using MW2 as a loss function. Proposition 3.1 quantifies the decrease of MW22(ˆµ, µ) when µˆ is an estimator of µ, and Proposition 3.2 quantifies the decrease of MW22(ˆµ0, ˆµ1) − MW22(µ0, µ1) when µˆi are estimators of µi. The empirical study shows that EM convergence translates into precision on the MW2 distance for well-separated cases (rate O(n−1/2)), but plateaus in weak separation regimes. **
**/Unbalanced GMM-OT: A novel unbalanced variant, UMW22(µ, ν; λ0, λ1), is introduced by penalizing marginal conditions instead of enforcing them. This variant is proposed as a possibly more stable alternative to the balanced formulation due to its amenability to optimisation and stability with respect to minibatch sampling. **
Applications of Differentiable EM
The paper illustrates the versatility of differentiable EM in several applications:
-
Barycentre Flow in 2D: Optimizing a point cloud X towards a barycentre of GMMs νi fitted from (Yi), solving min X∈Rn×2 MW22(µ(FTX (θ0)), νi).
-
Colour Transfer: Optimizing an image X to minimize the MW2 cost between a GMM fitted on X and a target GMM ν, using the Warm-Start EM method with fixed uniform GMM weights (Algorithm 2) to avoid local minima. The unbalanced variant is shown to be more robust to outliers in the target distribution.
Improvements for AI systems
As a fastidious and diligent researcher, I have thoroughly analyzed this paper on Differentiable Expectation-Maximisation (EM) and Optimal Transport (OT). The core contribution is enabling end-to-end gradient propagation for GMM fitting by differentiating the EM algorithm with respect to the input data.
Here are the specific improvements that can be made to AI systems, categorized by application:
),
-
Efficient and stable training of latent variable models (Gaussian Mixture Models).
-
Differentiable loss functions for generative tasks (Style Transfer, Image Generation).
-
Improved barycentre computation for complex data structures in machine learning.
Here are the specific improvements and capabilities:
- Efficient and stable training of latent variable models (Gaussian Mixture Models).
The system can now use a differentiable EM process to fit GMMs directly within deep learning pipelines, enabling optimization via gradient descent rather than relying on black-box iterative solvers. This is crucial for training complex generative models where the likelihood function is intractable.
- Differentiable loss functions for generative tasks (Style Transfer, Image Generation).
The paper provides a novel way to use the Mixture Wasserstein distance (MW2) as a differentiable loss function between GMMs and target distributions.
The improved AI system can perform:
-
Color Transfer with enhanced robustness against outliers in the target distribution.
-
Neural Style Transfer, allowing for more precise control over style application by minimizing MW2 loss between feature distributions across different layers of a VGG network.
-
Image Generation via MW2-based Generative Adversarial Networks (MW2-GANs). This allows the generator to be guided by a differentiable measure of distribution matching, potentially leading to higher quality and more consistent generated images compared to standard GANs relying solely on pixel loss.
- Improved barycentre computation for complex data structures in machine learning.
The system can solve challenging barycentre problems between multiple GMMs (e.g., three 2D images) by minimizing the MW2 distance between the EM estimates of point clouds towards these target GMMs. This allows for:
-
Generating synthetic data points that lie at the optimal geometric location (barycentre) of a set of learned distributions.
-
Computing generalized barycentres in higher dimensions, which is currently computationally prohibitive, by leveraging the discrete formulation of MW2 and differentiable EM flows.
This approach fundamentally shifts GMM fitting from a non-differentiable black box to a trainable component within a deep learning framework, unlocking new avenues for optimization in generative modeling and geometric data analysis.
Sources
- Nonlinear reduced basis using mixture Wasserstein barycenters: application to an eigenvalue problem inspired from quantum chemistry
- Minibatch optimal transport distances; analysis and applications
- Sinkhorn EM: An Expectation-Maximization algorithm based on entropic optimal transport
- Very Deep Convolutional Networks for Large-Scale Image Recognition
- Parsimonious Gaussian mixture models with piecewise-constant eigenvalue profiles
- Computing Barycentres of Measures for Generic Transport Costs
- Constrained Approximate Optimal Transport Maps
- A note on the relations between mixture models, maximum-likelihood and entropic optimal transport
- Causal Expectation-Maximisation
Related papers
- Polynomial-Augmented Neural Networks (PANNs) with Weak Orthogonality Constraints for Enhanced Function and PDE Approximation
- AIRL-S: Unifying Reinforcement Learning and Search-Based Test-Time Scaling via Adversarial Inverse Reinforcement Learning
- Transformers as Bayesian In-Context Experimenters: Smoothness-Adaptive Efficient ATE Estimation
- Convergence issues in Relational Concept Analysis based on AOC-posets
- Beliefs Beyond Posteriors: Local-Consistency Optimisation for Bayesian Neural Networks
- Understanding Diffusion Models via Ratio-Based Function Approximation with SignReLU Networks