cs.LGOct 7, 2026

Executing Causal Structure Learning with Linear-Attention Transformers

Authors: Amartya Roy, Sayar Karmakar

Organizations: School of Interdisciplinary Research Indian Institute of Technology Delhi · University of Florida

Abstract

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

Appendix figures & tables7 assets

Supplementary material from the paper’s appendix.

Appendix

Explore similar work

May 8, 2026cs.LG

Arrow: A Foundation Model for Causal Discovery

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.
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.
May 8, 2026cs.LG

Mask2Cause: Causal Discovery via Adjacency Constrained Causal Attention

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\textbf{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.