Source-linked AI summary
Foundation Models for Partial Causal Identification
Alexis Bellot, Anish Dhir
TL;DR
Partial causal identification lacks a general way to bound causal queries across varied structures and datasets when observational data leave multiple values compatible. The paper trains a Prior-Data Fitted Network under a full-support canonical prior over discrete structural causal models, showing that its predictive support asymptotically converges to the identified set. This reframes counterfactual bounding as amortized posterior prediction for broad classes of queries, datasets, and structural assumptions.
Problem
Existing partial-identification methods require bespoke derivations for specific causal structures, queries, and dataset types, while observational data may support multiple causal-query values.
Method
The paper trains a Prior-Data Fitted Network using a canonical prior with full support over structural causal models with discrete observables, mapping observational data to posterior distributions over causal queries.
Results
The predictive distribution converges asymptotically to the Bayesian posterior, and its support converges to the identified set for the causal query.
Takeaways & Limitations
Counterfactual bounding can be treated as amortized posterior prediction, potentially allowing one model to handle broad classes of queries, datasets, and structural assumptions.
Takeaways & Limitations
The analysis assumes endogenous variables have finite, discrete domains, although exogenous variables may take values in continuous domains.
Abstract
from arXiv · showhide
This paper investigates the development of causal foundation models for bounding the effect of interventions and counterfactuals from observational data. We show that a canonical prior can be defined with full support over the space of structural causal models with discrete observables. With this canonical prior, we translate the problem of bounding counterfactuals into that of learning distributions over functions that map data (and possibly structural assumptions) to a causal query of interest. This extends the promising causal foundational modelling paradigm to the estimation of partially-identifiable causal effects, i.e., under unobserved confounding, where multiple values are equally compatible with the observed data and prior structural assumptions.
1. Introduction
The paper addresses partial causal identification, where causal queries may not be uniquely determined from observational data, by developing a foundation-model approach that outputs posterior distributions encoding identified sets.
- Motivation: Partial causal identification seeks tight bounds containing all causal-query values compatible with observed data and additional assumptions when unique computation is unavailable.This generalizes causal estimation from identifying a single value to identifying a set.
- Motivation: Existing partial-identification methods are fragmented because different causal structures, queries, and dataset types require separate bespoke derivations and optimization runs.The paper contrasts this with Prior-Data Fitted Networks, which learn distributions over functions between general spaces such as datasets and causal queries.
- Approach: The proposed foundation model takes an observational dataset and returns a posterior distribution over a causal query whose support provably encodes the identified set.The framework is designed to accommodate observational data and potentially additional structural knowledge such as a causal diagram.
- Preliminaries: Structural causal models provide the hypothesis space by specifying endogenous and exogenous variables, generative functions, and the exogenous distribution underlying observed and counterfactual behavior.Interventions replace causal mechanisms, allowing counterfactual quantities such as P(yx, . . . , zw) to be formulated.
- Preliminaries: The paper defines optimal counterfactual bounds as the lower and upper limits over counterfactual probabilities induced by structural causal models compatible with an observational distribution.The identified set is denoted IΨ(PM), and the formulation encompasses causal effects such as P(yx).
- Preliminaries: The framework assumes finite, discrete domains for endogenous variables, while exogenous variables may have continuous domains; resulting counterfactual probabilities are categorical distributions.This assumption is an explicit scope condition for the paper’s analysis.
2. Universal Causal Estimation
The paper develops a universal solver for partial identification by learning posterior distributions over causal queries from observational data and canonical SCM priors. Canonical parameterization and full-support priors make the approach broadly expressive and asymptotically faithful to compatible counterfactual values.
- A trained model implicitly solves the optimal bounding problem for any query derived from an SCM rather than targeting one diagram or bound algorithm.
- The method approximates posterior support bounds, whose limits should capture the full range of counterfactual values compatible with observed data as sample size grows.
- Canonical SCMs: Canonical SCMs use categorical exogenous distributions with finite fixed support, while representing the same counterfactuals as arbitrary SCMs with discrete observables.
- Canonical Priors: The canonical parameterization is maximally expressive and can represent any SCM with discrete observables, motivating priors over causal diagrams and exogenous parameters.
- Canonical Priors: A Bernoulli graph model and Dirichlet exogenous distributions give every graph and categorical distribution positive probability under the canonical prior.
- Canonical Priors: Full support ensures that no SCM consistent with the observed distribution is assigned zero prior mass, preserving posterior faithfulness to the data.
- Prior-Data Fitted Networks: A Prior-Data Fitted Network learns the posterior from synthetic prior samples and can condition on partial or full causal-diagram specifications.
3. What do we converge to?
The PFN predictive distribution converges to the canonical-prior posterior over causal queries, whose support recovers the identified set and whose credible sets attain asymptotic frequentist coverage.
- The PFN predictive distribution converges weakly to the posterior induced by the canonical prior conditional on the observational distribution.
- With sufficient capacity, the PFN’s predictive distribution becomes the exact Bayesian posterior over the causal query.
- Asymptotically, the predictive distribution’s support converges to the identified set, leaving only partial-identification uncertainty.
- A Bayesian credible set derived from the PFN is an asymptotically valid frequentist confidence set for the identified set.
- This Bayesian credible set achieves asymptotic frequentist coverage for the identified set.
- The 100% Bayesian credible interval converges asymptotically to the true identified set for any observational dataset and counterfactual query.
4. Experiments
The experiments evaluate partial-identification inference on binary systems, comparing predicted intervals with analytical bounds and existing methods across coverage, width, and inference time.
- The two-variable experiment estimates interventional and counterfactual probabilities from observational samples in binary structural causal models.
- Figure 1 compares posterior distributions with tight analytical lower and upper bounds for randomly selected observational distributions.
- Table 1 reports coverage, average interval width, and inference time across context sizes for the CFM and a Gibbs sampler.
- The Gibbs-sampler comparison unions predicted sets across all two-variable causal diagrams because the sampler requires a causal diagram as input.
5. Conclusion
The paper introduces a causal foundation-model approach that amortizes partial identification through posterior prediction under a canonical prior with full support over discrete SCMs.
- A Prior-Data Fitted Network is trained under a canonical prior with full support over discrete structural causal models.
- The approach interprets counterfactual bounding as amortized posterior prediction over causal queries.
- A single model can potentially handle broad classes of queries, datasets, and structural assumptions without deriving problem-specific bounds each time.
A. Related Work
Partial-identification methods bound causal effects when observational data do not uniquely determine them, using structural or sensitivity assumptions to tighten those bounds.
- Causal effects may remain partially identifiable under unobserved confounding or unknown causal structure, yielding non-trivial bounds rather than unique values.
- Knowledge of a causal graph can tighten bounds by exploiting independencies implied by the graph in observational and interventional distributions.
- Sensitivity models quantify unobserved confounding using statistics such as odds ratios and propensity scores.
- The paper amortizes partial identification with a PFN trained on a canonical full-support prior instead of deriving problem-specific bounds for each dataset and query.
B. Proofs
The proofs establish that the canonical prior has full support over discrete SCMs and that the PFN posterior converges to the causal-query distribution among SCMs compatible with the observational law. Consequently, PFN credible sets asymptotically provide frequentist confidence sets for partially identified causal effects.
- Full-support prior: The canonical prior assigns positive probability to every open neighborhood of every SCM in the model space.Positive graph probability and full-support Dirichlet parameter priors combine with continuity to establish full support.
- Full-support prior: The prior factorization separates graph structure from canonical parameters, with every semi-Markov graph receiving positive prior mass.Directed and bidirected edges are independently sampled with probabilities strictly between zero and one.
- Posterior convergence: As n →∞, qω(Ψ | Dn) converges weakly to Π0(Ψ | PM), the prior distribution of the causal effect restricted to SCMs generating the true observational distribution.The observational distribution is point-identified even though multiple SCMs may induce it, leaving compatible-SCM uncertainty for partially identified queries.
- Coverage guarantee: The guarantee concerns conditional-prior recovery under partial identification rather than point recovery of an identified causal functional.The proof combines concentration of the observational parameter with PFN optimality and differs from results imposing identification restrictions.
- Posterior convergence: Within a fixed observational distribution, the posterior over SCMs remains equal to the conditional prior because the likelihood is identical across compatible SCMs.The data update the observational law but do not reweight SCMs that induce the same law.
- Coverage guarantee: PFN-derived Bayesian credible sets are asymptotically valid frequentist confidence sets for the identified set.When the causal effect is point-identified, the identified set collapses to a point and the result reduces to standard frequentist coverage.
C. Additional Experiments
Additional experiments examine posterior concentration and interval performance as context size increases. The evaluations use binary systems, including an instrumental-variable graph, and report coverage, interval width, and inference time.
- Posterior concentration: As context samples increase, posterior distributions for interventional and counterfactual queries approximately concentrate on the identified set.Figure 2 studies P(Yx=1 = 1) and P(Yx=1 = 0, Yx=0 = 0) for a randomly drawn SCM.
- Instrumental-variable experiment: The three-variable experiment uses binary X, Y, and Z under the instrumental-variable graph Z →X →Y, X ↔Y.It evaluates 95 percent credible intervals using observational samples from P(x, y, z).
- Instrumental-variable experiment: The instrumental-variable evaluation measures coverage, average interval width, and inference time across context-sample sizes.Results are averaged over evaluated models and reported in Table 2.
- Analytical-bound comparison: Figure 3 compares predicted posterior distributions for Ψ := P(Yx = 1) at x = 0 and x = 1 with analytical Balke–Pearl lower and upper bounds.The bounds are shown as dashed vertical lines in the three-variable binary SCM experiment.
D. Architecture and Training Details
The appendix provides a complete description of the neural-network architecture and training procedure used in the experiments.
- Implementation description: The appendix documents the neural-network architecture and training procedure used in the experiments.It serves as the implementation description for the experimental model.
D.1. Architecture
The architecture represents context data and causal queries jointly, processes them with alternating attention, and predicts a histogram distribution over the queried causal probability. Its query encoding supports unified interventional and counterfactual prediction but can blur overlapping events.
- Model components: A Neural Process uses a context/query encoder, outcome predictor head, and loss function to model qω(Ψ | D).A single forward path handles both interventional and counterfactual queries through a unified query specification.
- Query representation: Counterfactual queries are encoded as conjunctions of potential-outcome events, recording intervention and outcome nodes, values, and event indices.Queries exceeding the maximum number of events are padded.
- Query representation: Each query becomes a token whose node embeddings sum role, value, and event-index information across events.This preserves whether a node is used as an intervention or outcome variable.
- Attention design: Query nodes attend to context nodes for the same variables, providing an inductive bias and compressing complex queries into a single token.Summing multiple events into shared node slots can blur distinctions when events overlap on the same variables.
- Attention design: Alternating sample attention and node attention processes tokens across samples and variables, with causal masking preventing query-to-query attention.The encoder uses alternating attention blocks with residual connections and normalization.
- Prediction head: Outcome and intervention summaries are concatenated and mapped to B histogram logits over equal-width bins partitioning [0, 1].The predictive distribution is the histogram over Ψ, and the point prediction is its expectation.
- Training objective: Training assigns each ground-truth probability to its enclosing bin and maximizes predicted mass there using cross-entropy.The loss is averaged over all valid queries in the batch.
D.3. Training Procedure
Training uses episodic batches built from randomly drawn canonical structural causal models, with varying context sizes and structured counterfactual queries. The experiments’ model and training hyperparameters are summarized in Table 3.
- Each training step samples a fresh batch of episodes from randomly drawn canonical SCMs.Each episode contains observational context observations and structured queries paired with ground-truth counterfactual probabilities.
- Episodes combine Nc i.i.d. observational observations with Nq structured queries sharing the same context.Queries are paired with ground-truth counterfactual probabilities Ψq(M).
- Context size Nc is sampled uniformly from [Nmin, Nmax], exposing the model to varying amounts of observational evidence.
- Table 3 summarizes the model and training hyperparameters used in the experiments.