How Transformers Learn Causal Structure with Gradient Descent
Eshaan NichaniAlex DamianJason D. Lee
Proves that gradient descent enables two-layer transformers to learn latent causal graphs and form induction heads by showing that the attention matrix gradients naturally track token-level mutual information.
Modern sequence modeling relies heavily on transformer architectures, which exhibit an exceptional capacity for in-context learning—the ability to adapt to new tasks and infer relationships purely from prompt context without updating internal model parameters. While empirical studies show that trained models form specialized internal wiring such as induction heads to capture sequence dependencies, foundational understanding has lagged regarding why and how standard training algorithms actually discover these underlying causal mechanisms from scratch.
To bridge this gap, the article investigates the training dynamics of autoregressive transformers trained with gradient descent on synthetic sequence tasks governed by hidden causal structures. The primary objective is to prove mathematically and demonstrate empirically that standard gradient descent naturally forces a two-layer transformer to uncover latent causal graphs and execute accurate in-context predictions.
The authors analyze an in-context learning task where sequences are generated from latent, directed acyclic graphs over discrete alphabets, encompassing structures such as Markov chains and function classes. Using a simplified two-layer attention-only model, the study evaluates two-stage gradient descent on the population cross-entropy loss, tracing the evolution of attention weights across training iterations. These theoretical derivations are paired with empirical simulations across varied causal tree graphs and multi-parent structures using standard neural network optimization.
The article establishes several key findings. First, gradient descent provably drives the first attention layer to converge directly to the adjacency matrix of the hidden causal graph, successfully attending to true parent tokens with probability approaching one. Second, the mechanism driving this alignment is informational: the mathematical gradient of the attention matrix naturally computes the mutual information between token pairs, which, due to information theory constraints, strictly peaks at direct causal parent connections. Third, the second attention layer learns to aggregate matching historical contexts, forming an induction circuit that achieves near-optimal statistical predictions with vanishing loss errors scaling inversely with effective sequence length. Fourth, the resulting model generalizes effectively to out-of-distribution transition distributions, achieving bounded estimation error with high probability. Finally, the authors show that in tasks where tokens depend on multiple causal parents, multi-head architectures distribute individual parent associations across distinct heads.
These findings provide rigorous theoretical justification for why attention mechanisms excel at structured sequence modeling. Rather than acting as black boxes, transformer layers optimized via gradient methods systematically mirror classical tree-recovery algorithms by extracting maximum mutual information. This mechanistic insight reduces architectural uncertainty and confirms that standard optimization reliably guides attention layers toward robust, interpretable causal representations without requiring explicit graph supervision.
Organizations developing sequence models should leverage multi-layer and multi-head attention structures tailored to the dependency complexity of their domain, ensuring context length and head capacity match the underlying data-generating graphs. For complex dependency modeling, practitioners should ensure adequate prompt length relative to vocabulary size to allow mixing and accurate in-context estimation. Future engineering and research efforts should explore the optimization dynamics of multi-head symmetry breaking under multiple parent dependencies and extend these dynamic guarantees to deeper architectures with non-linear feedforward layers.
The analysis relies on specific theoretical boundary conditions, including tree-structured latent dependencies, two-layer attention-focused architectures, and stage-wise gradient updates on well-behaved stationary measures. Confidence in the core finding—that gradient descent naturally aligns attention weights with maximum mutual information—is very high for the studied regimes, though readers should exercise appropriate caution when extrapolating specific convergence rates directly to massive, highly non-linear production models without further empirical validation.
- Paper: Attention Is All You Need, Ashish Vaswani et al. (2017). Introduces the foundational Transformer architecture and multi-head self-attention mechanism whose optimization dynamics and causal structure learning are theoretically analyzed in the source.
- Paper: Quantifying Attention Flow in Transformers, Samira Abnar et al. (2020). Establishes how information propagates across Transformer layers using directed acyclic graph representations of attention, providing key context for understanding layer-wise causal graph encoding.
- Paper: What Does BERT Look at? An Analysis of BERT’s Attention, Kevin Clark et al. (2019). Analyzes how individual attention heads specialize into functional sub-circuits, directly relating to the induction heads and attention structures analyzed in the source.
- Paper: One-Layer Transformer Provably Learns Multiclass One-Nearest Neighbor in Context, Skanda Athreya et al. (2026). Extends the theoretical study of gradient descent dynamics on shallow attention models from learning latent causal graphs to learning in-context nearest-neighbor classification rules.
- Paper: Next-Latent Prediction Transformers Learn Compact World Models, Jayden Teoh et al. (2025). Builds upon how Transformers learn internal causal graphs and representations by introducing auxiliary latent prediction objectives to reinforce compact world model representations.
- Paper: Learn from your own latents and not from tokens: A sample-complexity theory, Daniel J. Korchinski et al. (2026). Provides a complementary sample-complexity theory for how neural networks learn latent hierarchical structures during gradient-based training rather than simple token prediction.
- Paper: Test-time regression: a unifying framework for designing sequence models with associative memory, Ke Alexander Wang et al. (2025). Generalizes the mechanisms behind in-context associative recall and induction phenomena into a formal test-time regression framework across sequence model architectures.
