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
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 between GMMs, allowing MW2 to be used as a differentiable loss in imaging and machine learning tasks. To complement our practical use of MW2, we contribute a novel stability result which provides theoretical justification for the use of MW2 with EM, and also introduce a novel unbalanced variant of MW2. 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
Figure 1 : Left: sensitivity of the EM output to perturbation. Right: particle flow with repulsive energy.
Method
Time
Memory
Full Automatic Differentiation (AD)
O(T(nKd2+Kd3))
O(TKd2+nd)
Approximate Implicit Gradient (AI)
O(TnKd2+K3d6+nK2d5)
O(nK2d4)
One-Step gradient (OS)
O(T(nKd2+Kd3))
O(Kd2+nd)
Table 1 : Complexities of the backward passes (gradient computations) for T EM iterations on n points in Rd with K components.
Figure 2 : Sample complexity of MW22 for three specific GMMs of “low” to “high” mode separation. Curves show the median across 40 repetitions, with shaded interquartile ranges.
Figure 3 : Comparison of experimental setup and flows of ET using different methods. The dark shades of purple correspond to earlier iterations, and the yellow shades to later iterations. (Warm-Start and AI are almost identical to AD)
Figure 4 : Flow of ET with the AD method and standard EM Algorithm 1 . We compare two settings: one with non-uniform GMM weights (NU) and one with uniform weights (U).
Figure 5 : Stochastic Flow of ET for Algorithm 2 with the full automatic differentiation method. We vary the sub-sampling ratio r∈(0,1] , which corresponds to performing EM on only [r×n] random points from the current point cloud at each step.
Figure 6 : Varying the number of samples n and the number of iterations T , we study the convergence of EM, the local contractivity of F , and the MSEs of the OS and AI gradients against the AD gradient.
Figure 7 : EM−MW22 flow displacing particles in order to make their EM output approach a barycentre of three target GMMs.
Figure 8 : Colour transfer experiments. Top row: EM−MW22 from source to target. Bottom row: unbalanced colour transfer with regularisations λ1=10 (source) and λ2=0.1 (target).
Figure 9 : Style transfer method inspired by [ GEB15 ] : setup and example result.
Figure 10 : Choosing the number of components K for texture synthesis.
Figure 11 : Multi-scale texture synthesis with K=4 components for 8×8 patches
Appendix figures & tables21 assets
Supplementary material from the paper’s appendix.
Appendix
Figure 12 : Local behaviour of the EM map in a symmetric two-component example. The grey dashed line is the identity map.
Figure 13 : Two fixed GMMs μ (top) and ν (bottom) used in Section 3.2 . Columns correspond to low ( σ=10 ), medium ( σ=0.5 ), and high ( σ=0.1 ) separation.
Figure 14 : 3D representation of the three GMMs used in Section 4.5 to compare EM gradient methods.
Figure 15 : Relative MSE of the full automatic differentiation gradient (AD) against the finite differences approximation (FD) for three different GMMs and varying the FD step size εFD .
Sample
T
MW22 (warm)
MW22 (cold)
Difference
#1
1
676.5595
662.2502
14.3093
#1
2
674.8875
670.1658
4.7217
#1
5
674.6132
674.6144
0.0012
#1
10
674.6132
674.6132
0
#2
1
680.5627
622.3420
58.2207
#2
2
680.4759
629.5430
50.9329
Appendix
Table 2 : Comparison between warm-start and cold-start runs.
Figure 16 : Setup for the local minimum of W22(μα,η,ν) . The support of μ is represented with blue squares, and its weights with vertical blue lines. For the target ν , its support is red squares and its weights red lines. We consider a specific region where the points x1=η1 and x2=η2 stay within (−21,21) and the weights a1=61+α1 and a2=61+α2 stay within [0,31] , as represented by the orange rectangle. Likewise, the point x3=1+η3 must stay within (21,23) and its weight a3=32−(α1+α2) must stay in [31,1] , as shown with the purple rectangle.
Figure 17 : With the data Xε:=(x1,⋯,x6) , one iteration of the EM algorithm initialised at θ⋆ (corresponding to the GMM μ⋆ ) yields approximately the same parameters θ⋆ . We shall see that the gradient of the energy EEM−MW22 at Xε is approximately zero, illustrating the vanishing gradient phenomenon.
Figure 18 : Local minimum in Example A.2 with fixed non-uniform weights.
Figure 19 : Impact of the treatment of source weights in colour transfer.
Figure 20 : Comparison of balanced and unbalanced particle flows for non-uniform weights.
Figure 21 : Varying the number of components K , we study the convergence of EM, the local contractivity of F , and the MSEs of the OS and AI gradients.
Figure 23 : 3D barycentre (left) and its projections (right).
Figure 24 : Final GMMs in RGB space with (a) variable weights and (b) fixed weights. Target is in blue and optimised mixture is in red. In (a) we are stuck in a local minimum, in (b) we converged.
Figure 25 : Colour transfer with (a) K=1 components and (b) K=10 components.
Figure 26 : Colour transfer: barycentric transfer with entropic OT, and our method.
Figure 27 : Four images and their corresponding log-likelihood vs. K plots, for VGG layers ℓ∈{1,⋯,3} .
School of Artificial Intelligence Jilin University No. 2699, Qianjin Street, Chaoyang District Changchun 130012, China · Zhongguancun Academy Daniufang 2nd Ring Road, Haidian District2026 Beijing 100094, China · School of Artificial IntelligenceMay Jilin University No. 2699, Qianjin Street, Chaoyang District14 Changchun 130012, China