Research
A Mean-Field Framework for Inference-Time Distributional Control of Diffusion Models
A Mean-Field Framework for Inference-Time Distributional Control of Diffusion Models Overview Research area: Machine learning / statistical machine learning (stat.ML), specifically inference-time stee
- arXiv
- 2608.08770
- Published
- 2026-08-09
- Authors
- Samuel Howard, Nikolas Nüsken
AI summary
A Mean-Field Framework for Inference-Time Distributional Control of Diffusion ModelsOverview
Research area: Machine learning / statistical machine learning (stat.ML), specifically inference-time steering of diffusion and flow-based generative models using interacting particle systems and mean-field (McKean–Vlasov) theory.
Technical level: Advanced. The paper is built around functional derivatives of measure-defined reward functionals, Feynman–Kac reweighting, McKean–Vlasov dynamics, and McKean–Vlasov convergence statements. The intuition is accessible, but the machinery assumes familiarity with stochastic processes and measure-theoretic optimisation.
Scope: The paper formalises distribution-level ("distributional") reward steering of diffusion models as targeting a tilted measure in a mean-field framework, derives a weighted interacting-particle scheme that provably targets it, and validates the scheme in low-dimensional tractable settings and in protein conformation tasks.
What This Paper Is About
Existing inference-time steering methods for diffusion and flow models mostly optimise a pointwise reward defined on individual samples, and recent work has placed these on a firm theoretical footing by combining reward-gradient steering with Feynman–Kac particle reweighting so that samples follow a prescribed tilted distribution. Many practical objectives are instead distributional — they are defined on the whole generated ensemble (calibration to population-level data, mode balancing, diversity) — and existing approaches to these are largely heuristic, with no clear characterisation of what distribution they actually sample. This paper fills that gap: it casts distributional steering as targeting a tilted measure under a mean-field framework and derives a weighted interacting-particle scheme that targets it in a principled manner, recovering the pointwise case as a special case.
Key Contributions
-
A formal target for distributional steering. The paper formulates distribution-level steering as sampling from a tilted distribution for a measure-defined reward $\mathcal{R}: \mathcal{P}(\mathbb{R}^d) \rightarrow \mathbb{R}$, where the tilt itself depends on the measure — a distributional analogue of the reward-tilted targets used with pointwise rewards. Proposition 3.1 shows the maximiser satisfies an implicit self-consistency relation $\mu^{}(\mathrm{d}x) = \frac{1}{Z}e^{\Psi(x,\mu^{})}p_1(x)\mathrm{d}x$, with $Z=\int_{\mathbb{R}^d} e^{\Psi(x,\mu^{*})}p_1(x)\mathrm{d}x$ and $\Psi$ the first variation (functional derivative) of $\mathcal{R}$.
-
Weighted McKean–Vlasov dynamics with corrective reweighting. Theorem 3.2 derives interacting (McKean–Vlasov) dynamics whose weighted marginals correctly track the desired tilted path $\mu_t^{*}$, where the log-weight dynamics replace the pointwise $\partial_t r(\cdot)$ term with $\dot{\Psi}_t(\cdot) = \frac{\mathrm{d}}{\mathrm{d}t}\Psi_t(\cdot,\mu_t)$ — a term whose time dependence arises through both $\Psi_t$ and the evolving measure. Proposition 3.3 gives an implicit characterisation of $\dot{\Psi}_t$ via the second variation $\Phi_t(x,y,\mu) = \frac{\delta^2 \mathcal{R}_t}{\delta\mu^2}(x,y,\mu)$ and its mean-centred version.
-
A finite-particle algorithm with a convergence guarantee. Algorithm 1 turns the idealised dynamics into a single-batch sampling procedure (a multi-batch version, Algorithm 2, is given in Appendix B), with two solvers for $\dot{\Psi}_t$ — a linear-system approach (Algorithm 3) and a lightweight Picard/fixed-point iteration (Algorithm 4). Theorem 3.4 states that the weighted empirical measure converges to $\mu_t^{}$ as the number of particles $N \rightarrow \infty$, in expectation of the sup over $t \in [0,1]$ of $|\hat{\mu}_t^{N}(\phi) - \mu_t^{}(\phi)|$, for every $\phi \in C_b^2(\mathbb{R}^d)$.
-
Unification and empirical validation. The framework recovers standard pointwise-reward steering as the special case $\mathcal{R}(\mu) = \int r(x),\mathrm{d}\mu(x)$, for which $\Psi(x,\mu) = r(x)$, and provides a principled counterpart to existing batch-level gradient-steering methods. Empirically, the procedure is verified to target the tilted distribution in tractable low-dimensional settings and examined on higher-dimensional protein conformation tasks.
Main Findings
-
The tilted target is implicit and mean-field. Because $\mu^{*}$ appears on both sides of the optimality condition, the problem lives in a mean-field regime: the target measure is characterised by a potential that depends on the measure itself. This is the structural difference from the pointwise case.
-
Gradient steering alone does not hit the target. In the 1-dimensional bimodal Gaussian mixture experiment (base $p_1$ with means $(-1.0, 1.0)$, both modes unit variance, mode weights $(1,3)$; target $\nu$ with weights $(3,1)$; squared MMD reward; $\lambda^{}=10$), standard gradient steering pushes particles toward the tilting measure $\nu$ but does not target $\mu^{}$, while the mean-field reweighting procedure accurately samples from $\mu^{*}$.
-
Objective gap is minimised at the correct $\lambda$. For $J(\mu)=\mathrm{KL}(\mu|p_1)+\lambda^{}\mathrm{MMD}^2(\mu,\nu)$, the gap $J(\mu)-J(\mu^{})$ is minimised when using mean-field steering with the correct $\lambda$, confirming correct targeting. This holds across three noise schedules: fixed $\sigma_t = 1.0$, decaying $\sigma_t = \sqrt{1-t}$, and the memoryless schedule $\sigma_t = \sqrt{2(1-t)/t}$. Gradient steering, by contrast, samples from a measure whose form depends on the steering strength and does not in general coincide with $\mu^{*}$.
-
Invariance to the noise schedule. The target distribution is reported to be invariant to the choice of noise schedule, indicating that the mean-field procedure's target is explicitly characterised by $\lambda$ rather than indirectly induced by steering hyperparameters.
-
Solver accuracy. Running $K_{FP}=1$ fixed-point iteration in the $\dot{\Psi}t$ solver gives a slight bias, as expected; $K{FP}=3$ and $K_{FP}=5$ are highly accurate and perform very similarly. The choice of solver for $\dot{\Psi}_t$ had little effect on empirical performance, and the authors primarily used the fixed-point iterative solver.
-
HIV-1 protease tilting. Under the base Boltz-2 model, HIV-1 protease states are sampled with proportions 83% closed and 17% semi-open; DEER analysis in Liu et al. (2016) suggests these proportions are 25% and 75%. Steering with the MMD objective over residue-55 C$\alpha$-C$\alpha$ distances pulls generations toward the observed data as $\lambda$ increases, and the induced residue-distance distribution accurately matches the ground-truth proxy $\mu^{*}$; the objective is minimised at the correct $\lambda$ value on the 55-residue distances.
-
4OLE electron-density steering. For PDB structure 4OLE (region of interest residues 423–431), the base model generates almost entirely from the helical conformation A, which does not explain the bimodal electron density. For the same steering strength, mean-field steering produced a stronger reward tilt than gradient-only steering — a larger cosine alignment between generated and target electron densities, and more samples closer to conformation B. Mean-field steering also generally placed more samples in the alternate conformation for the same cosine alignment. Because there is no known ground truth here, the authors discuss these as qualitative differences.
-
Scope of the protein aim. The authors state explicitly that their aim in the protein experiments is not necessarily to improve over gradient steering (which they describe as highly effective), but to understand how far the improved consistency and predictability of the theoretically grounded correction extends to higher-dimensional, challenging settings.
Methodology in Plain English
The starting point is a pretrained flow/diffusion model whose dynamics transport a Gaussian $p_0$ to a data distribution $p_1$. To steer generation, the standard recipe adds a reward-gradient term to the drift and then applies Feynman–Kac log-weight updates plus SMC resampling so that particles follow a tilted path of marginals.
The move here is to replace the pointwise reward $r(x)$ with a reward $\mathcal{R}$ defined on probability measures, and to ask: what does the KL-regularised objective $\arg\max_\mu {\mathcal{R}(\mu) - \mathrm{KL}(\mu|p_1)}$ actually target? Taking the functional derivative and setting it to zero (the paper's Proposition 3.1) gives an exponential tilt in which the potential is itself a function of the target measure — a self-consistency relation. The problem therefore becomes a mean-field one.
The authors then construct stochastic dynamics whose coefficients are chosen so that the law of the process matches the desired tilted path. The position update includes the gradient of the first variation, $\nabla_x \Psi_t(X_t,\mu_t)$, which makes the particles interact — this is exactly the form many existing batch-level steering heuristics take (repulsive potentials, ensemble likelihoods, expected information gain, MMD gradients). The log-weight update includes $b_t(X_t)\cdot\nabla_x\Psi_t(X_t,\mu_t) + \dot{\Psi}_t(X_t)$, where $\dot{\Psi}_t$ is the total time derivative of the potential along the target flow.
$\dot{\Psi}_t$ is the awkward part, because the weights depend on the measure and the measure depends on the weights. Proposition 3.3 turns this into a tractable implicit equation involving the second variation. In practice, the empirical measure $\hat{\mu}t = \sum_j w_t^j \delta{X_t^j}$ with $w_t^j \propto e^{A_t^j}$ replaces the true measure, and $\dot{\Psi}_t$ is computed either by solving a linear system or by a cheap Picard/fixed-point iteration. Positional updates and log-weight updates are then discretised at time steps $\delta t$, and residual resampling at fixed intervals prevents weight degeneracy. The whole loop is Algorithm 1 (single batch), with a multi-batch variant in the appendix. A convergence theorem states the finite-particle approximation converges to the idealised dynamics as $N \rightarrow \infty$, under mild assumptions on the dynamics and the $\dot{\Psi}_t$ solver.
Why This Matters
Impact on research. The paper connects two previously separate literatures — theoretically grounded pointwise-reward Feynman–Kac steering and heuristic batch-level distributional steering — under one mean-field framework. It gives the distributional case a characterisable target $\mu^{*}$, a correctness proof, and a convergence guarantee. It also supplies a theoretical account of what existing heuristic batch-steering methods are implicitly doing, and an explicit correction that fixes the mismatch.
Real-world applications.
- Protein conformational ensemble generation: steering a base structure model (e.g., Boltz-2) so the ensemble matches experimentally inferred state proportions, as demonstrated on HIV-1 protease with DEER-derived 25%/75% proportions.
- Fitting population-level experimental measurements: X-ray crystallography electron densities are spatial and temporal averages over a crystal lattice, so they are inherently distribution-level; the 4OLE example shows a correction that better covers both conformational modes.
- Mode balancing and diversity control: repulsive or diversity-promoting objectives are naturally measure-defined, so the framework applies directly where the goal is to avoid mode collapse.
- Calibration to auxiliary population data: settings where only aggregate statistics (not per-sample labels) are available for steering.
Industry relevance. Inference-time steering requires no retraining and allows different reward functions to be swapped without additional training, which is attractive for deployment. Structured-biology and drug-discovery pipelines are the clearest immediate beneficiaries, since experimental data in those domains is frequently population-level and indirect.
Future Directions
- Reducing the resampling burden. The authors identify reliance on resampling as the primary limitation — particularly problematic for image modalities, where final generations may include near-duplicates. They suggest methods to reduce weight variance (citing Ren et al., 2026) as a mitigation.
- Scaling to other high-dimensional modalities. The paper only examines proteins at scale; extending principled distributional control to images, video, and other modalities where the base literature already reports strong results is an open step.
- Removing the differentiable-reward requirement. The method requires a differentiable reward function, which constrains admissible $\mathcal{R}$.
- More expressive distributional objectives. The authors express hope that future work will draw on domain expertise to formulate richer distributional steering objectives inspired by practical applications — and they note that exact minimisation of a distributional objective is often not strictly necessary, so the framework is best read as a principled foundation for pragmatic gradient-steering approaches.
Target Audience
This paper is most useful to researchers in generative modelling and stochastic processes who work on inference-time control, guidance, or particle-based sampling; to statisticians interested in McKean–Vlasov and mean-field formulations of sampling problems; and to computational biologists or cheminformatics researchers applying diffusion models to protein conformational ensembles and experimental-data fitting. It is not a beginner-friendly introduction to diffusion models — the mathematical content sits at an advanced level, though the algorithmic presentation (Algorithm 1) and the experimental discussion are readable for practitioners already familiar with guidance and SMC-style resampling.
Authors’ abstract
Diffusion models are increasingly used as controllable samplers, whose generations can be steered at inference time according to a chosen reward function. While such rewards are typically defined on individual samples, for many applications it is desirable to steer according to distribution-level rewards, for example to calibrate with population-level information or to encourage diversity. In both cases, simply incorporating the reward gradient into the dynamics, while often effective, comes with few theoretical guarantees on the sampled distribution. For pointwise rewards, recent work has therefore sought to develop a principled framework for targeting a prescribed tilted distribution using particle reweighting. However, an analogous theoretically-grounded approach for distributional rewards is currently lacking. In this work, we formulate inference-time distributional control as targeting a tilted measure under a mean-field framework, and derive a weighted interacting particle scheme to target it in a principled manner. Our framework recovers pointwise-reward steering as a special case, while providing a theoretical foundation for existing batch-level steering methods. Empirically, we verify that the procedure correctly targets the prescribed distribution in tractable low-dimensional settings, and investigate its behaviour in higher-dimensional protein conformation tasks.