It Just Takes Two: Scaling Amortized Inference to Large Sets
AuthorsAntoine Wehenkel, Michael Kagan, Lukas Heinrich, Chris Pollard
Resources
The paper shows that you can train set-based neural posterior estimators on tiny sets of two and still make them work on huge observation sets, cutting compute while keeping accuracy strong.
Key results
PAIRS pretrains the mean-pool Deep Set encoder only on sets of size 1 and 2, which the theory says is enough to recover the same sufficient representation as training at larger cardinalities.
On the bivariate Gaussian model with a closed-form posterior, PAIRS tracks the analytical posterior up to N = 10^5 without further encoder training.
For the particle-physics bump hunt, PAIRS matches MCMC at N = 100, indicating the learned aggregator captures the shared nuisance coupling across observations.
In conditional flow matching for novel-view synthesis, the target-view mean square error decreases from 0.011 at N = 1 to 0.004 at N = 100.
Compared with training at N = 1–10, PAIRS is reported to be 3–4× cheaper while usually matching or outperforming those baselines.
End-to-end training at N = 1000 is reported to cost about 100× more compute than PAIRS, without consistent gains.
What the paper found
It Just Takes Two introduces PAIRS, a training scheme for neural posterior estimation that decouples representation learning from posterior modeling in set-structured inference problems with shared nuisance variables. The key result is that, for mean-pool Deep Sets, an encoder trained only on sets of size 1 and 2 recovers the same sufficient representation, up to an affine transform, as training at any larger cardinality; the proof reduces the N = 2 case to a Cauchy functional equation. This lets the encoder be pretrained once at N ≤ 2, cached on individual observations, and then a density head be finetuned on pre-aggregated embeddings, making training cost essentially independent of deployment set size. Across benchmarks with N in the thousands, including a bivariate Gaussian model with closed-form posterior, a particle-physics bump hunt, Circle Radius images, rotated MNIST Digit Expectation, multi-view 3D object properties from ModelNet40, GEOM-Drugs molecular lipophilicity, and novel-view synthesis with conditional flow matching, PAIRS matches or outperforms stronger baselines. On the Gaussian sanity check it tracks the analytical posterior up to N = 10^5; on the bump hunt it matches MCMC at N = 100; and on novel-view synthesis it reduces target-view MSE from 0.011 at N = 1 to 0.004 at N = 100. Compared with training at N = 1–10, PAIRS is 3–4× cheaper and usually better, while end-to-end training at N = 1000 costs about 100× more compute without consistent gains.
Original abstract
Neural posterior estimation has emerged as a powerful tool for amortized inference, with growing adoption across scientific and applied domains. In many of these applications, the conditioning variable is a set of observations whose elements depend not only on the target but also on unknown factors shared across the set. Optimal inference therefore requires treating the set jointly, which in turn requires training the estimator at the deployment set size -- a regime where memory and compute quickly become prohibitive. We introduce a simple, theoretically grounded strategy that decouples representation learning from posterior modeling. Our method trains a mean-pool Deep Set on sets of size at most two, producing an encoder that generalizes to arbitrary set sizes. The inference head is then finetuned on pre-aggregated embeddings, making training cost essentially independent of the deployment set size N. Across scalar, image, multi-view 3D, molecular, and high-dimensional conditional generation benchmarks with N in the thousands, our approach matches or outperforms standard baselines at a fraction of the compute.
Read the original paper