Transformers can execute algorithms on data given in their input. We ask whether they can do the same for causal discovery. We study a standard continuous method that repeatedly updates a candidate causal graph while enforcing acyclicity. We explicitly construct a fixed-weight transformer whose forward pass exactly reproduces one update of this method, so repeated blocks reproduce its optimization trajectory. The transformer carries the current graph and the algorithm's multiplier between updates. We show that retaining the multiplier is essential for exact execution, since different multiplier values can lead to different next updates. We also give conditions under which, within a fixed stage, the number of updates needed to reach a target accuracy can be computed in advance and rounding errors stay bounded as depth grows. Experiments show that the constructed block agrees with a reference update to floating-point precision, while arithmetic replay on synthetic data and seven published benchmark network topologies inherits the reference solver's successes and failures. This separates accurate algorithm execution from accurate causal recovery. In contrast, the ordinary attention models tested under our training budgets do not reliably execute the update or transfer to larger graphs. Whether gradient training can learn an executor in the architecture class of the construction remains open.
Figures & tables
Figure 1: One update, run two ways. The solver carries the edge matrix Wℓ and multiplier αℓ ; the fixed-weight transformer carries the same state in Zℓ . With identical controls, the two transitions agree exactly (Theorem 2 ). The surfaces are schematic.
Figure 2: One update, two implementations. Left: one Poly-prox iteration. Right: the same iteration as the fixed-weight transformer carries it out on the encoded state Zℓ of ( 5 ), line for line; blue is linear attention, orange bilinear/ReLU, green persistent state. Theorem 2 shows both return the same (Wℓ+1,αℓ+1) .
Figure 3: Benchmark topologies (top: ASIA, CANCER, EARTHQUAKE, SURVEY; bottom: SACHS, CHILD, ALARM). Median-SHD draw at m=104 , SHD beside each name; red dots are true edges; the column labelled “transformer” shows the direct arithmetic replay of the block. Last column, thresholded at 0.3 : green kept by both, red missed by both, gold added by both; the black disagreement category is empty on every row.
Figure 4: What the trained models produce. Red dots mark true edges. (1) A softmax model trained at p=5 and applied at p=10 after 12 updates. (2) At p=5 , a collapsed-input linear model withholds W,α , while a softmax model receives them; this illustration changes both architecture and state access. Table 4 separately tests write-back with identical inputs and architectures. (3) Reference-solver updates 1 , 5 and 40 . The right column is the solver’s returned estimate. Colour shows ∣W∣ , scaled within each panel; the rows use two illustrative problem instances.
Figure 5: How close small trained transformers get to the exact update. One-step error is normalized so copying scores 1 (dashed line). Circles: trained and tested at p=5 ; squares: transfer to p=10 . Top: smaller- and larger-corpus one-step errors, then smaller-corpus rollout distances. Bottom: identical state inputs with residual write-back or prediction from scratch; the latter stays above copying. Five training runs are attempted per condition; means and standard errors use finite run summaries (failure counts in Tables 3 and 4 ).
Appendix figures & tables7 assets
Supplementary material from the paper’s appendix.
Appendix
Protocol
Radius attempts
Passing snapshots
p=3
p=5
p=10
Saved V inside
Original centres, recorded γ
270
0/90
0
0
0
0
Polished centres, recorded γ
725
46/90
21
12
13
25
Polished centres, new γ
636
67/90
30
20
17
27
Appendix
Table 1: Verified interval fixed-stage tests. Each dimension has 30 snapshots from ten SEM instances; snapshots and radius attempts are not independent experimental units. The polished policies retain one passing ball per successful snapshot. The last column counts saved matching-control endpoints in that policy’s passing ball, not automatically successful continuations of the original adaptive controller. Only 2 and 4 original snapshot states, respectively, lie in the polished passing balls.
Starting policy
p
Seed
Snapshot
Certified updates
Final verified distance bound
Seeded half-radius perturbation
3
5
0
134
7.41⋅10−16
Seeded half-radius perturbation
5
0
0
43
6.54⋅10−14
Seeded half-radius perturbation
10
3
0
978
1.83⋅10−12
Saved matching-control endpoint
3
8
0
186
3.43⋅10−12
Saved matching-control endpoint
5
9
0
426
5.97⋅10−11
Saved matching-control endpoint
10
3
0
546
2.51⋅10−11
Appendix
Table 2: Repeated execution of the constructed block in serial float64 arithmetic. All use the recorded step size and b=0 , with controls fixed throughout. Within each starting policy and dimension, the selector prefers the recorded-step policy and selects the smallest positive predicted update count, up to 5000 . All six trajectories remain finite, stay inside their verified balls, retain exactly clean scratch/immutable coordinates, and meet the 10−6 fixed-point distance target by the predicted update count. Final bounds use exact rational a posteriori residuals divided by 1−q , or the verified triangle bound through the centre, whichever is smaller.
Figure 6: Accuracy of graph recovery. Lower structural Hamming distance (SHD) means fewer edge mistakes. Each bar averages 36 simulations with 1000 observations; error bars show standard errors. On raw data, the continuous solvers and sort-regress outperform PC and GES. Rescaling each variable to unit variance makes sort-regress the worst method at 20 variables. PC and GES are unchanged by this rescaling.
Model
p=5 error
p=10 error (zero-shot)
trajectory error, 12 updates
paper-size corpus
softmax, with memory
0.94±0.00
0.97±0.00
0.64±0.01
linear, with memory
0.98±0.01
1.01±0.01
24.24±22.50
linear, no memory (collapsed input)
10.48±0.04
16.48±0.05
1.1×103±0.7×103
learned-threshold unroll
3.94±0.06
11.81±0.25
2.8×102±1.7×102
∼ 10 × corpus, 2.5 × epochs
Appendix
Table 3: Trained update executors, mean ± standard error over five attempted training runs. Step error is normalized so that copying the input scores 1 ; the median over test cases is taken within a seed. A bracketed count gives the seeds whose value was non-finite, which are excluded from that mean; trajectory means are dominated by a few divergent seeds.
Attention
Residual write-back
p=5 error
p=10 error
trajectory error
softmax (paper-size)
persistent ( W=W+δ )
0.94±0.00
0.97±0.00
0.64±0.01
softmax (paper-size)
absent (output from scratch)
10.98±0.38
17.01±0.39
1.12±0.12
linear (paper-size)
persistent ( W=W+δ )
0.98±0.01
1.01±0.01
24.24±22.50
linear (paper-size)
absent (output from scratch)
10.44±0.05
16.49±0.03
1.37±0.22
softmax (large corpus)
persistent ( W=W+δ )
0.67±0.01
1.56±0.02
0.40±0.01
softmax (large corpus)
absent (output from scratch)
1.70±0.13
14.82±0.21
0.65±0.03
Appendix
Table 4: Residual write-back with identical inputs, five training runs. “Persistent” writes the update back into the state registers, W=W+δ ; “absent” gives the model the same inputs and the same parameter count but produces W and α from scratch. One-step errors are normalized so that copying scores 1 ; the trajectory column is the unnormalized ∥W12−W12∥F after twelve unrolled updates. A bracketed count gives the seeds whose value was non-finite, excluded from that mean. These are the numbers quoted in Section 8 .
Figure 7: Arithmetic replay and conditioning checks. (1) One-update replay error across eight states per dimension. (2) Effective-multiplier growth in the two-node stationary family. (3) Relative gradient perturbations remain below the bound of Corollary 5.4 ; fixed absolute perturbations are amplified as c grows. (4) Undamped multiplier error accumulates as ℓ⋅10−3 while W=0 .
Figure 8: Stationarity along saved accepted primal outputs, using each update’s pre-controller controls. At each of 101 normalized progress fractions f , trajectory j contributes its saved index ⌊f(nj−1)⌋ . Lines are medians and bands are minima/maxima over ten input-seed trajectories per dimension, not confidence intervals. The horizontal axes do not align outer-stage boundaries. The plotted first point is the first primal output; at the separately recorded initial W=0 , h=∇h=c=0 and ηmin=dM . Finite control changes can increase the augmented residual as the constraint and its gradient decrease.
We introduce Arrow, a foundation model for zero-shot causal discovery on observational tabular data. Arrow factorizes a directed acyclic graph into an undirected skeleton and a topological order, guaranteeing acyclicity by construction. Given a new dataset, it uses a transformer-based architecture to contextualize variables within and across observations, then predicts skeleton edge probabilities and node order scores that together define a graph. Arrow is trained in a supervised fashion on synthetic datasets with ground-truth graphs, using an end-to-end differentiable directed edge composite likelihood induced by the skeleton-order factorization. The training distribution spans diverse graph families, functional forms, noise models, and dataset shapes. Across in- and out-of-distribution synthetic, semi-synthetic, and real datasets, Arrow matches or outperforms existing causal discovery methods at substantially lower inference cost than competitive alternatives. Our results demonstrate that large-scale pretraining on diverse synthetic data can yield zero-shot causal discovery models that are fast, accurate, and reusable on new datasets.
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
Leveraging deep learning for causal discovery in time series remains challenging because existing neural methods predominantly rely on component-wise architectures that fail to capture shared system dynamics or employ decoupled post-hoc graph extraction that risks overfitting to spurious correlations. We propose Mask2Cause, an end-to-end framework that recovers the underlying causal graph directly during the forecasting forward pass. Our approach introduces an Inverted Variable Embedding and an Adjacency-Constrained Masked Attention mechanism, trained with homoscedastic or heteroscedastic objectives to capture causal influences in both mean and variance. Empirical results on diverse benchmarks, from synthetic chaotic dynamics to realistic biological simulations, demonstrate state-of-the-art causal discovery with significantly reduced parameter complexity compared to standard baselines. We further show that inferred causal structures can be used to reduce parameter count of forecasting models by more than 70% on average while maintaining predictive accuracy.
Omar Muhammad, Pasupuleti Dhruv Shivkant, Deepak N. Subramani