Search arXiv⌕ Search

arXiv · 2609.39525

Mitigating Representation Gaps in Amortized Bayesian Inference with Auxiliary Supervision

Abstract

Casting Bayesian inference as a neural network optimization problem targeting an amortized posterior is attractive, as it extends to otherwise intractable statistical models and offers near instantaneous inference for new datasets after prepaying the training cost. Although theory guarantees faithfulness under ideal convergence, practical amortized inference still requires iterating over architectures and optimization choices and ultimately ``satisficing'' under finite simulation, compute, and time budgets. Even the best-performing solution may thus retain avoidable representation gaps that typically require problem-specific fixes. Here, we propose a generic alternative which improves training dynamics with auxiliary guidance losses applied to internal representations. Specifically, we show how such guidance leads to faster convergence when training data is abundant and to better performance when it is scarce. We formalize representation gaps as getting stuck in a local optimum at the information bottleneck between the parts of the network tasked with feature learning and those tasked with conditional distribution learning, and offer a generic diagnostic to separate summary failures from inference failures. Finally, we demonstrate that auxiliary supervision improves convergence speed and accuracy on a range of challenging real-world inference problems.

Explore related subjects

Keep this discovery

Explore connections, maps & timelines

BibTeXRIS

Hans Olischläger, Svenja Jedhoff, Šimon Kucharský, Aayush Mishra, Stefan T. Radev, Paul Bürkner. 2026-09-30. Mitigating Representation Gaps in Amortized Bayesian Inference with Auxiliary Supervision. https://arxiv.org/abs/2609.39525

Cite the original work for its findings. Save a collection to share your selection of sources.

KEEP EXPLORING

Related papers

Convergence of Statistical Estimators via Mutual Information Bounds

Recent advances in statistical learning theory have revealed profound connections between mutual information (MI) bounds, PAC-Bayesian theory, and Bayesian nonparametrics. This work introduces a mutual information bound for statistical models, and derives from it convergence rates for fractional posteriors, for their variational approximations, and for the maximum likelihood estimator. The observations are assumed independent but not identically distributed, and the model is not assumed well-specified, so that the bounds are oracle inequalities and cover regression with a fixed or a conditioned design; the independent and identically distributed, well-specified case is recovered by dropping an index. We illustrate the method on two applications. In the Gaussian sequence model the rate is minimax in both the radius of the Sobolev ball and the sample size, which only appears through its product with the temperature. In logistic regression, where the model is not conjugate and the variational approximation is computed by a stochastic gradient method, the bound applies to the output of the algorithm rather than to an idealized minimizer, and attains the parametric order in the Renyi risk with no logarithmic factor, the statistical and the optimization error being separated.

stat.ML↗

Mitigating Over-squashing without Rewiring: A Sheaf Effective Resistance Perspective

Graph Neural Networks (GNNs) often struggle to capture long-range dependencies due to over-squashing -- a phenomenon in which the repeated compression of node embeddings into finite-size messages causes representations to collapse. Over-squashing is most often diagnosed as a property of the graph topology, with effective resistance serving as a principled measure of the bottleneck. We provide a complementary view on the matter: building on cellular sheaves, we introduce sheaf effective resistance, a generalization of effective resistance that depends on the sheaf attached to the graph, and we prove that for flat vector bundles, the over-squashing sensitivity in the Jacobian sense is upper bounded by a quantity related to the sheaf effective resistance between the nodes. The bottleneck thus need not lie in the graph itself: it can be relocated, and reduced, by adjusting the sheaf. We instantiate this idea in FlatNSD, a simple message-passing variant of Neural Sheaf Diffusion, and show that it implicitly learns to modulate total sheaf effective resistance, performing well on benchmarks designed to stress over-squashing without altering the original graph topology.

stat.ML↗

Risk-Calibrated Proposal Transport for Finite-Particle Diffusion Steering

Inference-time steering combines pretrained diffusion experts or rewards without retraining by changing the dynamics that transport noise to data. Feynman-Kac correction compensates for proposal mismatch through importance-weighted sequential Monte Carlo (SMC), whose finite-particle behavior depends on the proposal. Variance-controlling guidance (VCG) improves that proposal by fitting a linear drift correction to minimize empirical log-weight-rate variance. Although its population optimum cannot worsen residual variance, finite-particle VCG can nearly eliminate its fitting residual while increasing residual risk on new states by orders of magnitude. The resulting update can degrade unweighted generation or accelerate particle collapse. We show that the centered Feynman-Kac rate is the normalized transport residual and that expected out-of-fit benefit is exactly population headroom minus coefficient-estimation penalty. Under regularity assumptions, a Wasserstein analysis bounds the unweighted proposal's terminal error using this residual. These results motivate Risk-Calibrated Proposal Transport (RCPT), which uses deletion leave-one-out residuals to calibrate the retained fraction of the VCG update, adding no model calls and only small linear-algebra overhead. Experiments on 2D checker distributions, scaffold decoration, molecular property optimization, and class-conditional CIFAR-10 generation demonstrate recovery from harmful fitted updates. Across molecular and image domains, RCPT mitigates harmful fitted updates and improves a broad range of terminal metrics relative to uncalibrated VCG.

stat.ML↗