让EM算法可微分,实现高阶生成任务的端到端训练
Differentiable Expectation-Maximisation and Applications to Gaussian Mixture Model Optimal Transport
- 提出多种EM算法可微化方法,支持梯度反向传播
- 首次将混合高斯模型的Wasserstein距离作为可微损失使用
- 适用于图像生成、风格迁移等需要精确分布对齐的任务
期望最大化(EM)算法是统计学和机器学习中处理潜变量模型的核心工具,广泛应用于高斯混合模型(GMMs)。尽管应用广泛,传统EM通常被视为不可微的黑箱,难以融入现代端到端学习框架。本文系统研究并对比了多种EM的可微化策略,涵盖完全自动微分与近似方法,评估其精度与计算效率。作为关键应用,我们利用可微分的EM计算两个GMM之间的混合Wasserstein距离(MW₂),使该距离可作为可微损失函数用于图像生成与机器学习任务。为支持实际应用,我们提供了关于使用EM时MW₂稳定性的新理论证明,并提出一种新的非平衡型MW₂变体。数值实验在重心计算、色彩与风格迁移、图像生成及纹理合成等任务中验证了该方法的通用性。
原文摘要 · Abstract (English)
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 $\mathrm{MW}_2$ between GMMs, allowing $\mathrm{MW}_2$ to be used as a differentiable loss in imaging and machine learning tasks. To complement our practical use of $\mathrm{MW}_2$, we contribute a novel stability result which provides theoretical justification for the use of $\mathrm{MW}_2$ with EM, and also introduce a novel unbalanced variant of $\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.
Thank you to arXiv for use of its open access interoperability. PaperDance 不是 arXiv 官方产品;中文卡片由大模型生成,请以原文为准。