Source-linked AI summary

Mode Coverage in Normalizing Flow Boltzmann Generators via Log-Ratio Variation

Qi Feng, Rongjie Lai, Di Qi, Xuda Ye

arXiv:2609.09473v1stat.MLcs.LG

TL;DR

Forward-KL Boltzmann-generator training can miss target mass while showing high ESS because its target samples may omit modes. The paper introduces KLXX, combining forward KL with target- and QT–pushforward-weighted log-ratio variations, and reports improved mode coverage plus theoretically controlled staged inference.

  • Problem

    Forward-KL training depends on target surrogates that may omit modes, allowing incomplete coverage, surrogate error, and weak constraints on target-poor pushforward regions despite high ESS.

  • Method

    KLXX augments forward KL with target-weighted and QT–pushforward mixture log-ratio variations, then applies it in an adaptive-staging Boltzmann generator with importance reweighting.

  • Results

    KLXX improves mode coverage over forward KL, improves per-stage diagnostics against the schedule-building loss, and yields asymptotically unbiased inference under essentially bounded stage weights.

  • Takeaways & Limitations

    Log-ratio variations provide target, candidate-mode, and pushforward-support information that sampled forward KL usually omits.

  • Takeaways & Limitations

    The staged inference guarantee requires each stage weight to be essentially bounded, and the flow’s connected support can retain mass along paths between disconnected target basins.

Abstract

from arXiv · show

Normalizing flow Boltzmann generators retain a tractable pushforward density, but training with forward KL depends on target samples that may be biased or omit modes. As a result, a flow can miss target mass while its observed importance weights give a high effective sample size. We introduce the log-ratio variation $\X_ω$, the mean absolute pairwise difference of the target-to-pushforward log-density ratio under a weighting measure $ω$, and use it to define KLXX, a new loss function. Two log-ratio variations are added to the forward KL (denoted by the two X's): one weighted by the target to improve accuracy, the other by a mixture of quench and temper samples with pushforward samples to search candidate modes. We derive the Fisher--Rao gradient flow of KLXX, where both variations contribute nonpositive dissipation, and a fixed-surrogate error bound for KLXX. We use KLXX in an adaptive-staging Boltzmann generator, with importance reweighting at every stage. We bound the sampling error of its inference scheme when the stage weights are essentially bounded, and prove it asymptotically unbiased in the sample size. In the numerical tests, KLXX improves mode coverage over forward KL. It also improves the generator's per-stage diagnostics against the loss that built the schedule. The observables the generator recovers are close to independent references. The log-ratio variations thus supply information that the forward KL loss usually omits.

1 Introduction

Forward-KL training can miss target modes because its target surrogate may omit them, producing high ESS on incomplete support. KLXX adds log-ratio variations and QT-based candidate-mode information to improve coverage while retaining forward KL.

  • Problem: Forward KL can report high ESS even when the flow misses target modes, because ESS diagnoses only density agreement on reached support.The sampled loss also inherits bias from finite target surrogates and weakly constrains target-poor pushforward regions.
  • Contribution: QT melts source samples with Gaussian noise, quenches them into local basins, and tempers them to produce candidate-mode samples beyond the source support.This construction is used to search basins that ordinary source samples do not reach.
  • Contribution: KLXX retains forward KL and adds target-weighted and QT–pushforward mixture log-ratio variations to supply accuracy, candidate-mode, and support information.Pairwise log-ratio differences remove unknown normalizers while preserving the target as a minimizer.
  • Theory and inference: The paper derives nonpositive Fisher–Rao dissipation for both variations, a fixed-surrogate error bound, and an asymptotically unbiased staged inference scheme under bounded stage weights.The inference error is controlled at rate N^-1/2 for bounded observables.
  • Results: On Himmelblau and Sparse, both QT-based losses recover every mode, whereas forward KL, KL+Xπ, and FAB each miss modes.Mixing QT and pushforward samples also reduces visible intermodal leakage.
  • Application: An adaptive-staging Boltzmann generator uses KLXX with SMC and QT samples, importance reweighting, and per-stage ESS-based diagnostics.Training uses oracle access to U and ∇U and does not differentiate through MCMC transitions.

2 Construction of the KLXX loss

KLXX augments forward KL with two log-ratio variations whose weighting measures target reached regions, candidate QT modes, and pushforward support. The construction is compared with alternative loss and staging families while retaining a tractable flow density.

  • Loss construction: KLXX is forward KL augmented by two mean absolute pairwise log-ratio variations, with coefficients controlling their strengths and mixture proportions.Pairwise differences cancel unknown normalizers, and zero variation means the log-ratio is almost surely constant under its weighting measure.
  • Loss construction: The target variation reuses forward-KL target samples, while the mixture variation combines QT samples with detached pushforward samples.The detached pushforward distribution contributes support information without entering the differentiation graph.
  • QT construction: QT samples are melted beyond source support, quenched into potential basins, and tempered around basin centers rather than collapsed to point minima.Only melting provides the coverage expansion; quenching and tempering shape the resulting candidate-mode distribution.
  • Comparison to related works: Compared with FAB, KLXX retains forward KL and obtains coverage information from QT instead of replacing the loss with the α-divergence at α = 2.FAB targets importance-weight variance with AIS and a replay buffer, whereas KLXX uses QT and mixture variation evaluations.
  • Comparison to related works: Alternative families use annealed flow transport, ESS-gated staging, replay-buffer methods, diffusion or stochastic-control samplers, and autoregressive densities.These approaches differ in whether they alter the loss, train across annealing stages, or replace bijective flows.
  • Comparison to related works: Table 1 compares how samples enter each loss and does not assume equal computational budgets.Its forward-KL row uses this paper’s SMC target surrogate, while its FAB row describes the reference implementation.

3 Analysis of the generalized KLXX loss

The generalized KLXX analysis explains how QT-based mixture variation supports mode discovery, why pushforward samples reduce intermodal leakage, and how the loss behaves under Fisher–Rao flow and biased surrogates.

  • 3.1 Principle of mixture variation: QT melting reaches basins beyond the source, quenching maps samples to local minima, and tempering restores configuration-space spread.The resulting distribution covers target modes but assigns basin masses according to melted source mass rather than target probability.
  • 3.1 Principle of mixture variation: The mixture variation combines QT samples that constrain high log-ratios with detached pushforward samples that constrain low log-ratios, reducing intermodal leakage.QT-only variation can recover modes while leaving spurious mass between them because it does not inspect regions reached only by the pushforward.
  • 3.1 Principle of mixture variation: KLXX must retain the forward KL because log-ratio variations enforce pointwise flatness but do not by themselves drive the density globally toward the target.The forward KL is nonnegative and vanishes only at the target distribution.
  • 3.2 Fisher–Rao gradient of the KLXX loss: Under Fisher–Rao flow, both log-ratio variations contribute nonpositive dissipation, so adding them does not slow the forward-KL convergence bound in density space.The result concerns unconstrained density-space flow rather than parameter-training speed or optimizer behavior.
  • 3.3 Accuracy under biased target surrogate: With a fixed biased surrogate, both variations subtract nonnegative discrepancies from the surrogate error, improving the bound under target- and mixture-based weightings.The stationary-point result assumes a positive fixed surrogate and finite weighted importance expectations.
  • 3.3 Accuracy under biased target surrogate: The guaranteed improvement is capped at one and is mainly informative when the biased surrogate is already close to the target.The theorem applies to a fixed surrogate, whereas the algorithm rebuilds its surrogate at every gradient step.

4 Flow training with generalized KLXX loss

Generalized KLXX trains the flow using a target surrogate, a quench-and-temper sample set, and pushforward samples. Its two log-ratio variations add target accuracy and candidate-mode information to forward KL.

  • 4.1 KLXX training: KLXX combines forward KL with target and mixture log-ratio variations, using target-surrogate and quench-and-temper samples alongside pushforward samples.The target variation uses Xπ, while the mixture variation uses Xξ formed from QT and pushforward samples.
  • 4.1 KLXX training: Each gradient step pushes source samples through the flow, builds target-surrogate samples by SMC, forms the QT–pushforward mixture, and updates the flow.Only the log-ratio z(y) is differentiated during the gradient step.
  • 4.2 Constructing QT samples: QT melts samples with Gaussian noise, quenches them toward local minima using L-BFGS, and spreads them around modes with MALA.The resulting set supplies candidate modes beyond the source support, although its mode populations can differ from the target’s.
  • 4.3 Approximating π: SMC bridges pushforward ν to target π through M + 1 distributions, applying incremental weights w(y)^(1/M), resampling, and MALA rejuvenation.The ladder reduces weight degeneracy relative to direct importance sampling, while the π-invariant rejuvenation preserves the target limit.
  • 4.3 Approximating π: The SMC target-surrogate samples are asymptotically unbiased as sample size N grows under the stated regularity conditions.The resampled particle approximation converges to π, and π-invariant MALA preserves that target limit.

5 Adaptive-staging Boltzmann generator

The adaptive-staging Boltzmann generator interpolates between source and target distributions, selecting stage increments using training diagnostics and ESS gates. QT supplies candidate modes, while reweighting and rejuvenation update the carried sample set.

  • 5.1 Why staging is adaptive: Adaptive staging divides a difficult source-to-target transport into increments, because one flow may lack capacity in high dimensions or cannot split connected support.The schedule balances fit difficulty against the computational cost of adding stages.
  • 5.1 Why staging is adaptive: The interpolation uses Ut(x) = (1 − t)U0(x) + tU(x), with stage points selected during the run rather than fixed in advance.Later proposals extrapolate the previous accepted increment and are reduced when the ESS test fails.
  • 5.2 Adaptive-staging algorithm: Each candidate stage builds QT samples, trains KLXX from πk−1 to πk, compares the trained flow with the identity proposal, and selects the higher-ESS option.The stage point is accepted or shrunk using the ESS gate, and the selected weighted samples are rejuvenated under the interpolating target.
  • 5.2 Adaptive-staging algorithm: The stage ESS gate uses threshold τv = 0.4, shrinks rejected increments by γ = 0.7, and enlarges later proposals by Γ = 1.5.The first stage starts at tsafe = 0.2, while subsequent proposals extrapolate the last accepted increment.
  • 5.4 Propagation factor analysis: The inference error is bounded by C/N for bounded observables and vanishes as N grows, so the scheme is asymptotically unbiased under the stated stage-weight condition.The propagation factor is assembled from per-stage ESS values and grows with additional stages, making comparisons meaningful only on a shared schedule.

6 Numerical experiments

Across two-dimensional, high-dimensional, staged, and molecular tests, KLXX improves mode coverage and sampling diagnostics over forward-KL-based alternatives. Its target variation improves accuracy on reached modes, while QT and pushforward samples supply missing-mode and intermodal-support information.

  • 6.1 2D benchmark distributions: Both QT-based losses recover all three Three-Well modes, while forward KL and KL+Xπ leave the third well uncovered.Their ESS values are within 0.03 of each other.
  • 6.1 2D benchmark distributions: On Himmelblau, KLXX attains coverage 1 with less intermodal mass and higher ESS than KL+Xπ+Xˆπ, whereas forward KL and KL+Xπ capture only two wells at ESS above 0.97.The two losses using QT recover every well, but KLXX better suppresses mass between modes.
  • 6.1 2D benchmark distributions: On Sparse, KLXX attains coverage 1 with less intermodal mass and higher ESS than KL+Xπ+Xˆπ, while forward KL covers only the central pair.KL+Xπ settles on a single far mode; adding QT restores all four modes but leaves visible intermodal mass unless pushforward samples are included.
  • 6.1 2D benchmark distributions: KLXX reaches every mode on all four 2D targets and has the highest ESS among mode-complete methods on each target.The target variation improves accuracy where samples already exist, QT supplies missing modes, and pushforward samples reduce intermodal mass.
  • 6.3 Coefficient of the target variation at d = 100: The target-variation coefficient has a broad useful range: every tested λ from 0.5 to 5 exceeds ESS 0.81, while KLXX remains above the entire sweep.Forward KL ends at ESS 0.673, and the single-variation curve peaks at λ = 2.
  • 6.5 Adaptive staging and molecular tests: KLXX improves staged-generator diagnostics, with higher sample ESS on 20 of 24 shared stages, geometric-mean stage ESS gains of 0.04–0.06, and lower propagation factors.It reaches t = 1 in one fewer stage on NMA and glycerol, and has the smaller propagation factor on all three molecular targets.
  • 6.6 Boltzmann generator on achiral molecules: On molecular dihedral marginals, both losses populate every reference well, while KLXX halves glycerol’s excess trans mass relative to KL+Xπ.The largest visible departure occurs for glycerol; both methods also match the neutral-diethanolamine trans fraction.

7 Conclusions

KLXX addresses mode collapse by combining target, QT, and pushforward log-ratio variations while retaining forward KL. Across benchmarks it improves mode coverage and stage diagnostics, but most experiments use single training realizations.

  • 7 Conclusions: KLXX combines target and mixture log-ratio variations to expose candidate modes and pushforward-supported regions while retaining the exact target as a minimizer.The mixture uses QT samples for candidate-mode discovery and detached pushforward samples for regions occupied by the flow.
  • 7 Conclusions: Both QT-based losses recover every mode on the tested model problems, while adding pushforward samples reduces visible intermodal mass on Himmelblau and Sparse.Forward KL and KL+Xπ miss modes despite high ESS on several targets.
  • 7 Conclusions: KLXX has higher stage ESS than the schedule-building loss on all but the final shared stage and a smaller propagation factor at every reported batch size.It also reaches the final molecular stage with the smaller propagation factor on all three achiral molecules.
  • 7 Conclusions: The ϕ4, clock, and molecular experiments compare recovered observables against PT or exact references.These comparisons provide the paper’s accuracy evidence beyond ESS and coverage.
  • 7 Conclusions: Run-to-run uncertainty is estimated only for ϕ4, where three seeds produce final-ESS spread up to 0.031.ESS orderings elsewhere are therefore described as indicative rather than measured differences.

A Proof of the Fisher–Rao flow and its dissipation

The appendix derives the Fisher–Rao gradient flow for generalized KLXX and shows that both log-ratio variations contribute nonpositive dissipation. The resulting forward-KL bound preserves probability mass and gives exponential decay without standard geometric assumptions.

  • A Proof of the Fisher–Rao flow and its dissipation: The Fisher–Rao gradient of each log-ratio variation is expressed through a nonlocal correlation under its fixed weighting distribution.The flow at each point depends on the entire importance-weight profile under the target or mixture weighting measure.
  • A Proof of the Fisher–Rao flow and its dissipation: Antisymmetry of the correlation terms preserves total probability mass along the Fisher–Rao flow.The mixture is a probability density when α + β = 1, allowing its terms to be interpreted as expectations.
  • A Proof of the Fisher–Rao flow and its dissipation: The variation is convex in log ν but loses differentiability at weighting sets where equal log-ratios have positive pairwise mass.At ν = π, all importance-weight pairs tie and the derivation uses a least-norm subgradient convention.
  • A Proof of the Fisher–Rao flow and its dissipation: Each log-ratio variation contributes a nonpositive term to the derivative of KL(π∥ν_t), vanishing only when the importance weight is almost everywhere constant under its weighting distribution.The sign follows from the pairwise antisymmetric identity for the variation terms.
  • A Proof of the Fisher–Rao flow and its dissipation: KL(π∥ν_t) ≤ e^-t KL(π∥ν_0) without log-concavity, a spectral gap, or a growth condition on π.The variations change the flow structure rather than this forward-KL decay rate.

B Proof of accuracy under biased target surrogate

The biased-surrogate analysis characterizes stationary densities of KLXX and bounds their error relative to the target. The bound subtracts nonnegative discrepancy terms, subject to positivity and integrability conditions for a fixed surrogate.

  • B Proof of accuracy under biased target surrogate: The stationary density ν⋆ satisfies an identity obtained by combining the biased surrogate with the target and mixture variation correlations.The proof introduces the positive ratio ϱ = ν⋆/˜π and derives the stationary relation in that variable.
  • B Proof of accuracy under biased target surrogate: The biased-surrogate KL error is related to the stationary error by decomposing log(π/˜π) into log(π/ν⋆) and log(ν⋆/˜π).This decomposition connects surrogate bias to the stationary-density accuracy bound.
  • B Proof of accuracy under biased target surrogate: The fixed-surrogate error bound subtracts one nonnegative discrepancy for each log-ratio variation from the stationary error bound.The discrepancies sum to 1 − E_˜π[w⋆], and the subtraction is strictly below one when E_˜π[w⋆] > 0.
  • B Proof of accuracy under biased target surrogate: The theorem assumes a positive stationary solution and finite importance-weight expectations under the surrogate and mixture.A sufficient domination condition is conservative and excludes the coefficients used in the KLXX runs.

C Proof of the propagation factor analysis

The propagation-factor proof analyzes inference after the adaptive schedule and flow maps have been fixed. It proceeds stage by stage using importance weights and only requires invariance of the rejuvenation kernel.

  • C Proof of the propagation factor analysis: The stagewise error analysis uses identities linking consecutive target, pushforward, and weight triples, with kernel invariance as its only rejuvenation assumption.This establishes the propagation-factor analysis for the inference scheme with a general Markov rejuvenation kernel.
  • C Proof of the propagation factor analysis: The theorem concerns inference on a fresh particle set with a fixed accepted schedule and trained maps, excluding further training or selection.The proof tracks the consecutive-stage triple (π_k, ν_k, w_k).

C.1 Notations and what one stage does

The section defines the stage-wise measures, importance weights, particles, and Markov-kernel assumptions used to analyze the adaptive sampling scheme. It also explains how deterministic flow pushforwards, reweighting, resampling, and rejuvenation form each stage and propagate observable errors.

  • C.1 Notations and what one stage does: Each flow map is a diffeomorphism, so the pushforward ν_k has a density and the positive weight w_k = π_k/ν_k makes ν_k and π_k equivalent.The weight’s essential supremum is taken under ν_k.
  • C.1 Notations and what one stage does: At stage k, ν_k^N is the empirical measure of reweighted pushforward particles, while π_k^N is the empirical measure after resampling and rejuvenation.The returned estimate is π_K^N, and the analysis bounds its distance from the exact distribution.
  • C.1 Notations and what one stage does: Importance weighting distorts the measure according to previously accumulated error, while resampling contributes Monte Carlo error of order N^-1/2.Rejuvenation and resampling together produce one independent draw per particle conditional on the pushforward particles.
  • C.1 Notations and what one stage does: Because the pushforward matches null sets, errors transfer between corresponding stage measures with the same observable oscillation.This identity underlies the exact error decomposition used later.
  • C.1 Notations and what one stage does: The Markov kernel Q_k leaves π_k invariant and contracts observable oscillation, so rejuvenation preserves stage averages without increasing oscillation.For π_k-almost every x, ess inf h ≤ Q_kh(x) ≤ ess sup h.

C.2 Proof of Theorem 2

The proof establishes a stage-wise L2 error recurrence by separating fresh resampling error from inherited pushforward error. It then iterates the recurrence to obtain the final inference error bound under essentially bounded importance weights.

  • C.2 Proof of Theorem 2: The recurrence bounds stage-k error by a fresh sampling term plus 3∥w_k∥∞ times the preceding error, with the inherited term absent at k = 0.The normalization is almost surely finite and nonzero because weighted empirical averages lie in (0, ∥w_k∥∞].
  • C.2 Proof of Theorem 2: The importance-sampling lemma converts weighted ν_k-averages into π_k-averages and controls the normalized empirical ratio through centered weighted deviations.The centering function g_k has zero ν_k-average and oscillation bounded by 2∥w_k∥∞osc(h).
  • C.2 Proof of Theorem 2: The stage-k error decomposes exactly into Monte Carlo error from resampling and rejuvenation plus inherited error from previous stages.Minkowski’s inequality bounds the total L2 error by the sum of these two terms.
  • C.2 Proof of Theorem 2: Kernel invariance and oscillation contraction transfer the previous-stage error to the pushforward particles without increasing its observable dependence.Applying the lemma to Q_kh yields the inherited-error bound.
  • C.2 Proof of Theorem 2: Induction unrolls the recurrence into a weighted sum of stage-wise sampling errors, and evaluating the final stage yields the theorem’s inference bound.The final step applies the bound at k = K to the target observable φ.
Loading 2609.09473v1…