stat.MLOct 6, 2026

ProximalFM: Amortized Proximal Causal Inference under Hidden Confounding

Authors: Christophe Muller, Ayub Kharel, Alex Luedtke, Chan Park, Eric Tchetgen Tchetgen, Juan L. Gamella, Rahul Krishnan, Ricardo Silva, +1 more

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

Explore similar work

Sep 30, 2026stat.ML

Proximal Balancing for Causal Effect Estimation under Unmeasured Confounding

Estimating causal effects from observational data is central to science and policy, but the effects are not identified when confounders are unmeasured. Proximal causal inference addresses this problem with proxies of the unmeasured confounders. However, existing proxy-based approaches either designate proxy roles and solve an inverse problem, which is ill-posed and hard to estimate with high-dimensional proxies, or use a latent-variable model, which assumes that the learned latent variable matches the hidden confounder and leaves bias when it does not. To address these challenges, we introduce proximal balancing. It carries the classical idea of covariate balancing to confounders that are observed only through proxies: it learns a low-dimensional summary of the covariates and proxies that makes the treatment groups comparable, and then adjusts for this summary. It needs no designated proxy roles, inverse problem, or latent model. We give identification theory, finite-sample guarantees, and a practical algorithm, PROBE. We demonstrate the method on low-dimensional, high-dimensional, and image proxies and on real-world data.
Aug 2, 2026cs.LG

Spatiotemporal Proximal Causal Inference under Hidden Confounding and Interference

Estimating causal effects from real-world spatiotemporal data is challenging due to hidden confounders and interference. Standard causal identification methods assume conditional exchangeability given observed covariates, which fails whenever hidden confounders affect both treatment and outcomes - a common setting in domains such as climate, environmental policy, epidemiology, and regional economics. In this paper, we propose a novel spatiotemporal proximal causal inference framework that extends proximal identification theory to spatiotemporal settings. The proposed method jointly captures local and neighborhood-level confounding information by introducing treatment- and outcome-inducing proxies, and we derive a spatiotemporal outcome confounding bridge function that identifies the potential outcome without requiring direct recovery of the hidden confounder. We establish the identifiability of this bridge function under proxy exclusion restrictions and a spatiotemporal completeness condition, and show that the resulting estimator recovers the outcome through a proximal generalization of the g-computation formula. To operationalize this identification result, we propose a neural architecture that learns proxies via transformer-based spatiotemporal encoders - coupled with a conditional mutual information critic to enforce exclusion restrictions and a moment-matching network to guarantee that the learned bridge function satisfies the underlying identifying equation. We further introduce a stabilized weighting scheme to address treatment support imbalance. Experiments on synthetic datasets demonstrate that our approach achieves comparable performance to baseline causal inference methods, while providing, to our knowledge, the first theoretically grounded outcomes for the hidden confounding in the presence of spatiotemporal interference through a proximal causal inference framework.
Jun 16, 2026cs.LG

FoundCause: Causal Discovery with Latent Confounders from Observational Data

Causal discovery from observational data remains challenging due to the need to recover directed structure and latent confounding without interventions. We propose FoundCause, an amortized causal discovery model trained entirely on synthetic data that maps datasets directly to causal graphs in a single forward pass. By learning from large collections of simulated structural causal models, FoundCause captures transferable statistical patterns that generalize beyond individual datasets. The architecture incorporates several key inductive biases for causal discovery. It uses a permutation-invariant transformer encoder with alternating attention over samples and variables to jointly model cross-variable dependence and per-variable distributions. Pairwise statistical features derived from classical asymmetry measures are injected through statistics-conditioned attention, guiding the model toward known causal signals. A factorized decoder separates edge existence from direction, while a triangular refinement module enables reasoning over higher-order causal motifs such as chains and colliders. In addition, a dedicated confounder module based on learnable latent tokens explicitly models hidden common causes, and the model explicitly handles missing data via its masked input representation. To our knowledge, FoundCause is the first amortized causal discovery approach to explicitly model latent confounding. FoundCause outperforms 11 classical non-amortized methods (e.g., PC, GES, NOTEARS-style optimization) and 4 amortized causal discovery methods on 15 real-world datasets, achieving +9.6% improvement in F1F_1, +1.2% in AUROC, and an 18.9% reduction in structural Hamming distance relative to the strongest non-amortized methods, while performing inference in a single forward pass.