Attention-Based Pretraining for Unsupervised Amortized Causal Discovery
Abstract
Directed Acyclic Graphs (DAGs) are foundational to many areas of AI research, providing a formal language for causal reasoning and model interpretability. However, learning DAGs remains challenging, due to super-exponential computational cost and diminished accuracy in small sample regimes. To address these challenges, we introduce Attention-DAG (ADAG), a novel linear transformer architecture for unsupervised amortized DAG learning. Unlike traditional unsupervised DAG learning methods which recover the graph structure from data of each individual domain, ADAG leverages the knowledge from multiple domains. It provides a nonlinear mapping from observational data of each domain to both the corresponding graph structure and the underlying parameters of Structural Equation Models (SEMs). This enables efficient zero-shot inference of causal structures in new domains with unseen SEM parameters. Remarkably, we demonstrate that for DAG learning problem attention acts as a fixed-point iterative scheme, so the training on multiple domains effectively discovers an efficient solver for the constraint optimization problem. Extensive evaluations on synthetic and realistic benchmarks demonstrate that ADAG significantly outperforms existing baselines in accuracy and efficiency, particularly when data is scarce.