Approximate message passing on factor graphs underlies two dominant families of probabilistic inference algorithms: expectation propagation (EP) and variational message passing (VMP). Both methods approximate the marginal at each factor edge, forcing an iterative round-robin schedule, risking negative-precision messages, and, for VMP, collapsing to point estimates at Dirac-delta factors. We introduce Direct Message Approximation (DMA), which approximates factor-to-variable messages directly rather than the marginal. For normalisable factors, we define a consistency condition (requiring exactness when all other incoming messages are Dirac deltas) to guide message construction. We prove a master theorem (proper messages, any graph) bounding marginal KL from message KL, with three structural corollaries: Dirac-input consistency, no EP-style inner-loop iteration, and no negative-precision messages. Further, we prove a complementary O(1/r2) guarantee for the inherently improper backward message of the product factor, whose closed-form treatment has resisted prior work. As a concrete instantiation, we derive explicit DMA messages for the product and leaky-ReLU factors and assemble a Bayesian neural network (BNN) inference algorithm with one forward/backward sweep per training example and no gradient learning-rate hyperparameter, validating that the structural guarantees translate to predictive uncertainty that widens in data-sparse regions, including under model mismatch.
Figures & tables
Figure 1: Factor graph of a BNN (schematic; two learnable layers shown). Circles are variable nodes; filled squares are factor nodes. Blue-tinted circles represent the dl×dl−1 independent scalar weight variables Wij(l) (shown collectively per layer for clarity), each with a Gaussian prior factor π . Each × node is a matrix-vector product factor : shorthand for dl×dl−1 1D product factors δ(zij−Wij(l)xj(l−1)) (Sec. 3.1 ) combined with dl sum factors zi(l)=∑jzij(l) , together implementing the inner product (see Sec. 4.1 ). R factors are element-wise leaky-ReLU factors (Sec. 3.3 ); ℓ is the Gaussian likelihood factor (conjugate; δ=0 ); the double circle is the observed output y .
Figure 2: Left : Posterior predictive mean (solid) and ±2σ intervals (shaded) versus the true data-generating function (dashed), where σ=Var[f(x)]+β2 combines posterior variance with observation noise ( β=0.2 ). Training points shown as crosses ( N=200 , x∈[−2.5,1.5] ); dashed verticals mark the training boundaries. Right : Hinton diagram of posterior weight beliefs after training. Square area ∝ posterior mean magnitude; transparency ∝ posterior variance (opaque = certain).
Appendix figures & tables25 assets
Supplementary material from the paper’s appendix.
Appendix
ADF
EP
VMP
DMA (ours)
Approximates
marginal
marginal
marginal
message
Backward weight updates
×
✓
✓
✓
No cavity division
✓
×
×
✓
No inner-loop iteration
✓
×
×
✓
Consistency axiom
×
×
×
✓
Edge-local KL bound
×
×
×
✓
Appendix
Table 1: Comparison of approximate inference frameworks on properties relevant to BNN inference. Below the first row, ✓ indicates a desirable property. Approximates : whether the method targets the factor-to-variable message directly or approximates the marginal at each edge. No cavity division : whether updates avoid dividing out a stored incoming message (eq. 20 ). No inner-loop iteration : whether a single forward/backward sweep per example suffices, with no per-example fixed-point iteration. Backward weight updates : whether weight beliefs are updated via a backward message sweep. Consistency axiom / edge-local KL bound : see Definition 2.1 and Theorem 2.3 .
Factor
Direction
Proper?
Theorem
Gaussian prior
both
Yes
Cor. 2.4 (exact)
Gaussian likelihood
both
Yes
exact (conjugate)
Linear / copy
both
Yes
Cor. 2.4
Product ( z=xy )
fwd ( z )
Yes
Thm. 2.3
Product ( z=xy )
bwd ( x,y )
No
Thm. 3.3
Leaky-ReLU ( α>0 )
fwd
Yes
Thm. 2.3
Appendix
Table 2: Theorem coverage for each factor direction in the BNN. “Proper?” refers to whether the DMA message is a valid (positive-precision) Gaussian. All linear and copy factors satisfy Corollary 2.4 exactly; the Gaussian likelihood is conjugate and exact. See Remarks B.2 and B.6 for the improper cases.
Figure 3: Empirical validation of Theorem 2.3 across 192 leaky-ReLU factor configurations. Left : log-log scatter of message error δ versus marginal error; the dashed line Cδ (with C=∥mXj→f∥∞/Z ) lies above every point, confirming O(δ) scaling. Right : the marginal KL normalised by its per-configuration bound ∥mXj→f∥∞/Z⋅δ ; all 192 ratios lie below the bound line at 1 (maximum ratio 0.97 ), directly certifying the theorem. The bound is tight: the ∥mXj→f∥∞/Z prefactor evaluates the Gaussian sup-norm, which is largest where the incoming message is most concentrated and the marginal error is therefore small anyway.
Figure 4: Actual marginal KL (solid, coloured) versus the theoretical bound ∥mXj→f∥∞/Z⋅δ (dashed, black) as each parameter is swept individually (legend shown in panel (a)). (a) Input SNR r ( α=0.3 , σx=1 fixed): the bound decays monotonically; the actual KL is non-monotone, peaking near r≈2 – 3 where the ReLU kink is hardest to match, and small at both low and high r . (b) Leaky slope α ( r=2 , σx=1 fixed): both decrease as α→1 (linear factor, δ=0 ) and grow as α→0 (approaching the improper hard-ReLU limit). (c) Input width σx ( r=2 , α=0.3 fixed): the actual KL increases with σx (wider incoming message exposes more backward-message error in the marginal), while the bound decreases through its 1/σx prefactor; the conservatism is greatest at small σx . In all panels the bound lies strictly above the actual KL, with a maximum ratio of 0.68 across all three sweeps (width sweep, panel (c)).
Figure 5: KL divergence between IS reference marginal and DMA marginal as a function of input SNR r . Left : product factor backward message to X from Z=XY ; incoming messages Y∼N(μy,(μy/r)2) and Z∼N(μz,(μz/r)2) with μy=4 , μz=10 , and the incoming message from X fixed at N(3,1) . Right : ReLU backward message to X from Y=relu(X;α) with α=0.1 ; incoming message X∼N(μx,(μx/r)2) with μx=1 . Both panels show empirical O(1/r2) decay. The product backward (left) is covered by Theorem 3.3 ; the leaky-ReLU backward (right, α=0.1 , proper message) is consistent with Theorem 2.3 applied to the message KL.
Figure 6: Copy-factor KL divergence over the 2D parameter space, computed from the closed-form expression in Lemma B.3 (no Monte Carlo). Left : Step 1 error KL[LN(⋅;μY,σY2),m^f→X(⋅)] swept over (μY,σY) . Horizontal bands confirm that the KL depends only on the log-space variance σY2 , not on μY ; the analytical rate is 43σY2+O(σY4) . Right : Step 2 error KL[LNm(⋅;μX,σX2),N(⋅;μX,σX2)] swept over (μX,σX) . Dashed white lines are level curves of r=μX/σX ; diagonal banding confirms the KL depends only on r , consistent with the O(1/r2) rate.
Figure 7: Product factor f(X,Y,Z)=δ(Z−XY) . Each row shows marginals X , Y , Z (left to right). Blue: DMA. Gray: IS. (a) Nominal ( μX = 3, σX2 = 1, μY = 4, σY2 = 1, μZ = 10, σZ2 = 5): the Z forward marginal is well approximated by the DMA Gaussian; the X and Y backward marginals are already noticeably skewed in the IS reference. (b) Stress ( μX = 1, σX2 = 4, μY = -2.5, σY2 = 1, μZ = -5, σZ2 = 1): Z remains well approximated; the X and Y backward marginals become severely non-Gaussian.
Figure 9: Left : Training negative log-likelihood (NLL =−N1∑ilogp(yi∣xi) , nats/example) versus epoch for DMA (solid black) and Adam at four learning rates (dashed). DMA’s outer updates stop automatically at epoch 17; Adam’s speed and final NLL depend critically on η . Values above 8 are clipped to 8 for display. Right : Adam ( η=0.1 ) point prediction outside the training range [−2.5,1.5] (dotted verticals); the true function (dashed) can deviate arbitrarily from the point estimate, with no uncertainty quantification available. Compare with the widening DMA posterior in Figure 2 . DMA’s per-epoch cost is 0.80ms versus 0.45ms for Adam on this architecture ( 1.8× ), but DMA’s early stopping at epoch 17 versus Adam’s ∼80 gives an overall ∼5× reduction in total training time.
Figure 10: Left : Training NLL vs. epoch for DMA (solid black), Adam η=0.1 (dashed grey), AdamW η=0.1 at four weight decay values (solid coloured), and two poorly-tuned AdamW configurations ( η=0.01 and η=1.0 , both λ=1.0 , dotted) illustrating hyperparameter sensitivity. Right : AdamW η=0.1 , λ=0.1 (best weight decay on this seed) point prediction outside the training range [−2.5,1.5] (dotted verticals); the true function (dashed) is tracked well on this seed but no uncertainty quantification is available.
Method
Epochs
Extrapolation NLL
Median
IQR
DMA (posterior predictive)
16
00.80
1.95
Adam η=0.1
200
12.09
11.52
AdamW η=0.1 , λ=0.01
200
08.05
10.73
AdamW η=0.1 , λ=0.1
200
05.52
09.86
AdamW η=0.1 , λ=1.0
200
04.07
10.47
Appendix
Table 3: Extrapolation NLL ( ↓ better) on the 83-weight 1D regression task, averaged over 20 seeds (median [IQR]). Bayes-optimal ≈−0.69 nats.
Figure 11: Hyperparameter sweep: minimum training NLL achieved over 500 epochs for each of 160 IVON configurations ( 8×4×5 grid over η , β2 , δ ) on the 83-weight BNN ( N=200 , β=0.2 ). Only 26 of 160 configurations reach a final NLL below 1.0, concentrated in the δ=0.1 and δ=0.5 columns. DMA training NLL at convergence: −0.16 (epoch 17, no hyperparameter search); best IVON: −0.31 (epoch 500, requires δ=0.1 and extensive tuning).
Figure 12: Training NLL minus Bayes-optimal (log scale) vs. epoch for IVON at three learning rates (best β2 / δ per η from the sweep) and DMA, on the 83-weight BNN ( N=200 , β=0.2 ). DMA converges at epoch 17; η=0.001 descends slowly but does not converge within 500 epochs; η=0.01 converges to NLL ≈0.26 with δ=0.5 ; η=0.1 oscillates near the Bayes-optimal line.
Figure 13: Left : DMA posterior predictive (identical to Figure 2 ). Right : IVON predictive with the best sweep configuration ( η=0.6 , β2=0.9999 , δ=0.1 ) after 2000 epochs. Training range [−2.5,1.5] marked by dotted verticals. DMA extrapolation NLL: 0.54 ; IVON on this seed: −0.40 . Across 20 seeds IVON diverges on 14/20; median extrap NLL on finite seeds is 2.11 (IQR 5.27 ) vs. DMA median 0.79 (IQR 1.28 ).
Figure 14: Left : DMA posterior predictive (same as Figure 2 ). Right : Diagonal Laplace predictive with He-init prior (MAP via Adam with per-layer L2 decay, diagonal of exact Hessian, K=1000 MC samples). Training range [−2.5,1.5] marked by dotted verticals.
Figure 15: Extrapolation calibration curves for the 83-weight network, test points in [−4,−2.5]∪[1.5,3] . Δ=coverage−α : 0 = perfect, >0 = conservative, <0 = overconfident. On this seed, DMA (blue, Δ=+0.01 ) and diagonal Laplace (red, Δ=+0.05 ); Δ is the mean signed deviation — DMA’s smaller ∣Δ∣ reflects cancellation of over- and under-coverage rather than a uniformly tighter fit to the diagonal. Over 20 seeds both methods are comparable (median −0.09 vs. −0.10 ).
Method
Extrap NLL
Δ (extrap)
DMA
0.82[1.96]
−0.09[0.44]
Diagonal Laplace
0.77[1.91]
−0.10[0.35]
Bayes-optimal
−0.69
0
Appendix
Table 4: Extrapolation NLL and calibration error Δ on the 83-weight 1D regression task, over 20 independently drawn datasets and true functions (median [IQR]). Δ>0 : conservative; Δ<0 : overconfident; Δ=0 : perfect.
Figure 16: Model mismatch experiment ( N=200 , training range [−5,5] , β=0.2 ). The data-generating network has two wider hidden layers ( d=12,10 ); both methods learn with the two-hidden-layer model network ( d=6,5 ). Left : DMA posterior predictive mean (solid) and ±2σ intervals (shaded) versus the true function (dashed). Training points are shown as crosses; dotted vertical lines mark the training boundaries x=±5 . Right : Adam ( η=0.1 , red) and AdamW ( η=0.1 , λ=0.1 , purple) point predictions; neither carries uncertainty quantification.
Method
Median
Mean
IQR
DMA (posterior predictive)
−0.01
−0.29
0.44
Adam ( η=0.1 )
−0.50
−2.21
2.23
AdamW ( η=0.1 , λ=0.1 )
−0.36
−0.11
0.73
Bayes-optimal
−0.69
–
–
Appendix
Table 5: Extrapolation NLL over 20 independent seeds ( ↓ better; IQR = interquartile range). Bayes-optimal ≈−0.69 nats. AdamW uses η=0.1 , λ=0.1 (best of λ∈{0.01,0.1,1.0,10.0} by median).
Architecture
Weights
N
Out
Epochs
NLL
Per-epoch
8→6→5→1
83
200
1
17
−0.16
0.8ms
6→6→12→48→24→4
1932
1500
4
3
−0.73
67ms
Appendix
Table 6: DMA training summary: small baseline vs. large network. NLL is the per-example per-output training log-likelihood NK1∑i,klogp(yik∣xi) at the final epoch (lower / more negative is better). Per-epoch time excludes the first epoch (JIT warm-up).
Figure 17: DMA posterior predictive ( 6→6→12→48→24→4 , 1932 weights, N=1500 , 3 epochs, 0.72s total). Each panel shows one of the four output channels: posterior predictive mean (solid) with ±2σ bands against the true function (dashed); training points as faint dots.
Figure 18: Adam ( η=0.01 , 100 epochs, 1.7s total) point predictions with fixed ±2β bands on the same four output channels. The confidence interval is constant across the input range because a point estimate carries no epistemic uncertainty.
Figure 19: Hinton diagram of DMA posterior weight means after 3 epochs. Each block corresponds to one weight matrix; square size encodes ∣μw∣ , colour encodes sign.
Figure 20: Training NLL (per example per output) vs. epoch. DMA (solid black) converges at epoch 3 with no learning-rate tuning. Adam with η=0.01 reaches NLL −0.79 after 100 epochs; η=0.001 has not converged (NLL +0.84 ); η=1.0 diverges. Values above 8 are clipped to 8 for display. AdamW η=0.1 , λ=0.1 (solid purple) reaches NLL −0.47 : weight decay regularises the trajectory but the chosen η does not reach the same final NLL as the best-tuned Adam.
Figure 21: Extrapolation calibration curve for DMA on the 1932-weight network ( N=1500 ), averaged across all four output channels. Test points in [−6,−4]∪[4,6] . Δ=+0.026 (slightly conservative); the 83-weight result has a 20-seed median Δ=−0.09 (Appendix E.7 ).
Stacking probabilistic building blocks into deeper architectures typically breaks closed-form inference. We show that closed-form inference can be preserved. We identify five factor-graph primitives: a bilinear factor, an exponential link, a Gamma prior, a Gaussian likelihood, and an equality node, and prove that any model composed from them admits closed-form variational message passing. The construction works because each primitive preserves a small set of message families: under mean-field factorization, messages on Gaussian variables remain Gaussian and messages on precision variables remain Gamma, while the only non-conjugate interface, the exponential link, remains tractable through the Gaussian moment-generating function and the sufficient statistics of the Gamma family. We demonstrate composition at increasing depth, from static ensembles through input-dependent gating to split-branch routing, and show that stacking routing layers encodes arbitrary decision trees, establishing universal function approximation with closed-form inference. Applied to ensemble time-series forecasting, the framework yields a Bayesian mixture of experts in which gating functions are inferred rather than learned, providing calibrated uncertainty over expert selection across five benchmark datasets.
Mykola Lukashchuk, Kyrylo Yemets, Wouter M. Kouw +4
Eindhoven University of Technology, the Netherlands · Lviv Polytechnic National University, Lviv, Ukraine · Lazy Dynamics, Utrecht, the Netherlands +1
Amortized inference promises fast test-time Bayesian inference, but existing methods are inherently tied to fixed models. Extending amortization to unseen models typically requires retraining or costly test-time finetuning. In this paper, we ask: is it possible to build a single inference network capable of generalizing across varying priors, likelihoods, and dimensionality? We introduce Amortized Factor Inference Networks (AFINs), a family of encode-merge-decode inference networks built on dimension-independent modules that map a model specification and its observations to the parameters of a variational posterior. Experimentally, a single trained AFIN achieves posterior accuracy comparable to NUTS and several variational inference methods, while requiring 2 to 4 orders of magnitude less test-time compute. Code is available at https://github.com/joohwanko/AFINs.
Joohwan Ko, Justin Domke
Manning College of Information and Computer Sciences University of Massachusetts Amherst
We present a new algorithm for amortized inference in sparse probabilistic graphical models (PGMs), which we call Δ-amortized inference (Δ-AI). Our approach is based on the observation that when the sampling of variables in a PGM is seen as a sequence of actions taken by an agent, sparsity of the PGM enables local credit assignment in the agent's policy learning objective. This yields a local constraint that can be turned into a local loss in the style of generative flow networks (GFlowNets) that enables off-policy training but avoids the need to instantiate all the random variables for each parameter update, thus speeding up training considerably. The Δ-AI objective matches the conditional distribution of a variable given its Markov blanket in a tractable learned sampler, which has the structure of a Bayesian network, with the same conditional distribution under the target PGM. As such, the trained sampler recovers marginals and conditional distributions of interest and enables inference of partial subsets of variables. We illustrate Δ-AI's effectiveness for sampling from synthetic PGMs and training latent variable models with sparse factor structure.
Jean-Pierre Falet, Hae Beom Lee, Esmeralda S. Whitammer +6
Mila – Qu´ebec AI Institute, Universit´e de Montr´eal · Mila – Québec AI Institute, Université de Montréal