Executing Causal Structure Learning with Linear-Attention Transformers
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
| Protocol | Radius attempts | Passing snapshots | Saved 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 |
| Starting policy | Seed | Snapshot | Certified updates | Final verified distance bound | |
|---|---|---|---|---|---|
| Seeded half-radius perturbation | 3 | 5 | 0 | 134 | |
| Seeded half-radius perturbation | 5 | 0 | 0 | 43 | |
| Seeded half-radius perturbation | 10 | 3 | 0 | 978 | |
| Saved matching-control endpoint | 3 | 8 | 0 | 186 | |
| Saved matching-control endpoint | 5 | 9 | 0 | 426 | |
| Saved matching-control endpoint | 10 | 3 | 0 | 546 |
| Model | error | error (zero-shot) | trajectory error, 12 updates |
|---|---|---|---|
| paper-size corpus | |||
| softmax, with memory | |||
| linear, with memory | |||
| linear, no memory (collapsed input) | |||
| learned-threshold unroll | |||
| 10 corpus, 2.5 epochs | |||
| Attention | Residual write-back | error | error | trajectory error |
|---|---|---|---|---|
| softmax (paper-size) | persistent ( ) | |||
| softmax (paper-size) | absent (output from scratch) | |||
| linear (paper-size) | persistent ( ) | |||
| linear (paper-size) | absent (output from scratch) | |||
| softmax (large corpus) | persistent ( ) | |||
| softmax (large corpus) | absent (output from scratch) |