ProximalFM: Amortized Proximal Causal Inference under Hidden Confounding
Organizations: University of Oxford · Harvard University · University of Illinois Urbana-Champaign · Perelman School of Medicine, University of Pennsylvania · Causal Chamber · U. of Toronto & Vector Institute · University College London · SMARTbiomed, University of Oxford
Abstract
Standard causal identification methods often assume no unmeasured confounding and can fail when relevant confounders are unobserved. Proximal causal inference instead uses proxy variables to identify effects under hidden confounding. However, nonparametric proximal estimation can be challenging in practice: recovering causal estimands such as the conditional average treatment effect (CATE) requires solving an ill-posed integral equation that is data-hungry, hyperparameter-sensitive, and optimization-unstable. Bayesian inference for such models provides a desirable alternative, mitigating these difficulties by regularizing through the prior. However, computing a posterior is itself challenging, as a typical likelihood function will include latent variables. Following the recent success of tabular foundation models in backdoor, instrumental variable, and frontdoor settings, we propose that prior-data fitted networks (PFNs) are uniquely suited to resolve this bottleneck. Indeed, by training on synthetic data sampled from compliant structural causal models with access to oracle counterfactuals, we simplify the task substantially, amortizing the implied Bayesian operator inversion into a single transformer forward pass. Compared to prior literature that focuses primarily on point estimation, our model, ProximalFM, explicitly targets the Bayesian posterior distribution of the CATE. One unique aspect of this problem is that we need to provide Monte Carlo estimates of the oracle CATEs, leading to a novel variation of PFNs that accounts for the added stochastic error. Across a diverse suite of proximal regimes, ProximalFM achieves consistently strong CATE-estimation performance without dataset-specific tuning, with its largest advantage when latent confounding is substantial and the proxies are weakly informative; it also provides fast inference through a single amortized forward pass.
Figures & tables
Appendix figures & tables52 assets
Supplementary material from the paper’s appendix.
Appendix
| hyperparameter | value | description |
|---|---|---|
| Block dimensions | ||
| dim_X | observed baseline covariate dimension | |
| dim_U | triangular on | latent confounder dimension |
| dim_W | outcome-inducing proxy dimension | |
| dim_Z | treatment-inducing proxy dimension | |
| Dimension bounds | ||
| hyperparameter | value | description |
| Mechanism | ||
| noise_std | tiny fixed per-layer dither so continuous columns realise distinct values | |
| Feature corruption | ||
| max_categories | cap on categories per feature | |
| proxy_cat_cover_prob | chance a proxy mirrors a category (vs. a continuous wildcard) | |
| Positivity | ||
| hyperparameter | sampling law | description |
| Mechanism | ||
| num_layers | log-scaled trunc. normal, mean , int | mechanism MLP depth |
| hidden_dim | log-scaled trunc. normal, mean , int | mechanism MLP width |
| mlp_activations | categorical over 54 random activation functions | activation, drawn from a random library |
| init_std | log-scaled trunc. normal, mean , real | weight-initialisation scale |
| block_wise_dropout | categorical (random weights) | block-sparse vs. dense weight init |
| Proxy profile | Heuristic? | Reason |
|---|---|---|
| ok | exact match, column for column | |
| ok | categorical columns carry more categories than required | |
| ok | a numerical column covers ’s -cat column | |
| ok | numerical columns dominate every requirement | |
| no | too few columns ( ) | |
| no | the -cat column cannot cover ’s -cat column |
| conditioning set | role | datasets | median | 90th pct | share |
| excluded from | |||||
| conditions on | all | -0.002 | 0.007 | 0.021 | |
| conditions on | top confounded | -0.005 | 0.004 | 0.012 | |
| positive control | all | 0.004 | 0.099 | 0.313 | |
| positive control | top confounded | 0.029 | 0.222 | 0.598 | |
| excluded from | |||||
| rung | adjustment set | role | median error | 90th pct | share |
|---|---|---|---|---|---|
| naive | none | unadjusted | 0.041 | 0.195 | 0.241 |
| backdoor X | correct if no latent confounder | 0.037 | 0.163 | 0.211 | |
| backdoor XWZ | DAG-blind: everything measured | 0.030 | 0.121 | 0.148 | |
| oracle XU | identification: minimal sufficient set | 0.028 | 0.104 | 0.108 | |
| oracle XUW | efficiency: adds , a parent of | 0.025 | 0.105 | 0.107 |
| aggregation over | median | share | noise | |||
|---|---|---|---|---|---|---|
| mean | 0.273 | 0.226 | 0.166 | -0.002 | 0.208 | -0.202 |
| weakest | 0.053 | 0.593 | 0.034 | -0.249 | 0.277 | -0.140 |
| mean | 0.271 | 0.235 | 0.166 | -0.023 | 0.226 | -0.208 |
| weakest | 0.052 | 0.588 | 0.017 | -0.272 | 0.281 | -0.095 |
| Stage | Hyperparameter | Value |
| Shared | Embedding dimension | 128 |
| Feedforward expansion factor | 2 | |
| Activation | GELU | |
| Normalization | pre-norm | |
| Dropout | ||
| Column embedder | Blocks | 3 |
| Stage | Hyperparameter | Value |
| Prior | Confounder dimension | – |
| Treatment proxy dimension | – | |
| Outcome proxy dimension | – | |
| Covariate dimension | – | |
| Context units | – (uniform) | |
| Query units | 128 |
| Hours | % | ||
| Compute | Prior generation | ||
| Forward / backward | |||
| Validation, prior set (376 passes) | |||
| Validation, realistic set (751 passes) | |||
| Checkpointing, logging, startup | |||
| Total | GPU-hours ( H100 NVL) |
| Predictor | Family | Fitted estimator and prediction rule | Details |
|---|---|---|---|
| S-learner (ridge) | Backdoor meta-learner | S-learner with ridge regression. | Section D.1.1 |
| S-learner (RF) | Backdoor meta-learner | S-learner with random forest. | Section D.1.1 |
| S-learner (TabICL) | Backdoor meta-learner | S-learner with zero-shot TabICL. | Section D.1.1 |
| S-learner (TabICL) + kNN marg. | Backdoor meta-learner | S-TabICL fit on , then 16-neighbour proxy marginalisation. | Section D.1.1 |
| S-learner (TabICL) + kernel marg. | Backdoor meta-learner | S-TabICL fit on , then 16-neighbour kernel proxy marginalisation (scale 1). | Section D.1.1 |
| T-learner (ridge) | Backdoor meta-learner | T-learner with ridge regression. | Section D.1.1 |
| Component | Setting | Value |
| Cross-fitting | Folds (stratified on ) | 4 |
| Bridge class | Kernel | RBF, blockwise median heuristic |
| Nyström rank | 256 | |
| Bridge regularization | sweep : | |
| Selection criterion | Held-out proximal moment loss | |
| Validation fraction |
| Component | Setting | Value |
| Bridge class | MLP, ReLU, linear output | |
| Hidden width | searched : | |
| Depth (hidden layers) | searched : | |
| Moment kernel | RBF over | |
| Length scale | fixed at or median heuristic | |
| Loss | U- or V-statistic ( 50 ) | one method each |
| Component | Setting | Value |
| Kernels | Bandwidths baseline | per-column median heuristic |
| Kernel bandwidth multiplier (W,X) | searched : | |
| Kernel bandwidth multiplier (Z) | searched : | |
| Regularisation | (RKHS norm) | swept : |
| Cholesky jitter | ||
| Prediction | strategy | “ind. W”, “cond. W” |
| Component | Setting | Value |
|---|---|---|
| Data split | (stage-1 rows) | |
| (stage-2 rows) | ||
| Kernels | Kernel family | Column-wise Gaussian RBF for ; binary for |
| Base bandwidths | Per-column median heuristic, estimated on the context | |
| Shared multiplier on | Shared across ; searched: | |
| Regularization | searched by LOOCV; 10-point log grid in |
| Source | Rows | Features |
|---|---|---|
| pol | 26 | |
| MiniBooNE | 50 | |
| default-of-credit-card-clients | 20 | |
| Higgs | 24 | |
| jannis | 54 | |
| heloc | 22 |
| Linear | Nonlinear | |||||||
| Low confounding | High confounding | Low confounding | High confounding | |||||
| Method | Low proxy | High proxy | Low proxy | High proxy | Low proxy | High proxy | Low proxy | High proxy |
| S-learner (ridge) | 1.027 | 1.027 | 1.458 | 1.458 | 1.059 | 1.059 | 1.459 | 1.459 |
| S-learner (RF) | 0.868 | 0.868 | 1.229 | 1.229 | 0.986 | 0.986 | 1.542 | 1.542 |
| S-learner (TabICL) | 0.782 | 0.782 | 1.336 | 1.336 | 0.937 | 0.937 | 1.457 | 1.457 |
| S-learner (TabICL) + kNN marg. | 0.885 | 0.895 | 1.070 | 0.921 | 0.973 | 0.982 | 1.259 | 0.986 |
| Linear | Nonlinear | |||||||
|---|---|---|---|---|---|---|---|---|
| Low confounding | High confounding | Low confounding | High confounding | |||||
| Method | Low proxy | High proxy | Low proxy | High proxy | Low proxy | High proxy | Low proxy | High proxy |
| S-learner (ridge) | 1.013 | 1.013 | 1.564 | 1.564 | 1.023 | 1.023 | 1.580 | 1.580 |
| S-learner (RF) | 0.669 | 0.669 | 1.423 | 1.423 | 0.849 | 0.849 | 1.546 | 1.546 |
| S-learner (TabICL) | 0.499 | 0.499 | 1.294 | 1.294 | 0.762 | 0.762 | 1.432 | 1.432 |
| S-learner (TabICL) + kNN marg. | 0.527 | 0.581 | 1.053 | 0.683 | 0.793 | 0.822 | 1.206 | 0.865 |
| Linear | Nonlinear | |||||||
|---|---|---|---|---|---|---|---|---|
| Low confounding | High confounding | Low confounding | High confounding | |||||
| Method | Low proxy | High proxy | Low proxy | High proxy | Low proxy | High proxy | Low proxy | High proxy |
| S-learner (ridge) | 1.011 | 1.011 | 1.567 | 1.567 | 1.024 | 1.024 | 1.593 | 1.593 |
| S-learner (RF) | 0.568 | 0.568 | 1.384 | 1.384 | 0.793 | 0.793 | 1.497 | 1.497 |
| S-learner (TabICL) | 0.379 | 0.379 | 1.256 | 1.256 | 0.653 | 0.653 | 1.393 | 1.393 |
| S-learner (TabICL) + kNN marg. | 0.393 | 0.426 | 0.994 | 0.552 | 0.694 | 0.732 | 1.169 | 0.783 |
| Linear | Nonlinear | |||||||
|---|---|---|---|---|---|---|---|---|
| Low confounding | High confounding | Low confounding | High confounding | |||||
| Method | Low proxy | High proxy | Low proxy | High proxy | Low proxy | High proxy | Low proxy | High proxy |
| S-learner (ridge) | 1.011 | 1.011 | 1.560 | 1.560 | 1.022 | 1.022 | 1.606 | 1.606 |
| S-learner (RF) | 0.513 | 0.513 | 1.356 | 1.356 | 0.740 | 0.740 | 1.476 | 1.476 |
| S-learner (TabICL) | 0.310 | 0.310 | 1.239 | 1.239 | 0.564 | 0.564 | 1.371 | 1.371 |
| S-learner (TabICL) + kNN marg. | 0.291 | 0.298 | 0.953 | 0.441 | 0.580 | 0.624 | 1.111 | 0.695 |
| Method | Context size | ||||||
|---|---|---|---|---|---|---|---|
| S-learner (ridge) | 1.459 | 1.499 | 1.580 | 1.593 | 1.606 | 1.597 | 1.608 |
| S-learner (RF) | 1.542 | 1.536 | 1.546 | 1.497 | 1.476 | 1.602 | 1.473 |
| S-learner (TabICL) | 1.457 | 1.449 | 1.432 | 1.393 | 1.371 | 1.340 | 1.317 |
| S-learner (TabICL) + kNN marg. | 1.259 | 1.235 | 1.206 | 1.169 | 1.111 | 1.065 | 1.020 |
| S-learner (TabICL) + kernel marg. | 1.245 | 1.231 | 1.196 | 1.159 | 1.111 | 1.064 | 1.016 |