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
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
Figure 1: Examples of DAGs compatible with proximal causal identification assumptions. In every panel, Z and W are informative proxies for the unmeasured confounder U . Direct Z – A edges (in either direction) and a direct W→Y edge are optional.
Figure 2: Synthetic-prior episode: each covariate xi is fixed across K downstream SCM draws. One draw forms the observational context; replicate potential outcomes yield Monte Carlo CATE targets and their sampling variance. varK is the unbiased sample variance of the K paired treatment effects.
Figure 3: ProximalFM architecture. eX , eZ , and eW are learned structural-role embeddings; ey×a encodes the factual outcome according to treatment.
Figure 4: Median CATE nPEHE across benchmark episodes versus context size. Rows show linear and nonlinear mechanisms; columns vary latent confounding and proxy reliability. The dotted line marks the constant-ATE baseline ( nPEHE=1 ); lower is better.
Figure 5: Accuracy and speed at context size n=2048 , shown separately for linear and nonlinear mechanisms. Points show median runtime per episode against median paired CATE nPEHE rank (IQR whiskers). Filled and hollow points respectively include and exclude hyperparameter selection. The orange segment compares ProximalFM on CPU and MPS. More details in Section D.3.1 .
Appendix figures & tables52 assets
Supplementary material from the paper’s appendix.
Appendix
Figure B.1: Detailed diagram of the proximal prior’s per-dataset generation pipeline.
hyperparameter
value
description
Block dimensions
dim_X
Unif{0,…,max_dim_X}
observed baseline covariate dimension
dim_U
triangular on {1,…,max_dim_U}
latent confounder dimension
dim_W
Unif{dU,…,max_dim_W}
outcome-inducing proxy dimension
dim_Z
Unif{dU,…,max_dim_Z}
treatment-inducing proxy dimension
Dimension bounds
Appendix
Table B.1: Structural parameters of the prior: the block dimensions and the dataset shape. Each dimension is drawn per dataset from the law shown, stated in terms of the ceilings listed below it. The proxy dimensions are drawn conditionally on dU so that the completeness heuristic of Section B.2.2 .
hyperparameter
value
description
Mechanism
noise_std
0.001
tiny fixed per-layer dither so continuous columns realise distinct values
Feature corruption
max_categories
50
cap on categories per feature
proxy_cat_cover_prob
0.5
chance a proxy mirrors a U category (vs. a continuous wildcard)
Positivity
Appendix
Table B.2: Fixed hyperparameters of the prior, held constant across every dataset. Listed are the constants that pin down behaviour the surrounding text describes only qualitatively; switches whose shipped value is that behaviour itself are stated in the text instead.
hyperparameter
sampling law
description
Mechanism
num_layers
log-scaled trunc. normal, mean ∼logUnif[1,6] , int ≥2
mechanism MLP depth
hidden_dim
log-scaled trunc. normal, mean ∼logUnif[5,130] , int ≥4
mechanism MLP width
mlp_activations
categorical over 54 random activation functions
activation, drawn from a random library
init_std
log-scaled trunc. normal, mean ∼logUnif[0.01,10] , real ≥0
weight-initialisation scale
block_wise_dropout
categorical {True,False} (random weights)
block-sparse vs. dense weight init
Appendix
Table B.3: Per-dataset sampled hyperparameters of the prior. Each is drawn once per dataset from the listed meta-distribution; log-scaled truncated normals draw their centre log-uniformly over the stated mean range and add the lower bound as an offset.
Figure B.2: Delivered marginal distributions of the block dimensions dU,dW,dZ across the prior, for the shipped (triangular dU ) and the alternative (uniform dU ) sampling, overlaid with their analytic sampling laws (dashed: triangular; dotted: uniform). Bars are empirical frequencies over 1024 datasets. Uniform dU pushes proxy mass against the ceiling dmax=10 ; the triangular law spreads it into the linear ramp of ( 24 ). Delivered bars track the theoretical curves, confirming the generator samples the intended laws.
Figure B.3: Every dial: each input channel’s share of fA ’s (top) and fY ’s (bottom) total input variance , against its own sampled knob ( s for the U -channel, the covariate and proxy keep-probabilities for X~ and Z~A/W~Y , the noise share for ε ). Solid line: moving average. Each knob steers its channel’s delivered variance share.
Figure B.4: Output side of the composition: what each mechanism’s output depends on , measured by a cross-fitted TabICL regression on the delivered data. Columns 1–3 plot each parent channel’s unique leave-one-out ΔR2 (drop in out-of-fold R2 when that block is removed from the predictors) against its own knob, for fA (top, target logit(π) ) and fY (bottom, target Y0 ); column 4 plots the unexplained share 1−R2 against the noise knob. Solid line: moving average. Each channel’s dependence rises with its knob, as at the input level. The U -channel reads modestly because leave-one-out credits only signal unique to U , not the part shared with its child proxy.
Proxy profile P
Heuristic?
Reason
[num,5-cat,3-cat]
ok
exact match, column for column
[num,8-cat,4-cat]
ok
categorical columns carry more categories than required
[num,num,3-cat]
ok
a numerical column covers U ’s 5 -cat column
[num,num,num]
ok
numerical columns dominate every requirement
[num,num]
no
too few columns ( dP=3<dU=4 )
[num,4-cat,3-cat]
no
the 4 -cat column cannot cover U ’s 5 -cat column
Appendix
Table B.4: Worked examples of the categorical-completeness heuristic for the fixed confounder profile U=[num,5-cat,3-cat] (one numerical column and two categorical columns, with 5 and 3 categories). A proxy block P satisfies the heuristic when its columns can be matched one-to-one to those of U so that each is at least as informative: a numerical column dominates anything, and a categorical column dominates a categorical requirement only if it has at least as many categories. A numerical proxy column may thus cover a categorical requirement, but a categorical proxy column can never cover a numerical one.
Figure B.5: Joint distribution of the latent dimension dU against each feature block ( dX , dW , dZ ), one row per prior variant ( triangular , uniform ). Colour encodes the number of delivered datasets falling in each (dU,⋅) cell. The diagonal line on the dW / dZ panels marks the completeness heuristic dW,dZ≥dU ; both variants place their mass on or below it, i.e. constant-column filtering does not measurably break the heuristic imposed at sampling time. dX carries no such constraint and is shown for reference.
Figure B.6: Column-type composition of each block ( U , W , Z ), one 100% -stacked bar per variant ( default : categorical completeness enforced; no_enforcement : proxies left unconstrained). Each bar partitions a block’s columns into continuous (grey, “cont.”) and categorical, the latter binned by category count ( 2 – 5 singletons, then coarser 6 – 10 , 11 – 20 , 21 – dcatmax tail bins). Continuous columns dominate every block; contrasting W / Z across the two variants shows how much enforcement pushes proxy columns categorical relative to the unconstrained no_enforcement arm.
Figure B.7: How enforcement makes proxy structure track U ’s categorical demand, pooling W and Z across the default and no_enforcement variants; shaded bands are ± one SEM. (a, b) condition on a single U column’s own cardinality: the x -axis bins a U column by its category count (continuous “ num ”, then 2,…,9 , and coarser tail bins up to “ 30+ ”), and each point averages a statistic of the co-occurring proxy block over every dataset containing a U column of that kind – (a) the fraction of proxy columns that are continuous, essentially flat in U ’s cardinality but higher under default , and (b) the mean category count of the proxy’s categorical columns, which rises with U ’s cardinality under default but stays low under no_enforcement . (c) moves to the dataset level: the mean fraction of categorical proxy columns against U ’s total categorical demand ∣cats(U)∣ , with bins of fewer than 5 datasets dropped. Both arms rise (the per-dataset categorisation propensity is shared across blocks) and default sits below no_enforcement : enforcement deepens the proxy’s categorical columns rather than multiplying them. Because proxy statistics in (a, b) are block-level rather than tied to a matched column, those trends reflect association with “the dataset contains a U column of this kind,” not a one-to-one mapping.
Figure B.8: The treatment-assignment pipeline of ( 25 )–( 29 ) on one representative dataset, tracking a single unit (red line) across all five stages. (a) The raw logit ℓ=fA(⋅) , whose scale and offset are an uncontrolled by-product of the sampled architecture. (b) The standardised ℓ~ of ( 26 ), pinned to zero mean and unit variance. (c) The affine ℓ′=sℓ~+b of ( 27 ), shifted and sharpened by the two sampled overlap knobs. (d) The propensity π=σ(ℓ′) of ( 28 ), confined to [ϵ,1−ϵ] by the clamp of ( 29 ) (dotted bounds, drawn at an enlarged ϵ for visibility; production ϵ=10−3 ). (e) The realised treatment A∼Bernoulli(π) across the batch (fraction in each arm), with the tracked unit’s realised arm marked in red.
Figure B.9: The heterogeneity dial on the delivered prior, over 1024 datasets. All statistics are ratios, since each SCM’s outcome carries its own arbitrary scale. (a) Two datasets’ effect densities on a shared outcome-scaled axis, each centred on its own τˉ (dashed line), one gathered ( γ=0.09 ) and one spread ( γ=0.91 ); the pair is chosen with matched raw dispersion c=sd(τraw)/sd(Y) ( 1.127 vs. 1.125 ), so the contrast is attributable to γ rather than to SCM variation. (b) The outcome coupling corr(Y0,Y1) against γ (solid line: moving average). At γ=0 the effect is the constant τˉ , so Y1=Y0+τˉ and the coupling is exactly 1 (dashed reference). (c) Distribution of the delivered dispersion sd(τ)/sd(Y) per γ bin (Spearman 0.78 overall), against the proportionality γ⋅median(c) implied by the collapse; within-bin spread is dominated by variation in c . Because the reference line is anchored at zero and uses the overall median c for its slope, the panel verifies strict proportionality to γ —linear scaling starting from the origin—rather than predicting the absolute magnitude of dispersion.
Figure B.10: Gallery of feature marginals: one column per block ( X , U , W , Z , the propensity π=P(A=1∣⋅) for A , both potential outcomes for Y ), one row per randomly sampled dataset. Each panel is one feature column’s histogram, solid for continuous columns and dashed for categorical ones; the Y column overlays Y0 (blue) and Y1 (red). Feature blocks share the standardised axis Corrupt leaves; π∈[0,1] and Y keep their natural scale.
Figure B.11: Within-block dependence: one column per block, one randomly sampled dataset per row, each panel a scatter of two features from that block. For X , U , W , Z the two features are random columns of the block (sampled from datasets of width ≥2 ); for A the panel plots the propensity π against the treatment draw A ; for Y the two potential outcomes Y0 vs Y1 . Rows are subsampled for legibility; feature axes inherit the standardised scale.
Figure B.12: Cross-block dependence: one column per chosen block pair, one sampled dataset per row, each panel a scatter of one feature from each block. U – W and U – Z are the parent → child edges; W – Z is the proxy pair, associated only through the shared latent U (no direct edge); U – π is the confounder against the propensity π=P(A=1∣⋅) . Rows are subsampled for legibility; feature axes inherit the standardised scale.
Figure B.13: Delivered overlap across the prior, from the stored true propensity π=Pr[A=1∣X,U] (no estimation involved). (a) Pooled distribution of the overlap margin min(π,1−π) over all rows of 1024 datasets, on a logarithmic axis: the margin measures each unit’s distance from a degenerate assignment, folding both tails into one axis, and the log scale is what makes the mass piled on the clamp ε=10−3 (dotted) visible at all. Bars are per-bin shares of all rows on a logarithmic grid, so bin width grows to the right. (b) Per-dataset quantile function of π , summarised across datasets by the median curve with the 25 – 75% and 10 – 90% inter-dataset bands. Dashed lines mark the practical-overlap band [0.05,0.95] in both panels.
Figure B.14: Exclusion-restriction audit ( Assumption 2 ) over 1024 datasets. Each panel is one clause – (a) Z excluded from fY , (b) W excluded from fA – and each column a cross-fitted incremental ΔR2 of the excluded block: conditioning on U is the check (expected ≈0 ), control: U dropped the positive control (expected clearly positive). Colour and horizontal position both encode u_dependence=R2(withU)−R2(withoutU) , i.e. how much confounding that draw carries (linear ramp, floored at 0 , capped at the 95 th percentile). Both checks stay on zero regardless of colour; both controls fan upward with warmer colour, most visibly in the top tercile of u_dependence ( Table B.5 ).
conditioning set
role
datasets
median ΔR2
90th pct
share >0.02
Z excluded from fY
{A,X,U,Z}
conditions on U
all
-0.002
0.007
0.021
{A,X,U,Z}
conditions on U
top confounded
-0.005
0.004
0.012
{A,X,Z}
positive control
all
0.004
0.099
0.313
{A,X,Z}
positive control
top confounded
0.029
0.222
0.598
W excluded from fA
Appendix
Table B.5: Exclusion-restriction audit ( Assumption 2 ): cross-fitted incremental ΔR2 of the excluded block over the stated conditioning set, for the check (conditions on U ) and the positive control ( U dropped). Top confounded is the third of datasets with the largest u_dependence ( R2 with U minus R2 without it).
Figure B.15: Confounding-strength comparison ( Table B.6 ) over 1024 datasets. Each violin shows normalised absolute ATE error ∣ATE−ATE∣/sd(Yobs) on a log axis, scored against mean(Y1−Y0) : naive (no adjustment), backdoor X and (X,W,Z) (T-learners on measured covariates), and oracle (X,U) and (X,U,W) (T-learners also given the latent confounder).
Figure B.16: Backdoor- X and oracle- (X,U) error over the delivered confounding-share grid: 1024 datasets binned 3×3 on the quantiles of the delivered U share of fA (rows) and fY (columns) ( Figure B.3 ). Cells report the median normalised error on a shared colour scale. The backdoor panel’s top row (delivered U share of fA above 0.495 ) is elevated at 0.045 – 0.046 against 0.028 – 0.041 in the two rows below, and is nearly flat across fY ’s share. The oracle panel shows no comparable pattern.
rung
adjustment set
role
median error
90th pct
share >0.1
naive
none
unadjusted
0.041
0.195
0.241
backdoor X
X
correct if no latent confounder
0.037
0.163
0.211
backdoor XWZ
(X,W,Z)
DAG-blind: everything measured
0.030
0.121
0.148
oracle XU
(X,U)
identification: minimal sufficient set
0.028
0.104
0.108
oracle XUW
(X,U,W)
efficiency: adds W , a parent of fY
0.025
0.105
0.107
Appendix
Table B.6: ATE error across adjustment sets. Errors are absolute deviations from the in-sample mean(Y1−Y0) , divided by sd(Yobs) . All adjusted estimates use cross-fitted T-learners; the oracle rows include the latent U .
Figure B.17: Proxy relevance measured by the cross-fitted increment ΔRP2(uj)=R2(uj∣X,P)−R2(uj∣X) , shown separately for W and Z . (a) Distribution of the mean and weakest-column ΔR2 over U ; boxes show the IQR and whiskers the 10 th– 90 th percentiles. (b) Median ΔR2 by proxy dimension slack dP−dU , with slack ≥7 pooled. (c) Median ΔR2 by sampled noise share of fP , in six equal-frequency bins per proxy. Bands in (b) and (c) show the IQR.
aggregation over U
median ΔR2
share <0.1
ρdP
ρdU
ρdP−dU
ρ noise
W
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
Z
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
Appendix
Table B.7: Proxy-relevance audit ( Assumption 4 ): cross-fitted incremental ΔR2=R2(uj∣X,P)−R2(uj∣X) of each proxy P for the latent columns, aggregated over U by the mean or by the weakest column. The counting heuristic ( Section B.2.4 ) passes for every dataset, so the share column is the rate at which it and functional recoverability disagree. ρ is Spearman’s correlation with each candidate driver: the proxy’s own width dP , dU , the slack dP−dU that conflates the two, and the sampled noise share of the fP input.
Stage
Hyperparameter
Value
Shared
Embedding dimension E
128
Feedforward expansion factor
2
Activation
GELU
Normalization
pre-norm
Dropout
0.0
Column embedder
Blocks
3
Appendix
Table C.1: Architectural hyperparameters of ProximalFM . The three stages are those of TabICLv2, kept at its values so that its pre-trained weights initialize every component we did not modify; the output head is ours. Parameter counts are per stage, the in-context learner’s including the head.
Stage
Hyperparameter
Value
Prior
Confounder dimension dU
1 – 10
Treatment proxy dimension dZ
1 – 10
Outcome proxy dimension dW
1 – 10
Covariate dimension dX
0 – 10
Context units N
64 – 2048 (uniform)
Query units nq
128
Appendix
Table C.2: Training configuration of ProximalFM : the prior the model is trained on, the objective and its smoothing schedule, the optimizer, the pre-trained backbone it is initialized from, and the validation sets used for checkpoint selection.
Figure C.1: Training curves for ProximalFM . (a) The objective and its two terms against log step, with the σ -annealing window shaded; the dashed line is the loss of a uniform density over the CATE grid, matching the value at initialisation. (b) The same curves on a linear step axis after the anneal. (c) The σ and learning-rate schedules. (d) A magnified view of the weighted mean term, λmeanLmse , with the σ -annealing window again shaded. Stochastic curves show the median and interquartile band within log-spaced step bins in panels (a) and (d), and linearly spaced bins in panel (b); the deterministic schedules in panel (c) are shown without binning.
Hours
%
Compute
Prior generation
43.46
65.7
Forward / backward
18.12
27.4
Validation, prior set (376 passes)
1.27
1.9
Validation, realistic set (751 passes)
2.34
3.5
Checkpointing, logging, startup
0.96
1.5
Total
GPU-hours ( 1× H100 NVL)
66.15
100.0
Appendix
Table C.3: Cost of the ProximalFM training run, as measured GPU-hours. Validation is the two frozen evaluation datasets, scored at their logging intervals throughout training.
Figure C.2: Evolution of point-estimate performance during training on the fixed realistic-mechanism validation suite ( Section C.4.1 ). Curves show the validation metric at each evaluation step. The shaded regions in panels (b)–(d) indicate the interquartile range across validation episodes. (a) ATE RMSE, a pooled metric and therefore shown without an episode-level band. (b) normalized ATE bias, where zero denotes no bias. (c) CATE nPEHE, with the dashed line marking the constant-ATE baseline ( nPEHE=1 ). (d) CATE shape R2 , with the dashed line marking the corresponding constant-ATE baseline ( Rshape2=0 ). Lower values are preferred in panels (a) and (c), while panel (b) is centered at zero and higher values are preferred in panel (d). For readability, curves and shaded bands are lightly smoothed by taking medians within log-spaced training-step bins.
Figure C.3: Evolution of distributional predictive performance during training on the fixed realistic-mechanism validation suite ( Section C.4.1 ). (a) CRPS, for which lower values indicate better probabilistic predictions. (b) predictive standard deviation divided by RMSE, with the dashed line at one denoting calibrated predictive spread. (c) empirical 90% and 50% coverages, with the dashed lines marking the nominal coverage levels and the shaded regions showing the interquartile ranges across validation episodes. (d) widths of the central 90% and 50% credible intervals, pooled over validation rows. Panels (b) and (d) use logarithmic y-axes. For readability, curves are lightly smoothed by taking medians within log-spaced training-step bins.
Figure C.4: Illustration of the spiky posterior-density representation for one selected query point from the validation suite. (a) Raw posterior density induced by the model’s histogram output. (b) Corresponding raw CDF. (c) CDF after Savitzky–Golay smoothing with a 50-point window and polynomial degree three. (d) Density obtained by differentiating the smoothed CDF. The smoothed CDF closely follows the raw CDF while yielding a more visually interpretable density. The dashed line marks the true CATE. Smoothing is used for visualization only; all validation metrics are computed from the original posterior representation.
Figure C.5: Posterior-width contraction during training on the fixed realistic-mechanism validation suite ( Section C.4.1 ). Each cell reports the mean predictive posterior standard deviation averaged over all validation query rows, for a given training checkpoint and context size. Widths are expressed in standardized units and use the piecewise-uniform posterior variance. The shared colour scale is logarithmic.
Figure C.6: Evolution of the predictive CATE posterior for one fixed query point from the realistic-mechanism validation suite ( Section C.4.1 ). Curves show the density obtained from the posterior CDF at four training checkpoints, with training steps indicated by colour. The same query point is evaluated with (a) nctx=64 and (b) nctx=1024 . The dashed line marks the true CATE. Densities are shown in standardized units. The CDFs are smoothed for visualization, following Section C.4.4 .
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 (X,W,Z) , then 16-neighbour proxy marginalisation.
Section D.1.1
S-learner (TabICL) + kernel marg.
Backdoor meta-learner
S-TabICL fit on (X,W,Z) , 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
Appendix
Table D.1: Catalogue of CATE predictors selected in the main benchmark.
Component
Setting
Value
Cross-fitting
Folds (stratified on A )
4
Bridge class
Kernel
RBF, blockwise median heuristic
Nyström rank r
256
Bridge regularization
λh,λq
sweep : {10−6,10−5,10−4,10−3,10−2,10−1}
Selection criterion
Held-out proximal moment loss
Validation fraction
0.25
Appendix
Table D.2: P-learner settings. Bridge regularization parameters are selected separately within each cross-fitting fold using held-out moment loss. Final-stage hyperparameters are selected per episode using four-fold cross-validated prediction MSE.
Component
Setting
Value
Bridge class
MLP, ReLU, linear output
Hidden width
searched : {16,32,64}
Depth (hidden layers)
searched : {2,4,6}
Moment kernel
RBF over (A,Z/σ^Z,X/σ^X)
Length scale ℓ
fixed at 1.0 or median heuristic
Loss
U- or V-statistic ( 50 )
one method each
Appendix
Table D.3: NMMR settings. Values marked searched are selected per episode by the procedure described below; the remainder are fixed across all episodes. The U- and V-statistic variants are reported as separate methods, so the loss is fixed within a run rather than searched.
Component
Setting
Value
Kernels
Bandwidths baseline
per-column median heuristic
Kernel bandwidth multiplier (W,X)
searched : {0.8,1,1.2}
Kernel bandwidth multiplier (Z)
searched : {0.8,1,1.2}
Regularisation
λ (RKHS norm)
swept : {10−5,10−4,10−3,10−2,10−1}
Cholesky jitter
10−8
Prediction
W strategy
“ind. W”, “cond. W”
Appendix
Table D.4: PMMR settings. Bandwidth multipliers are sampled ; the regularisers are swept , i.e. every combination is evaluated for each bandwidth draw. Both are selected per episode on the held-out V-statistic.
Component
Setting
Value
Data split
m1 (stage-1 rows)
ncontext
m2 (stage-2 rows)
ncontext
Kernels
Kernel family
Column-wise Gaussian RBF for X,W,Z ; binary for A
Base bandwidths
Per-column median heuristic, estimated on the context
Shared multiplier c on σd2
Shared across X,W,Z ; searched: {0.8,1.0,1.2}
Regularization
λ1,λ2
searched by LOOCV; 10-point log grid in [10−5,10]
Appendix
Table D.5: KPV settings. Both stages use all context rows. The regularizers λ1 and λ2 are selected within each fit by their closed-form leave-one-out criteria. The shared kernel-scale multiplier is selected per episode using the stage-2 LOOCV loss.
Figure D.1: Construction of the generated proximal benchmarks. (a) A shared real source row supplies dependent observed covariates and hidden source features. Independent Gaussian innovation augments the hidden features. Both X and U enter the proxy measurements and the treatment and outcome mechanisms. (b) Independent Gaussian X and U define the controlled synthetic study. The joint proxy information parameter changes the measurement channels while preserving the causal variables within each dimension. The diagrams display the generated causal structure, with independent disturbances suppressed where indicated.
Source
Rows
Features
pol
10082
26
MiniBooNE
72998
50
default-of-credit-card-clients
13272
20
Higgs
940160
24
jannis
57580
54
heloc
10000
22
Appendix
Table D.6: Original sizes of the twelve selected datasets. Feature counts exclude the original prediction target, although that column remains eligible for selection as a source variable in our construction.
Figure D.2: Learning curves and error decomposition for the nonlinear, high-confounding, low-proxy-reliability setting (36 episodes: 12 sources with three realizations each). (a) Median CATE nPEHE versus context size; the dotted line marks the constant-ATE benchmark ( nPEHE=1 ). ProximalFM-FT continues training the main checkpoint on the same prior and objective with contexts up to 5000 rows, updating only part of the in-context predictor ( Section C.3.1 ). (b) For query-row errors ei=τ^i−τi , their mean eˉ , and true CATE variance sτ2=Vari(τi) , the episode-level decomposition is nPEHE2=eˉ2/sτ2+Vari(ei)/sτ2 . Panel (b) plots the first term horizontally (squared bias in the query-average treatment effect) and the second vertically (heterogeneity error after removing that bias). Dotted contours show constant x+y ; plotted coordinates are separate component medians, so they need not reproduce panel (a)’s median nPEHE. Arrows and larger markers indicate increasing context size ( 128 – 8192 ). ProximalFM-FT moves mainly leftward: post-training improves the average effect estimate, with little improvement in CATE shape.
Figure D.3: Posterior calibration and contraction. The top row evaluates 400 datasets drawn from the pre-training prior; the bottom row evaluates 120 semi-synthetic datasets with nonlinear mechanisms, high confounding, and low proxy reliability. In every episode, the same 256 query points are evaluated as the context grows through nested prefixes from 128 to 8192 observations. (a,d) Coverage by context: each solid curve gives empirical coverage of a central CATE credible interval at one nominal level (1%, 5%, 10%, …, 95%, 99%); the faint dotted line of the same colour marks that level’s nominal coverage. (b,e) Calibration: each curve fixes a context size, coloured from purple ( 128 ) to yellow ( 8192 ), and plots empirical against nominal coverage. The dotted diagonal denotes perfect calibration; curves below it indicate undercoverage. (c,f) Posterior contraction: the black curve is the mean predicted CATE posterior standard deviation across query rows. The grey band is the interquartile range of episode-level mean standard deviations, showing variation between datasets. These panels use logarithmic vertical axes. Coverage stays close to nominal on the prior, whereas the semi-synthetic intervals increasingly undercover as context grows, even though their predicted uncertainty contracts.
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
Appendix
Table D.7: Exhaustive semisynthetic CATE nPEHE results at nctx=128 . Each entry first averages the three replicates of each source dataset, then averages equally over the 12 sources; lower is better. The lowest mean in each column is bold. Other underlined methods are not significantly worse than that empirical best under one-sided exact paired sign-flip tests across source-level means, with Holm correction within the column ( α=0.05 ).
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
Appendix
Table D.8: Exhaustive semisynthetic CATE nPEHE results at nctx=512 . Each entry first averages the three replicates of each source dataset, then averages equally over the 12 sources; lower is better. The lowest mean in each column is bold. Other underlined methods are not significantly worse than that empirical best under one-sided exact paired sign-flip tests across source-level means, with Holm correction within the column ( α=0.05 ).
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
Appendix
Table D.9: Exhaustive semisynthetic CATE nPEHE results at nctx=1,024 . Each entry first averages the three replicates of each source dataset, then averages equally over the 12 sources; lower is better. The lowest mean in each column is bold. Other underlined methods are not significantly worse than that empirical best under one-sided exact paired sign-flip tests across source-level means, with Holm correction within the column ( α=0.05 ).
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
Appendix
Table D.10: Exhaustive semisynthetic CATE nPEHE results at nctx=2,048 . Each entry first averages the three replicates of each source dataset, then averages equally over the 12 sources; lower is better. The lowest mean in each column is bold. Other underlined methods are not significantly worse than that empirical best under one-sided exact paired sign-flip tests across source-level means, with Holm correction within the column ( α=0.05 ).
Method
Context size nctx
128
256
512
1,024
2,048
4,096
8,192
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
Appendix
Table D.11: Extended-context semisynthetic CATE nPEHE results for nonlinear mechanisms with high confounding and low proxy reliability. Each entry first averages the 3 replicates of each source dataset, then averages equally over the 12 sources; lower is better. The lowest mean in each column is bold. Other underlined methods are not significantly worse than that empirical best under one-sided exact paired sign-flip tests across source-level means, with Holm correction within the column ( α=0.05 ).
Figure D.4: CATE RMSE against joint proxy information, measured by the population R2 for predicting U from (W,Z) . Columns vary proxy dimension d ; rows vary context size. All methods are evaluated on the same 256 held-out query rows at each setting. The vertical axis is logarithmic. Within each dimension, the underlying draws are shared across information levels; results come from one realization per setting.
Figure E.1: The Light Tunnel Mk2 and a simplified diagram with the variables relevant to our setup.
Figure E.2: Physical and software pathways in the optical experiment. Solid arrows represent the predetermined physical pathways and dashed green arrows represent effects resulting from the software control functions, i.e., software mappings . Independent Gaussian noise enters the two proxy mappings. The lower sequence shows acquisition of the factual outcome and the two intervention measurements. The graph describes the intended within-observation structure. Its causal interpretation requires stable device state and the proxy restrictions discussed in the text.
Figure E.3: Software control functions for the physical experiment. Panels (a)–(c) show the deterministic mappings before addition of proxy noise and final actuator clipping. Panel (d) shows the treatment probabilities for Bernoulli assignment and the corresponding hard-threshold protocols. The plots display the specified control functions over illustrative sensor ranges.
Figure E.4: Absolute error of estimated average treatment effects in the Causal Chambers Light Tunnel experiment as the factual context grows. Each panel is one physical operating condition, crossing control-function family (linear or non-linear), confounding strength, and nominal proxy informativeness. Curves show ∣θn−θref,n∣ in raw infrared-sensor counts, where both the estimate and the paired-intervention reference are evaluated on the same nested context prefix of size n . Both axes use logarithmic scales; lower is better.
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.
Yonghan Jung
School of Information Sciences University of Illinois Urbana-Champaign
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.
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 F1, +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.
Patrick Blöbaum, Krishnakumar Balasubramanian, Shiva Prasad Kasiviswanathan
Amazon Web Services · Department of Statistics, University of California, Davis