Scalable Diffusion SBI for Compositional Inference under Simulator Misspecification

Vincent D. Zaballa and Elliot E. Hui

University of California, Irvine

Paper on arXiv Code (coming soon) BibTeX
Four-minute summary of the paper (no sound).

Simulation-based inference rarely means conditioning on one observation from a perfect simulator. Real datasets contain many heterogeneous observations, latent variables shared across some groups and specific to others, and systematic gaps between the simulator and the system it describes. We treat a pretrained diffusion model as a reusable inference engine that handles all three without being retrained for each new arrangement of the data.

Three-panel schematic. Left: per-observation scores are combined by a composition-aware SDE with a new diffusion coefficient. Middle: shared parameters exchange updates with group-specific latent states. Right: a heat map over two design axes shows path divergence concentrated where simulator and data disagree.
Compositional sampling combines observations with a derived diffusion coefficient \(g_n(t)\) (left). Hierarchical blockwise diffusion sampling alternates updates of shared parameters \(\theta\) and group-specific states \(R_k\) (middle). Path divergence \(\mathcal{A}(\theta^{\mathrm{ref}},\xi)\) maps how much fine-tuning changed the model at each experimental design (right).

What does the framework do?

We start from a Simformer (Gloeckler et al., 2024), a masked diffusion transformer trained with denoising score matching on simulator draws of parameters \(\theta\), observations \(y\), and experimental designs \(\xi\). Because any subset of variables can be conditioned on, the same network answers posterior queries \(p(\theta\mid y,\xi)\) and likelihood queries \(p(y\mid\theta,\xi)\). Everything below happens after pretraining.

The application that motivated this combination is a mechanistic model of Bone Morphogenetic Protein (BMP) signaling, where ligands bind pairs of receptors and the resulting complexes drive a transcriptional response.

How the inference variables map onto the BMP model
SymbolMeaning in the BMP modelDimension
\(\xi\)Experimental design, the doses of five ligands5
\(\theta\)Binding affinities and phosphorylation efficiencies shared by all cell lines60
\(R^k\)Receptor state of cell line \(k\)5
\(y\)Steady-state SMAD response1

The diffusion coefficient depends on the number of observations

Under conditional independence, the posterior given \(n\) observations factorizes into single-observation posteriors, so one score network evaluated \(n\) times can stand in for a model trained on the whole set. F-NPSE (Geffner et al., 2023) uses this in discrete time, and its reverse transitions contract as \(n\) grows. We carry that contraction into continuous time by keeping the variance-preserving drift and prescribing a smaller accumulated noise variance, which fixes the diffusion coefficient.

\[ g_n(t)^2=\beta(t)\,\frac{1+(n-1)\bigl(1-\alpha(t)\bigr)^2}{\bigl[1+(n-1)\bigl(1-\alpha(t)\bigr)\bigr]^2} \]

Here \(\alpha(t)\) is the cumulative squared signal retention and \(\beta(t)\) the usual noise rate. One observation recovers the ordinary VP-SDE. With more, the coefficient falls toward \(\beta(t)/n\) at high noise. The resulting sampler, F-NPSE-SDE, needs only score evaluations. It skips the score Jacobians of JAC and the auxiliary covariance estimates of GAUSS (Linhart et al., 2026).

How does the schedule change with the number of observations?

Diffusion variance rate against reverse progress for the ordinary VP schedule and the composition-aware schedule

Compositional samplers on an exact-score Gaussian (\(n=100\), 400 predictor steps) and on SLCP with learned scores (\(n=30\)). Mean ± standard deviation over five seeds for the Gaussian and over 25 test instances with one network seed for SLCP. \(L\) is the number of Langevin corrections per step.
MethodGaussianSLCP
max-sW ↓sW ↓MMD ↓C2ST → 0.5
GAUSS0.24 ± 0.230.45 ± 0.180.05 ± 0.160.97 ± 0.04
JAC0.26 ± 0.220.53 ± 0.190.05 ± 0.140.95 ± 0.05
Langevin0.52 ± 0.500.81 ± 0.280.13 ± 0.220.97 ± 0.03
F-NPSE0.29 ± 0.240.83 ± 0.290.11 ± 0.130.98 ± 0.03
F-NPSE-SDE, \(L=1\)0.29 ± 0.250.74 ± 0.310.10 ± 0.150.98 ± 0.03
F-NPSE-SDE, \(L=5\)0.31 ± 0.270.73 ± 0.280.10 ± 0.150.98 ± 0.03
Two line charts against number of observations. The left shows the reduction in sliced Wasserstein distance from the derived coefficient, positive for every n above 1. The right shows step size relative to the stability limit, which crosses 1 for the ordinary schedule and stays near 0 for the derived one.
Exact-score Gaussian benchmark with 400 predictor steps and one Langevin correction per step. Left, the reduction in endpoint sliced Wasserstein distance from using \(g_n(t)^2\) in place of \(\beta(t)\) at matched cost, with 95% intervals over five paired seeds. Right, step size relative to the local contraction limit. The ordinary schedule crosses 1 as \(n\) grows and the derived coefficient stays below it.

HBDS introduces hierarchical structure at sampling time

Pooling every observation into one compositional posterior assumes they all share the same latent variables, and grouped data, such as the receptor states \(R^k\) of the BMP model, break that assumption. Each of the four cell lines has its own \(R^k\), while the biophysical parameters \(\theta\) are common to all of them.

Hierarchical Blockwise Diffusion Sampling (HBDS) keeps one shared parameter state and a separate latent state per group. At every reverse-diffusion time each group state is updated from its own data and the current shared parameters, then the shared parameters are updated from all groups. The pretrained network is reused unchanged, with no hierarchical estimator to train and no need to tokenize the whole dataset.

The toy model below isolates this structure. Each group \(g\) has a sign-valued latent \(z_g\), all groups share one parameter \(\theta\), and \(x_{gi}\) is observation \(i\) in group \(g\).

\[ z_g\sim\mathrm{Rad}\!\left(\tfrac12\right),\qquad \theta\sim\mathcal{N}(0,1),\qquad x_{gi}\mid\theta,z_g\sim\mathcal{N}(5z_g+\theta,\,1). \]

\(\mathrm{Rad}(\tfrac12)\) gives the two signs equal probability, so the two groups sit ten units apart when their signs differ.

Two scatter panels of sampled latent state against shared parameter. Pooled sampling puts most samples at one latent value. HBDS gives one cluster per group at the same shared parameter values.
A toy with shared \(\theta\) and a sign-valued latent per group, using learned scores and 20 observations, ten from each group. Pooled sampling ties both groups to one latent. HBDS keeps a latent per group at a common \(\theta\).

Path divergence shows where fine-tuning changed the model

A surrogate trained on simulations inherits the simulator's errors. We adapt it to observed data by sampling responses from the learned likelihood at a fixed reference parameter \(\theta^{\mathrm{ref}}\), scoring them against observations with a negative mean-squared-error reward, and differentiating through the sampler. A path-space KL penalty holds the fine-tuned (FT) model near the pretrained (PT) one.

Schematic of likelihood fine-tuning. Generated responses are compared with observations, and gradients flow back through sampling into token embeddings, FiLM layers, and the score network.
Gradients through sampling update token embeddings, FiLM layers, and the score network. Likelihood and posterior queries share those embeddings, which is how a likelihood-side correction reaches the posterior.

Fine-tuning changes the law of the whole sampling trajectory, and Girsanov's theorem gives that change in closed form.

\[ \mathcal{A}(\theta,\xi):=\mathrm{KL}\bigl(\mathbb{P}_{\phi}^{\theta,\xi}\,\big\|\,\mathbb{P}_{\phi_0}^{\theta,\xi}\bigr)=\tfrac12\,\mathbb{E}_{\mathbb{P}_{\phi}^{\theta,\xi}}\!\left[\int_0^T g(t)^2\,\bigl\|s_{\phi}^{y}-s_{\phi_0}^{y}\bigr\|_2^2\,dt\right] \]

Evaluated across designs at the reference parameter, \(\mathcal{A}\) is a map of how far the model was moved at each experimental condition. It measures the size of the adaptation and says nothing on its own about whether predictions improved, so we always read it next to predictive error.

Two curved surfaces over axes labelled parameters and conditions. The lower blue surface is the simulator-trained family of path laws and the upper green surface is the data-adapted family. They touch at one corner, labelled no adaptation, and separate elsewhere. Arrows from the lower to the upper surface are labelled model adaptation. A small tangent plane sits on the upper surface at one condition.
Each pair of parameters \(\theta\) and conditions \(\xi\) gives one sampling path law, so each model traces out a surface of path laws. Fine-tuning moves the simulator-trained surface \(\mathcal{M}_{\phi_0}\) to the data-adapted surface \(\mathcal{M}_{\phi_*}\). \(\mathcal{A}(\theta,\xi)\) is the size of that move at one point, and it is zero where the model was left unchanged. The tangent plane marks the local Fisher geometry on the adapted surface.
Reading path divergence together with the change in predictive error at the same design
\(\mathcal{A}(\theta,\xi)\)Error improvesError does not improve
SmallAdequate simulator or mild correctionUnder-correction
LargeCorrected misspecificationOver-correction or forgetting

Does path divergence find a known discrepancy?

A linear-Gaussian simulator over a two-dimensional design space omits a localized bump \(h(\xi)\) that is present in the observations. We fine-tune on ten noisy observations and compare predictions over the full design grid.

Four heat maps over design space: the known discrepancy, pretrained absolute error, fine-tuned absolute error, and path divergence. Path divergence is concentrated where the discrepancy is.
The known discrepancy \(h(\xi)\), PT and FT absolute error on a shared scale, and path divergence \(\mathcal{A}(\theta^{\mathrm{ref}},\xi)\), with \(\lambda=0.1\).
Fine-tuning on the analytic toy. Improvement is PT error divided by FT error, so values below 1× mean fine-tuning made things worse.
MetricPTFTImprovement
Posterior \(W_2\) ↓0.090.071.35×
Posterior mean error ↓0.080.032.88×
Likelihood RMSE, correctly specified region ↓0.050.260.18×
Likelihood RMSE, misspecified region ↓0.450.301.49×

Two controls trade correction against preservation. The penalty strength \(\lambda\) acts during training. Autoguidance acts at sampling time and needs no retraining.

Two rows of heat maps across six regularization strengths. The top row shows fine-tuned absolute error and the bottom row path divergence, which fades as regularization increases.
Increasing \(\lambda\) from \(10^{-4}\) to \(10\). Stronger regularization keeps more of the PT model and weakens the correction near the bump. Circles mark observed designs and dashed contours enclose \(h(\xi)\ge 0.1\).

Fine-tuning corrects the model near the discrepancy and adds error elsewhere, while the pretrained model does the opposite. Autoguidance (Karras et al., 2024) combines the two during sampling with the score

\[ s_w^{y}=s_{\mathrm{PT}}^{y}+w\,\bigl(s_{\mathrm{FT}}^{y}-s_{\mathrm{PT}}^{y}\bigr), \]

where \(w=0\) recovers the PT model and \(w=1\) the FT model. An intermediate weight keeps the correction near the bump and reduces the error that fine-tuning introduced at designs far from the observations, which lowers error over the whole design grid. Neither model is retrained, and the only added cost is one more score evaluation per sampling step.

Two rows of heat maps across five autoguidance weights from pretrained to fine-tuned. Intermediate weights show lower average error and smaller path divergence.
Autoguidance with weight \(w\) from PT (0) to FT (1). At \(w=0.5\) the full-grid mean absolute error is 0.12, against 0.14 for PT and 0.15 for FT. The weight was chosen retrospectively in a single-seed comparison.

Fine-tuning with HBDS improves predictive fit on 940 BMP measurements

We combine fine-tuning with HBDS on 940 measurements from a parent cell line (NMuMG) and three receptor knockdowns. The model has 60 shared biophysical parameters and five receptor quantities per cell line. The fine-tuning reference is the least-squares fit (LSR) that is standard for this model, used as an anchor and not as ground truth.

Measurements
940
Cell lines
4
Shared parameters
60
Receptor quantities per line
5
Three columns of plots for a BMP4 and BMP10 competition series in NMuMG, the same series in the BMPR2 knockdown, and a BMP4 titration in NMuMG. Rows show predictions against observations, best-fitting least-squares cost group, path divergence, and change in prediction error.
From top to bottom, LSR and FT predictions against observations with parameters held at the LSR reference, the LSR cost group that fits each design best, path divergence relative to the median of its series, and the reduction in absolute error from LSR to FT on a symmetric log scale.

With parameters held at the reference, FT improves on the simulator throughout the BMPR2-knockdown competition series and at lower BMP4 doses. In the titration those gains generally coincide with larger path divergence. In the NMuMG competition series adaptation often increases error, which is the case the table above labels over-correction.

Posterior-predictive fit with pooled sampling and HBDS for the PT model, the FT model, and an FT model trained without FiLM. 250 draws and one sampling seed.
Model and samplerRMSE ↓Median distance ↓
PT, pooled0.581,594.55
FT, pooled0.3666.77
FT without FiLM, pooled12.304,538.82
PT, HBDS0.78194.01
FT, HBDS0.3414.94
FT without FiLM, HBDS0.91274.61

We benchmark fine-tuning against the most relevant existing method, Flow Matching Corrected Posterior Estimation (FMCPE). FMCPE trains a neural posterior estimator on simulations and then corrects it with a smaller calibration set, which we build here from LSR fits. It needs a dedicated posterior model and correction flows for each observation set, so every row with a different number of observations below was retrained. The Simformer reuses one set of weights across observation sets through masking and compositional sampling. FMCPE still gives a direct comparison, because both methods can be scored by the same predictive errors through the BMP simulator.

Comparison with FMCPE (Ruhlmann et al., 2026), which corrects a posterior estimator using calibration pairs built here from LSR fits. Predictive errors through the original BMP simulator on all 940 measurements, 250 parameter draws per run, mean ± SE over three seeds. FMCPE rows use the largest LSR budget at each observation count.
MethodObservationsLSR fitsRMSE ↓Median distance ↓
FMCPE-LSR2354,8160.168 ± 0.011136.460 ± 77.855
FMCPE-LSR4704,8160.177 ± 0.00331.505 ± 10.206
FMCPE-LSR7054,8160.168 ± 0.00750.355 ± 16.256
FMCPE-LSR9404,8160.175 ± 0.00531.051 ± 8.460
PT BMP Simformer94000.252 ± 0.00311.111 ± 0.116
FT BMP Simformer94010.245 ± 0.00310.376 ± 0.058
Show all 16 FMCPE budget settings
ObservationsLSR fitsRMSE ↓Median distance ↓
2351,2040.242 ± 0.004203.661 ± 38.307
2352,4080.231 ± 0.026203.451 ± 70.019
2353,6120.184 ± 0.01077.887 ± 19.521
2354,8160.168 ± 0.011136.460 ± 77.855
4701,2040.256 ± 0.008283.935 ± 26.484
4702,4080.202 ± 0.006108.803 ± 19.193
4703,6120.195 ± 0.003154.748 ± 46.537
4704,8160.177 ± 0.00331.505 ± 10.206
7051,2040.246 ± 0.003296.322 ± 88.535
7052,4080.197 ± 0.014161.322 ± 41.237
7053,6120.193 ± 0.00298.697 ± 18.042
7054,8160.168 ± 0.00750.355 ± 16.256
9401,2040.270 ± 0.012203.950 ± 20.409
9402,4080.210 ± 0.025188.117 ± 73.589
9403,6120.186 ± 0.02362.011 ± 26.766
9404,8160.175 ± 0.00531.051 ± 8.460
Show the regularization and Langevin-correction sweep
Each setting used 250 HBDS draws and one sampling seed. \(L\) is the number of corrector steps per diffusion step. Each cell gives RMSE and median distance.
\(\lambda\)\(L=0\)\(L=1\)\(L=3\)\(L=5\)
\(0\)0.29 / 12.160.36 / 12.960.36 / 12.970.36 / 12.97
\(10^{-4}\)0.30 / 28.150.32 / 89.550.32 / 90.630.32 / 90.87
\(5\times10^{-4}\)0.28 / 14.300.34 / 14.890.33 / 14.620.34 / 14.94
\(10^{-3}\)0.34 / 29.970.36 / 41.200.36 / 41.360.36 / 41.86
\(10^{-2}\)0.33 / 45.260.40 / 93.510.42 / 94.780.41 / 99.30
\(10^{-1}\)0.30 / 52.170.52 / 157.550.53 / 153.480.48 / 159.86
\(1\)0.30 / 47.000.40 / 130.730.41 / 134.900.41 / 139.27

Likelihood fine-tuning is anchored at a single reference \(\theta^{\mathrm{ref}}\), the LSR fit, so one might expect posterior inference to reduce to that point estimate. It does not. The FT posterior marginals generally shift toward the reference relative to pretraining while retaining substantial spread, which indicates that the likelihood-side correction transfers to posterior queries without collapsing inference onto the fine-tuning anchor.

Grid of sixty small histograms comparing pretrained and fine-tuned posterior marginals for binding affinities and phosphorylation efficiencies, with a red line at the least-squares reference in each.
Posterior marginals for the 30 binding affinities (upper block) and 30 phosphorylation efficiencies (lower block). Purple is PT, green is FT, and red lines mark the LSR reference. FT marginals generally move toward the reference and keep their spread.

What does the path-space Fisher metric say about parameters and designs?

Path divergence compares two models at the same \((\theta,\xi)\). The same construction also compares one model with itself at nearby conditioning values. Writing \(z=(\theta,\xi)\), the KL divergence between path laws at \(z\) and \(z+\delta z\) is, to second order, \(\tfrac12\,\delta z^{\top}\mathcal{I}_{\phi}^{\mathrm{path}}(z)\,\delta z\) with

\[ \mathcal{I}_{\phi}^{\mathrm{path}}(z)=\mathbb{E}_{Y\sim\mathbb{P}_{\phi}^{z}}\!\left[\int_0^T g(t)^2\,J_z s_{\phi}^{y}(Y_t,t;z)^{\top}J_z s_{\phi}^{y}(Y_t,t;z)\,dt\right]=\begin{bmatrix} I_{\theta\theta} & I_{\theta\xi}\\ I_{\xi\theta} & I_{\xi\xi}\end{bmatrix}. \]

This is a Fisher metric on the family of path laws \(\mathcal{M}_{\phi}\). It can be singular where the conditioning variables are not locally identifiable, for example when different parameter combinations give indistinguishable responses.

A curved surface of path laws with two marked points that share parameters and differ in design. Below each point is its tangent plane with one arrow for a parameter perturbation and one for a design perturbation. At the first point the two arrows are nearly parallel and span a thin sliver. At the second point they point in clearly different directions and span a wide parallelogram.
Tangent planes at two designs \(\xi_a\) and \(\xi_b\) with the same parameters. \(v_\theta\) and \(v_\xi\) are the first-order changes in the path law from perturbing the parameters and the design, and their inner product is \(\delta\theta^{\top}I_{\theta\xi}(z)\,\delta\xi\). At \(\xi_a\) the two are nearly aligned, so a change in parameters looks like a change in design. At \(\xi_b\) they are distinct.

Each block answers a different question.

  • \(I_{\theta\theta}\) measures how much the path law changes with the mechanistic parameters at a fixed design. Designs where it is large are locally informative about \(\theta\).
  • \(I_{\xi\xi}\) measures how quickly the path law changes across designs. For the pretrained model, large values mark regions where denser simulator coverage would help most.
  • \(I_{\theta\xi}\) measures whether a parameter change and a design change have the same first-order effect.

The paper proposes using these blocks to move a population of reference parameters during fine-tuning, to choose designs, and to weight a fine-tuning curriculum. These are proposed extensions. None of them is implemented or tested in the experiments above.

What are the main results?

What are the limitations?

Compositional sampling and HBDS both rely on approximations to intermediate-time posteriors. Fine-tuning at a fixed reference can improve misspecified regions while degrading regions that were already well modeled, and on the BMP data the comparison with the least-squares fit remains mixed. These results point toward adaptation that varies across experimental designs instead of one global regularization strength.

Citation

@article{zaballa2026scalable,
  title   = {Scalable Diffusion {SBI} for Compositional Inference
             under Simulator Misspecification},
  author  = {Zaballa, Vincent D. and Hui, Elliot E.},
  journal = {arXiv preprint arXiv:2609.36950},
  year    = {2026},
  url     = {https://arxiv.org/abs/2609.36950}
}

References

  1. Geffner, T., Papamakarios, G., and Mnih, A. (2023). Compositional score modeling for simulation-based inference. International Conference on Machine Learning, PMLR 202, 11098–11116.
  2. Gloeckler, M., Deistler, M., Weilbach, C. D., Wood, F., and Macke, J. H. (2024). All-in-one simulation-based inference. International Conference on Machine Learning, PMLR 235, 15735–15766.
  3. Karras, T., Aittala, M., Kynkäänniemi, T., Lehtinen, J., Aila, T., and Laine, S. (2024). Guiding a diffusion model with a bad version of itself. Advances in Neural Information Processing Systems.
  4. Linhart, J., Victorino Cardoso, G., Gramfort, A., Le Corff, S., and Rodrigues, P. L. C. (2026). Diffusion posterior sampling for simulation-based inference in tall data settings. Transactions on Machine Learning Research.
  5. Ruhlmann, P.-L., Arbel, M., Forbes, F., and Rodrigues, P. L. C. (2026). Flow matching calibration for simulation-based inference under model misspecification. International Conference on Machine Learning.