Stability of Finite-Batch Particle Mean-Field Variational Inference Beyond Strong Convexity

arXiv:2608.11486 · math.NA, cs.NA, math.OC, math.PR, stat.ML · Submitted 2026-08-11 · Read on arXiv

Vinh Nguyen, Truong Vu

math.NA, cs.NA, math.OC, math.PR, stat.ML

Submitted: 2026-08-11

Updated: 2026-08-13

Comments: 39 pages, 8 figures

Code: https://github.com/tvu25/MFVI

License: http://creativecommons.org/licenses/by/4.0/

Importance score: 75/100

The gist: This paper studies the stability of finite-batch particle mean-field variational inference (MFVI) beyond strong convexity.

Terminology

Summary

This paper studies the stability of finite-batch particle mean-field variational inference (MFVI) beyond strong convexity. The authors analyze the implementable finite-batch particle algorithm for mean-field variational inference as a fully discrete stochastic approximation of the projected Wasserstein dynamics.

Problem setting and key assumptions. The target potential V is globally smooth (Assumption 2.1: V ∈ C2(Rm) with ‖∇2V(x)‖ op ≤ L for all x) but need not be strongly convex. The departure from contractivity is quantified by the curvature defect defined in Definition 2.2:

d α(x,y) = [α‖x−y‖2 − ⟨∇V(x) − ∇V(y), x−y⟩]+,

which is the additive loss in the one-step Euler contraction estimate. The uniform defect condition d α(x,y) ≤ β is equivalent to ⟨∇V(x) − ∇V(y), x−y⟩ ≥ α‖x−y‖2 − β.

Main result. The principal result is a non-asymptotic error estimate for the finite-batch scheme. Corollary 7.3 gives, up to universal numerical constants:

EW22(q Xn, q⋆) 1/2 ≲ (1 − cαh) n/2 EW22(q X0, q⋆) 1/2 + (2 + √(40κ/α)) ε N(q⋆) + √(12β/α) + √(8hΓ⋆/(αB)) + √(80L2h/α2)(hΓ⋆ + m),

where Γ⋆ = E q⋆‖∇V‖2 and κ measures the sensitivity of the projected drift to perturbations of the product law. The sharper estimate in Theorem 7.1 replaces β by a geometrically weighted sum of the defects encountered by the stationary coupling.

Key components of the proof. The proof uses a stationary comparison array whose population law is an MFVI minimizer but whose particle-level law is a random product empirical measure. The authors construct an array whose rows have the marginals of q⋆ and couple it to the PAVI array by predictable rowwise optimal matchings. The product of the stationary row empirical measures is random and is not equal to the deterministic product law q⋆, producing a nonvanishing discrepancy that yields the κε N(q⋆) term.

Variational foundations. Under the uniform defect condition, the authors prove quadratic coercivity of V (Lemma 3.1), existence of MFVI minimizers (Theorem 3.3), the implication from global minimizers to coordinatewise Gibbs equations (Proposition 3.6), centered moment bounds for stationary marginals without log-concavity (Lemma 3.9), and the stationary gradient moment bound Γ⋆ ≤ mL2(1+β)/α (Lemma 3.10).

Continuous-time analysis. The independent-projection McKean–Vlasov diffusion is shown to have unique strong solutions preserving product structure (Theorem 4.2). A synchronous coupling gives the continuous-time stability estimate W22(μ t, ν t) ≤ e−2αtW22(μ0, ν0) + 2∫0t e−2α(t−s)Δ(s)ds (Theorem 4.3). Corollary 4.4 shows that all MFVI minimizers lie within √(β/α) of one another in W2.

Coordinatewise defects and dimension dependence. The paper introduces coordinatewise defects d α,i and shows that the full-dimensional defect is bounded by their sum. After normalizing by m−1/2W2, the curvature residual is controlled by the average coordinate defect. Although the generic estimate for the projected-drift sensitivity is κ ≤ √(mL), sparse or bounded cross-coordinate interactions can give a dimension-independent bound via the Schur bound κ ≤ √(K1K∞).

Computational cost. Proposition 5.1 shows that a direct implementation uses O(mNB·C∂V) work and O(m(N+B)) storage per iteration, where C∂V is the cost of one coordinate derivative evaluation. Over a fixed physical time horizon T, this becomes O(mNBT·C∂V/h).

Nonconvex benchmark. The authors construct an arbitrary-dimensional smooth nonconvex interaction model V a,J(x) = Σi ν a(xi) + (1/2)s(x)TJs(x) with ν a(x) = x2/2 + a·cos(x), s(x) = (tanh x1,..., tanh x m)T, and J symmetric with zero diagonal and ‖J‖ op < 1. This benchmark has a closed-form unique MFVI minimizer q⋆ = π a⊗m (Proposition 8.2), permitting direct measurement of particle, step-size, batch, defect, and dimension effects without unknown optimization error.

Polynomial drift obstruction. The paper explains why polynomially growing drifts require modification of the untamed explicit scheme. For V(x) = x4/4 + x2/2, the one-particle Euler update has E[X n+12X n = x] growing like h2x6, so no global quadratic Foster–Lyapunov inequality can hold. Taming the stochastic batch estimator generally changes its conditional mean and may destroy the unbiasedness used in the orthogonality argument of Lemma 6.4.

Numerical validation. The experiments reproduce the predicted B−1/2 batch scaling (fitted slope −0.505), show nearly N−1/2 empirical scaling for the smooth benchmark (fitted slope −0.458, faster than the worst-case N−1/4 bound), and the coupled time-discretization study is approximately first order (fitted slopes 1.110 and 1.010). The dimension dependence is consistent with √m-type growth in the unnormalized metric, while the per-coordinate error remains nearly constant.

Improvements for AI systems

Improvement 1: Non-Convex Optimization with Provable Stability Guarantees

The improved AI system can optimize objectives that are smooth but not strongly convex (e.g., neural network losses, variational autoencoder objectives) with explicit finite-time error bounds. It can handle batch-based stochastic gradient updates while maintaining convergence to a neighborhood of the optimum, with the neighborhood size controlled by batch size, step size, and curvature defect. This enables reliable deployment in non-log-concave settings (e.g., Bayesian deep learning) where traditional convex analysis fails.

Improvement 2: Adaptive Batch-Size Selection for Mean-Field Variational Inference

The AI system can automatically tune batch size B in particle-based variational inference to balance computational cost (O(mNB·C∂V) per iteration) against the ε N(q⋆) and √(8hΓ⋆/(αB)) error terms. Given a target accuracy, it can compute the minimal B needed to achieve that accuracy, reducing memory usage (O(m(N+B)) storage) and wall-clock time in large-scale models (e.g., latent variable models with millions of parameters).

Improvement 3: Dimension-Aware Error Control in High-Dimensional Sampling

The system can detect when cross-coordinate interactions are sparse or bounded (via the Schur bound κ ≤ √(K1K∞)) and switch to dimension-independent error guarantees. It can then scale to high-dimensional problems (m > 104) without the √m penalty, enabling efficient posterior sampling in high-dimensional Bayesian models (e.g., spatial-temporal models, deep generative models) where naive methods suffer from curse of dimensionality.

Improvement 4: Robustness to Non-Strong-Convexity via Curvature Defect Monitoring

The AI system can monitor the curvature defect d α(x,y) during optimization and adaptively adjust step size h and batch size B to keep the bias term √(12β/α) below a user-specified threshold. This allows it to handle saddle points, flat regions, and non-convex landscapes (e.g., in reinforcement learning policy gradients) while maintaining a contraction rate (1−cαh) n/2 toward the optimal distribution.

Improvement 5: Tamed Stochastic Gradient Estimators for Polynomial Drift

For objectives with polynomially growing gradients (e.g., quartic potentials like x4/4 + x2/2), the system can implement a tamed version of the stochastic batch estimator that preserves unbiasedness (via the orthogonality argument in Lemma 6.4) while preventing variance explosion. This enables stable optimization in heavy-tailed or high-moment settings (e.g., robust statistics, financial modeling) where standard Euler updates diverge.

Improvement 6: Benchmark-Driven Hyperparameter Tuning for Nonconvex MFVI

Using the closed-form benchmark V a,J(x) with known MFVI minimizer q⋆ = π a⊗m, the AI system can perform automatic hyperparameter search (step size h, batch size B, particle count N) by directly measuring the error W22(q Xn, q⋆) without needing to solve the optimization problem first. This yields calibrated hyperparameters that transfer to similar nonconvex problems, reducing manual tuning effort in production systems.

Improvement 7: Early Stopping Criteria with Non-Asymptotic Guarantees

The system can compute a rigorous stopping time n* based on the explicit error bound in Corollary 7.3: it stops when the contraction term (1−cαh) n/2 times the initial error falls below the combined bias terms (ε N, β, h, batch effects). This provides a principled termination criterion for iterative variational inference, avoiding over-iteration waste and under-iteration inaccuracy.

Improvement 8: Coordinatewise Defect Decomposition for Sparse Models

For models with sparse or structured interactions (e.g., graphical models, factor analysis), the system can decompose the curvature defect into coordinatewise components d α,i and allocate computational resources (more particles or smaller step sizes) to coordinates with high individual defects. This yields faster convergence in settings where only a few coordinates are strongly nonconvex, improving efficiency over uniform treatment.

Improvement 9: Uncertainty Quantification in Non-Log-Concave Posteriors

The improved system can provide calibrated uncertainty estimates (via the W2 error bound) for Bayesian inference in non-log-concave models, even when the posterior is multimodal or heavy-tailed. It can output both point estimates and error bars that account for particle discretization, batch noise, and step-size bias—critical for decision-making in medical imaging, climate modeling, and autonomous systems.

Improvement 10: Memory-Efficient Streaming Inference for Large-Scale Data

Leveraging the O(m(N+B)) storage requirement, the system can process streaming data (e.g., online learning, federated learning) by maintaining a fixed-size particle ensemble and batch buffer, updating them sequentially. The non-asymptotic bounds guarantee that the error remains controlled even with non-stationary data streams, enabling real-time Bayesian updating in resource-constrained edge devices.

Abstract

We study the implementable finite-batch particle algorithm for mean-field variational inference as a fully discrete stochastic approximation of the projected Wasserstein dynamics. The target potential is globally smooth but need not be strongly convex. The departure from contractivity is quantified by the curvature defect [d alpha(x,y) = [alphax-y 2- grad V(x)-grad V(y),x-y]+,] which is the additive loss in the one-step Euler contraction estimate. We prove a non-asymptotic Wasserstein stability bound that separates initialization, product-empirical approximation, finite-batch drift error, time discretization, and the defects accumulated along the coupled trajectories. Under the uniform bound d alpha at most beta, the particle iterates remain within O(sqrt beta/alpha) of any MFVI minimizer, up to explicit errors in the particle number, batch size, and step size. The proof uses a stationary comparison array whose population law is an MFVI minimizer but whose particle-level law is a random product empirical measure, and it controls the resulting projected-drift discrepancy explicitly. We also give coordinatewise defect estimates and structural conditions for dimension-independent projected-drift sensitivity, construct an arbitrary-dimensional smooth nonconvex benchmark with a closed-form MFVI minimizer, and explain why polynomially growing drifts require a modification of the untamed explicit scheme.

Sources

Related papers