Geometric Convergence Analysis of Variational Inference via Bregman Divergences
Sushil Bohara, Amedeo Roberto Esposito
TL;DR
This work addresses convergence of variational inference under non-convex ELBO by developing a geometric framework based on the exponential-family Bregman divergence induced by the log-partition function $A(\phi)$. It proves that $L(\phi)=D_A(\phi^* \| \phi)$ with $\nabla L(\phi)=H(\phi)(\phi-\phi^*)$, enabling a monotonicity property and a ray-wise analysis that yields two-sided quadratic bounds via ray-wise spectral envelopes $\alpha(\phi)$ and $\beta(\phi)$. The authors derive non-asymptotic convergence rates for natural gradient descent (NGD) and Euclidean gradient descent (GD): NGD contracts along a fixed ray with rate $|1-\eta|^k$ (independent of conditioning), while GD's rate depends on the local condition number $\beta(\phi)/\alpha(\phi)$ and can be slower. Numerical experiments on Bernoulli and Gaussian VI validate the theory, showing NGD's robustness to conditioning and the advantage of the geometric approach for understanding and guiding VI optimization.
Abstract
Variational Inference (VI) provides a scalable framework for Bayesian inference by optimizing the Evidence Lower Bound (ELBO), but convergence analysis remains challenging due to the objective's non-convexity and non-smoothness in Euclidean space. We establish a novel theoretical framework for analyzing VI convergence by exploiting the exponential family structure of distributions. We express negative ELBO as a Bregman divergence with respect to the log-partition function, enabling a geometric analysis of the optimization landscape. We show that this Bregman representation admits a weak monotonicity property that, while weaker than convexity, provides sufficient structure for rigorous convergence analysis. By deriving bounds on the objective function along rays in parameter space, we establish properties governed by the spectral characteristics of the Fisher information matrix. Under this geometric framework, we prove non-asymptotic convergence rates for gradient descent algorithms with both constant and diminishing step sizes.
