How Transformers Learn Causal Structure with Gradient Descent

Eshaan NichaniAlex DamianJason D. Lee

article2024ICML132 citations

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.

Listen

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.
Cover for How Transformers Learn Causal Structure with Gradient Descent

Abstract

The incredible success of transformers on sequence modeling tasks can be largely attributed to the self-attention mechanism, which allows information to be transferred between different parts of a sequence. Self-attention allows transformers to encode causal structure which makes them particularly suitable for sequence modeling. However, the process by which transformers learn such causal structure via gradient-based training algorithms remains poorly understood. To better understand this process, we introduce an in-context learning task that requires learning latent causal structure. We prove that gradient descent on a simplified two-layer transformer learns to solve this task by encoding the latent causal graph in the first attention layer. The key insight of our proof is that the gradient of the attention matrix encodes the mutual information between tokens. As a consequence of the data processing inequality, the largest entries of this gradient correspond to edges in the latent causal graph. As a special case, when the sequences are generated from in-context Markov chains, we prove that transformers learn an induction head (Olsson et al., 2022). We confirm our theoretical findings by showing that transformers trained on our in-context learning task are able to recover a wide variety of causal structures.

Table of Contents

  • 1. Introduction
  • 1.1. Our Contributions
  • 1.2. Related Work
  • 2. Setup
  • 2.1. Transformer Architecture
  • 2.2. Problem Setup: Random Sequences with Causal Structure
  • 2.3. Examples
  • 3. What does the Transformer Learn?
  • 3.1. Experiments
  • 3.2. Construction
  • 3.3. The Reduced Model
  • 4. Main Results
  • 4.1. Training Algorithm
  • 4.2. Main Theorem
  • 5. Proof Sketch
  • 5.1. Stage 1: Learning the Causal Graph
  • 5.1.1. THE ORACLE ALGORITHM
  • 5.1.2. THE GRADIENT DESCENT DYNAMICS
  • 5.2. Stage 2: Decreasing the Loss
  • 6. Causal Graphs with Multiple Parents
  • Impact Statement
  • Acknowledgements
  • References
  • A. Disentangled Transformer Equivalence
  • B. Multiple Parents Construction
  • C. Additional Experiments and Details
  • D. Analyzing the Dynamics
  • D.1. Proof of Lemma 3.2
  • D.2. Notation
  • D.3. Heuristic Derivation of Lemma 5.3
  • D.4. Gradient Computations
  • D.5. Gradient of A (1) (Stage 1)
  • D.6. Gradient of A (2) (Stage 2)
  • D.7. Proof of Theorem 4.4
  • E. Markov Chain Preliminaries
  • F. Concentration
  • G. Lemmas for Stage 1
  • G.1. Strong Data Processing Inequality
  • G.2. Auxiliary Dynamics Lemmas
  • G.3. Concentration
  • H. Lemmas for Stage 2
  • H.1. Idealized Gradient
  • H.2. Auxiliary Dynamics Lemmas
  • H.3. Concentration
  • H.4. Proof of Theorem 4.5
  • I. Finite Sample Analysis
  • I.1. Stage 1
  • I.2. Stage 2

Knowls

  1. Knowl 1 — Convergence and Causal Structure Recovery Guarantee for Two-Layer Transformers

    theoretical result

    Let G=([T],E)G = ([T], E) be a directed acyclic graph representing a latent causal structure where each node i∈[T]i \in [T] has at most one parent p(i)∈[i−1]p(i) \in [i-1], and R={i∈[T]:p(i)=∅}R = \{i \in [T] : p(i) = \emptyset\} denotes the set of root nodes. Let Teff:=T/max⁡iLiT_{\text{eff}} := T / \max_i L_i be the effective sequence length, where LiL_i is the number of leaves in tree TiT_i of the disjoint tree decomposition G=⋃i=1kTiG = \bigcup_{i=1}^k T_i.

    Consider the reduced two-layer autoregressive transformer fθ(s1:T)=X⊤S(S(A(1))XA(2)xT)∈RSf_\theta(s_{1:T}) = X^\top S(S(A^{(1)}) X A^{(2)} x_T) \in \mathbb{R}^S, where X=[es1,…,esT]⊤∈RT×SX = [e_{s_1}, \dots, e_{s_T}]^\top \in \mathbb{R}^{T \times S} is the sequence token embedding matrix, A(1)∈RT×TA^{(1)} \in \mathbb{R}^{T \times T} is a lower-triangular attention matrix, A(2)∈RS×SA^{(2)} \in \mathbb{R}^{S \times S} is a token-token attention matrix, and S(⋅)S(\cdot) denotes the row-wise softmax operator. The model is trained to minimize the perturbed population cross-entropy loss

    L(θ)=−Eπ,s1:T[∑s′∈[S]π(s′∣sT)log⁡(fθ(s1:T)s′+ϵ)]L(\theta) = -\mathbb{E}_{\pi, s_{1:T}} \left[ \sum_{s' \in [S]} \pi(s' \mid s_T) \log (f_\theta(s_{1:T})_{s'} + \epsilon) \right]

    with perturbation ϵ=Teff−1/2\epsilon = T_{\text{eff}}^{-1/2}, where the sequence s1:Ts_{1:T} is generated from a Markov transition matrix π\pi sampled from prior PπP_\pi.

    Assume the prior PπP_\pi satisfies regularity conditions with parameter γ>0\gamma > 0, the root fraction satisfies ∣R∣/T≤1−γ|R|/T \le 1 - \gamma, and Teff≥poly(γ−1,S)T_{\text{eff}} \ge \text{poly}(\gamma^{-1}, S). When training fθf_\theta via two-stage gradient descent initialized at A(1)(0)=0T×TA^{(1)}(0) = 0_{T \times T} and A(2)(0)=β0ISA^{(2)}(0) = \beta_0 I_S with β0≤cγ,STeff−3/2\beta_0 \le c_{\gamma, S} T_{\text{eff}}^{-3/2}, there exist step sizes η1,η2\eta_1, \eta_2 and step counts τ1,τ2\tau_1, \tau_2 such that the resulting parameter θ^=(A^(1),A^(2))\hat{\theta} = (\hat{A}^{(1)}, \hat{A}^{(2)}) satisfies:

    1. Graph Recovery in First Layer: For all non-root nodes i∉Ri \notin R, the first attention layer places almost all attention weight on the true parent p(i)p(i):
    S(A^(1))i,p(i)≥1−O(1T).S(\hat{A}^{(1)})_{i, p(i)} \ge 1 - O\left(\frac{1}{T}\right).
    1. Population Loss Convergence: The population loss reaches the minimal possible loss L∗:=−Eπ[1S∑s,s′π(s′∣s)log⁡π(s′∣s)]L^* := -\mathbb{E}_\pi [\frac{1}{S} \sum_{s, s'} \pi(s' \mid s) \log \pi(s' \mid s)] up to a polynomial error in effective context length:
    L(θ^)−L∗≤Cγ,S(log⁡2TTeff)cγ/24L(\hat{\theta}) - L^* \le C_{\gamma, S} \left( \frac{\log^2 T}{T_{\text{eff}}} \right)^{c\gamma / 24}

    for a universal constant c>0c > 0.

  2. Knowl 2 — Random Sequences with Latent Causal Structure

    definition

    Let [S]={1,…,S}[S] = \{1, \dots, S\} be a finite alphabet, and let G=([T],E)G = ([T], E) be a directed acyclic graph over positions [T]={1,…,T}[T] = \{1, \dots, T\} representing a latent causal structure, where directed edges (j→i)∈E(j \to i) \in E only exist if j<ij < i. For each position i∈[T]i \in [T], p(i):={j:(j→i)∈E}p(i) := \{j : (j \to i) \in E\} denotes the set of parent nodes, and R:={i∈[T]:p(i)=∅}R := \{i \in [T] : p(i) = \emptyset\} denotes the set of root nodes, with ∣p(i)∣≤1|p(i)| \le 1 for all i∈[T]i \in [T].

    Let PπP_\pi be a prior distribution over irreducible and aperiodic Markov transition matrices π\pi on [S][S], and let μπ\mu_\pi denote the unique stationary distribution of π\pi. A random sequence with causal structure GG and target token yy is generated according to the following procedure:

    1. Sample a transition matrix π∼Pπ\pi \sim P_\pi.
    2. For i=1,…,T−1i = 1, \dots, T - 1:
      • If p(i)=∅p(i) = \emptyset, sample si∼μπs_i \sim \mu_\pi.
      • If p(i)={j}p(i) = \{j\}, sample si∼π(⋅∣sj)s_i \sim \pi(\cdot \mid s_j).
    3. Draw the context query token sT∼Unif([S])s_T \sim \text{Unif}([S]) and the target token y=sT+1∼π(⋅∣sT)y = s_{T+1} \sim \pi(\cdot \mid s_T). (Hence T∈RT \in R is always a root node).
    4. Return prompt sequence x=s1:T=(s1,…,sT)x = s_{1:T} = (s_1, \dots, s_T) and target label y=sT+1y = s_{T+1}.

    Special cases include:

    • In-context Markov chain estimation (induction heads): p(i)=i−1p(i) = i - 1 for all i>1i > 1.
    • In-context function learning: p(2k−1)=∅p(2k - 1) = \emptyset and p(2k)=2k−1p(2k) = 2k - 1 for k≥1k \ge 1, where transitions encode input-output pairs (xk,f(xk))(x_k, f(x_k)).
  3. Knowl 3 — Attention Layer Gradients Compute Conditional Chi-Squared Mutual Information

    theoretical result

    For a two-layer transformer with first-layer lower-triangular attention matrix A(1)∈RT×TA^{(1)} \in \mathbb{R}^{T \times T} and initialization scale A(2)=β0ISA^{(2)} = \beta_0 I_S, the population gradient with respect to row ii of A(1)A^{(1)} is given by

    ∇Ai(1)L(θ)=−β0STJ(S(Ai(1)))(gi+O(Teff−1/2))\nabla_{A^{(1)}_i} L(\theta) = -\frac{\beta_0}{S T} J(S(A^{(1)}_i)) \left( g_i + O(T_{\text{eff}}^{-1/2}) \right)

    where J(v)=diag(v)−vv⊤J(v) = \text{diag}(v) - v v^\top is the Jacobian matrix of the softmax operator S(⋅)S(\cdot), and the entries of the idealized gradient vector gi∈Rig_i \in \mathbb{R}^i are

    gi,j:=Eπ[∑s,s′∈[S]π(s′∣s)μπ(s′)P[sj=s,si=s′]]−1.g_{i, j} := \mathbb{E}_\pi \left[ \sum_{s, s' \in [S]} \frac{\pi(s' \mid s)}{\mu_\pi(s')} \mathbb{P}[s_j = s, s_i = s'] \right] - 1.

    For any non-root position i∉Ri \notin R with true parent p(i)p(i):

    1. The gradient entry at the true parent equals the conditional χ2\chi^2-mutual information between sis_i and sp(i)s_{p(i)} given π\pi:
    gi,p(i)=Iχ2(si;sp(i)∣π):=Eπ[∑s,s′(P[si=s′,sp(i)=s∣π])2μπ(s′)μπ(s)−1].g_{i, p(i)} = I_{\chi^2}(s_i ; s_{p(i)} \mid \pi) := \mathbb{E}_\pi \left[ \sum_{s, s'} \frac{(\mathbb{P}[s_i = s', s_{p(i)} = s \mid \pi])^2}{\mu_\pi(s') \mu_\pi(s)} - 1 \right].
    1. For any non-parent position j≠p(i)j \ne p(i), the sequence conditioned on π\pi forms a Markov chain sj→sp(i)→sis_j \to s_{p(i)} \to s_i. By the Data Processing Inequality for χ2\chi^2-divergence, the mutual information is strictly strictly contracted:
    gi,p(i)−gi,j≥1−α(π)2∥Bπ∥F2≥γ32Sg_{i, p(i)} - g_{i, j} \ge \frac{1 - \alpha(\pi)}{2} \|B_\pi\|_F^2 \ge \frac{\gamma^3}{2 S}

    where α(π)≤1−γ\alpha(\pi) \le 1 - \gamma is the Markov contraction coefficient and ∥Bπ∥F2≥γ2/S\|B_\pi\|_F^2 \ge \gamma^2/S is the non-degeneracy of the chain.

    1. For root positions i∈Ri \in R, sis_i is independent of all preceding tokens sjs_j given π\pi, leading to gi,j=0g_{i, j} = 0, which preserves an approximately uniform attention distribution S(Ai(1))j≈1/iS(A^{(1)}_i)_j \approx 1/i for all j≤ij \le i.
  4. Knowl 4 — Out-of-Distribution Generalization of Learned Causal Transformers

    theoretical result

    Let θ^=(A^(1),A^(2))\hat{\theta} = (\hat{A}^{(1)}, \hat{A}^{(2)}) be the two-layer transformer parameters obtained from running two-stage gradient descent (Algorithm 1) on sequences generated with prior distribution PπP_\pi.

    Let π~∈RS×S\tilde{\pi} \in \mathbb{R}^{S \times S} be an arbitrary target Markov transition matrix that may lie completely outside the support of the training prior PπP_\pi, satisfying only the entrywise lower bound

    min⁡s,s′∈[S]π~(s′∣s)≥γS.\min_{s, s' \in [S]} \tilde{\pi}(s' \mid s) \ge \frac{\gamma}{S}.

    If a sequence s1:Ts_{1:T} is generated according to the latent causal DAG GG using the out-of-distribution transition matrix π~\tilde{\pi}, then with probability at least 0.990.99 over the draw of s1:Ts_{1:T}, the prediction fθ^(s1:T)∈RSf_{\hat{\theta}}(s_{1:T}) \in \mathbb{R}^S uniformly approximates the true transition distribution π~(⋅∣sT)\tilde{\pi}(\cdot \mid s_T):

    sup⁡s′∈[S]∣fθ^(s1:T)s′−π~(s′∣sT)∣≤Cγ,Slog⁡TTeffcγ\sup_{s' \in [S]} \left| f_{\hat{\theta}}(s_{1:T})_{s'} - \tilde{\pi}(s' \mid s_T) \right| \le C_{\gamma, S} \frac{\log T}{T_{\text{eff}}^{c\gamma}}

    where TeffT_{\text{eff}} is the effective sequence length of GG, and c>0,Cγ,S>0c > 0, C_{\gamma, S} > 0 are constants independent of TT.

  5. Knowl 5 — Two-Stage Gradient Descent for Causal Transformer Training

    algorithm

    The training procedure uses stage-wise gradient descent on the population cross-entropy loss L(θ)L(\theta) using a reduced two-layer transformer model fθf_\theta, parameterized by lower-triangular position attention matrix A(1)∈RT×TA^{(1)} \in \mathbb{R}^{T \times T} and token attention matrix A(2)∈RS×SA^{(2)} \in \mathbb{R}^{S \times S}.

    In Stage 1, A(1)A^{(1)} is trained with learning rate η1\eta_1 for τ1\tau_1 steps while A(2)A^{(2)} is frozen at small initialization β0IS\beta_0 I_S, driving A(1)A^{(1)} to encode the causal adjacency matrix. In Stage 2, A(2)A^{(2)} is trained with learning rate η2\eta_2 for τ2\tau_2 steps while A(1)A^{(1)} is held fixed at A(1)(τ1)A^{(1)}(\tau_1), growing the magnitude of A(2)A^{(2)} along IS−1S1S1S⊤I_S - \frac{1}{S} \mathbf{1}_S \mathbf{1}_S^\top to minimize predictive cross-entropy.

    Input: initialization scale β0\beta_0, learning rates η1,η2\eta_1, \eta_2, step counts τ1,τ2\tau_1, \tau_2
    Initialize A(1)(0)=0T×TA^{(1)}(0) = 0_{T \times T}, A(2)(0)=β0⋅IS×SA^{(2)}(0) = \beta_0 \cdot I_{S \times S}
    for t=1,…,τ1t = 1, \dots, \tau_1 do
        A(1)(t)←A(1)(t−1)−η1∇A(1)L(A(1)(t−1),A(2)(0))A^{(1)}(t) \leftarrow A^{(1)}(t-1) - \eta_1 \nabla_{A^{(1)}} L(A^{(1)}(t-1), A^{(2)}(0))
    end for
    for t=τ1,…,τ1+τ2−1t = \tau_1, \dots, \tau_1 + \tau_2 - 1 do
        A(2)(t+1)←A(2)(t)−η2∇A(2)L(A(1)(τ1),A(2)(t))A^{(2)}(t+1) \leftarrow A^{(2)}(t) - \eta_2 \nabla_{A^{(2)}} L(A^{(1)}(\tau_1), A^{(2)}(t))
    end for
    θ^←(A(1)(τ1),A(2)(τ1+τ2))\hat{\theta} \leftarrow (A^{(1)}(\tau_1), A^{(2)}(\tau_1 + \tau_2))
    Output: θ^\hat{\theta}
  6. Knowl 6 — Reduced Two-Layer Disentangled Transformer Model

    model/method

    Let [S][S] be a finite alphabet of size SS, and let s1:T=(s1,…,sT)∈[S]Ts_{1:T} = (s_1, \dots, s_T) \in [S]^T be an input sequence. Let X=[es1,…,esT]⊤∈RT×SX = [e_{s_1}, \dots, e_{s_T}]^\top \in \mathbb{R}^{T \times S} denote the one-hot token embedding matrix of the sequence.

    The reduced two-layer transformer model fθ:[S]T→RSf_\theta : [S]^T \to \mathbb{R}^S parameterized by θ=(A(1),A(2))\theta = (A^{(1)}, A^{(2)}) with lower-triangular matrix A(1)∈RT×TA^{(1)} \in \mathbb{R}^{T \times T} and symmetric matrix A(2)∈RS×SA^{(2)} \in \mathbb{R}^{S \times S} is defined as:

    fθ(s1:T)=X⊤S(S(MASK(A(1)))XA(2)⊤esT)f_\theta(s_{1:T}) = X^\top S\left( S(\text{MASK}(A^{(1)})) X A^{(2)\top} e_{s_T} \right)

    where MASK(M)i,j=Mi,j\text{MASK}(M)_{i,j} = M_{i,j} for i≥ji \ge j and −∞-\infty for i<ji < j, and S(⋅)S(\cdot) is the softmax function applied row-wise.

    To prevent infinite loss when tokens do not appear in context, predictions are perturbed by ϵ=Teff−1/2\epsilon = T_{\text{eff}}^{-1/2}. The perturbed population cross-entropy loss under prior PπP_\pi is

    L(θ)=−Eπ,s1:T[∑s′∈[S]π(s′∣sT)log⁡(fθ(s1:T)s′+ϵ)].L(\theta) = -\mathbb{E}_{\pi, s_{1:T}} \left[ \sum_{s' \in [S]} \pi(s' \mid s_T) \log (f_\theta(s_{1:T})_{s'} + \epsilon) \right].
  7. Knowl 7 — Equivalence of Decoder-Based and Disentangled Transformers

    theoretical result

    A standard decoder-based attention-only transformer TFθ\text{TF}_\theta of depth LL with embedding dimension dd, token embeddings E∈Rd×SE \in \mathbb{R}^{d \times S}, positional embeddings P∈Rd×TP \in \mathbb{R}^{d \times T}, head parameters (Qi(ℓ),Ki(ℓ),Vi(ℓ))∈(Rd×d)3(Q_i^{(\ell)}, K_i^{(\ell)}, V_i^{(\ell)}) \in (\mathbb{R}^{d \times d})^3 for i∈[mℓ]i \in [m_\ell], and output projection WO∈Rdout×dW_O \in \mathbb{R}^{d_{\text{out}} \times d} computes residual updates h(ℓ)=h(ℓ−1)+∑i=1mℓattn(h(ℓ−1);Qi(ℓ)Ki(ℓ)⊤)Vi(ℓ)⊤h^{(\ell)} = h^{(\ell-1)} + \sum_{i=1}^{m_\ell} \text{attn}(h^{(\ell-1)}; Q_i^{(\ell)} K_i^{(\ell)\top}) V_i^{(\ell)\top}.

    A disentangled transformer TF~θ~\widetilde{\text{TF}}_{\tilde{\theta}} initializes h(0)=X~=[est,et]t=1T∈RT×d0h^{(0)} = \tilde{X} = [e_{s_t}, e_t]_{t=1}^T \in \mathbb{R}^{T \times d_0} with d0=S+Td_0 = S + T, and at each layer ℓ∈[L]\ell \in [L] concatenates the attention outputs along the feature dimension to produce h(ℓ)=[h(ℓ−1),attn(h(ℓ−1);A~1(ℓ)),…,attn(h(ℓ−1);A~mℓ(ℓ))]∈RT×dℓh^{(\ell)} = [h^{(\ell-1)}, \text{attn}(h^{(\ell-1)}; \tilde{A}_1^{(\ell)}), \dots, \text{attn}(h^{(\ell-1)}; \tilde{A}_{m_\ell}^{(\ell)})] \in \mathbb{R}^{T \times d_\ell} with dℓ=(1+mℓ)dℓ−1d_\ell = (1 + m_\ell)d_{\ell-1}, parameterized by single attention matrices A~i(ℓ)∈Rdℓ−1×dℓ−1\tilde{A}_i^{(\ell)} \in \mathbb{R}^{d_{\ell-1} \times d_{\ell-1}} and output projection W~O∈Rdout×dL\tilde{W}_O \in \mathbb{R}^{d_{\text{out}} \times d_L}.

    For any standard transformer TFθ\text{TF}_\theta, there exists a disentangled transformer TF~θ~\widetilde{\text{TF}}_{\tilde{\theta}} of the same depth and number of heads such that TFθ(s1:T)=TF~θ~(s1:T)\text{TF}_\theta(s_{1:T}) = \widetilde{\text{TF}}_{\tilde{\theta}}(s_{1:T}) for all s1:T∈[S]Ts_{1:T} \in [S]^T. Conversely, any disentangled transformer TF~θ~\widetilde{\text{TF}}_{\tilde{\theta}} can be exactly represented by a standard transformer TFθ\text{TF}_\theta with hidden dimension d=dLd = d_L.

  8. Knowl 8 — Exact Construction for In-Context Causal Transition Estimation

    model/method

    For a latent causal DAG G=([T],E)G = ([T], E) with parent function p(i)p(i) for i∉Ri \notin R and root nodes R={i:p(i)=∅}R = \{i : p(i) = \emptyset\}, there exists a two-layer disentangled transformer whose output approximates the empirical transition estimator

    π^s1:T(s′∣sT):=∣{(j→i)∈E:(sj,si)=(sT,s′)}∣∣{(j→i)∈E:sj=sT}∣.\hat{\pi}_{s_{1:T}}(s' \mid s_T) := \frac{|\{(j \to i) \in E : (s_j, s_i) = (s_T, s')\}|}{|\{(j \to i) \in E : s_j = s_T\}|}.

    The construction sets attention matrices A~(1)\tilde{A}^{(1)} and A~(2)\tilde{A}^{(2)} as:

    A~(1)=[0S×S0S×T0T×SA(1)],A~(2)=[0S×S0S×TA(2)0S×T0T×S0T×T0T×S0T×T0S×S0S×T0S×S0S×T0T×S0T×T0T×S0T×T]\tilde{A}^{(1)} = \begin{bmatrix} 0_{S \times S} & 0_{S \times T} \\ 0_{T \times S} & A^{(1)} \end{bmatrix}, \quad \tilde{A}^{(2)} = \begin{bmatrix} 0_{S \times S} & 0_{S \times T} & A^{(2)} & 0_{S \times T} \\ 0_{T \times S} & 0_{T \times T} & 0_{T \times S} & 0_{T \times T} \\ 0_{S \times S} & 0_{S \times T} & 0_{S \times S} & 0_{S \times T} \\ 0_{T \times S} & 0_{T \times T} & 0_{T \times S} & 0_{T \times T} \end{bmatrix}

    where Ai,j(1)=β11(j=p(i))A^{(1)}_{i, j} = \beta_1 \mathbf{1}(j = p(i)) for i∉Ri \notin R (and 00 for i∈Ri \in R), and A(2)=β2ISA^{(2)} = \beta_2 I_S, with β1,β2→∞\beta_1, \beta_2 \to \infty.

    In the forward pass:

    1. The first attention layer copies the parent token embedding x~p(i)\tilde{x}_{p(i)} into position ii for non-root positions i∉Ri \notin R, and averages all prior tokens 1i∑j≤ix~j\frac{1}{i}\sum_{j \le i} \tilde{x}_j for root positions i∈Ri \in R.
    2. The second attention layer compares the final token sTs_T against the copied parent tokens in the residual stream, attending equally to all tokens ii where sp(i)=sTs_{p(i)} = s_T.
    3. The output layer W~O=[0S×d,0S×d,IS,0S×T,0S×d]\tilde{W}_O = [0_{S \times d}, 0_{S \times d}, I_S, 0_{S \times T}, 0_{S \times d}] extracts the averaged token vector, yielding fθ(s1:T)s′=π^s1:T(s′∣sT)f_\theta(s_{1:T})_{s'} = \hat{\pi}_{s_{1:T}}(s' \mid s_T).
  9. Knowl 9 — In-Context Estimation for Multi-Parent Causal Graphs and n-Grams

    model/method

    Let G=([T+1],E)G = ([T+1], E) be a directed acyclic graph where each non-root node i∈[T+1]∖Ri \in [T+1] \setminus R has exactly kk parent nodes p(i)={p(i)1,…,p(i)k}⊂[i−1]p(i) = \{p(i)_1, \dots, p(i)_k\} \subset [i-1] with p(i)1<⋯<p(i)kp(i)_1 < \dots < p(i)_k. Transitions are governed by kk-parent transition tensors π(⋅∣a1,…,ak)\pi(\cdot \mid a_1, \dots, a_k) sampled from prior PπkP_\pi^k. For nn-gram models, k=n−1k = n - 1 and p(i)={i−n+1,…,i−1}p(i) = \{i - n + 1, \dots, i - 1\}.

    A two-layer transformer with kk heads in the first layer computes the empirical multi-parent transition estimate

    π^s1:T(s′∣sp(T+1)1,…,sp(T+1)k)=∣{i<T:si=s′,sp(i)1=sp(T+1)1,…,sp(i)k=sp(T+1)k}∣∣{j<T:sp(j)1=sp(T+1)1,…,sp(j)k=sp(T+1)k}∣.\hat{\pi}_{s_{1:T}}(s' \mid s_{p(T+1)_1}, \dots, s_{p(T+1)_k}) = \frac{|\{i < T : s_i = s', s_{p(i)_1} = s_{p(T+1)_1}, \dots, s_{p(i)_k} = s_{p(T+1)_k}\}|}{|\{j < T : s_{p(j)_1} = s_{p(T+1)_1}, \dots, s_{p(j)_k} = s_{p(T+1)_k}\}|}.

    In this construction:

    • The ℓ\ell-th attention head in layer 1 uses scaled position-position weights (Aℓ(1))i,j=β11(j=p(i)ℓ)(A_\ell^{(1)})_{i, j} = \beta_1 \mathbf{1}(j = p(i)_\ell) for i<Ti < T and (Aℓ(1))T,j=β11(j=p(T+1)ℓ)(A_\ell^{(1)})_{T, j} = \beta_1 \mathbf{1}(j = p(T+1)_\ell), copying the ℓ\ell-th parent of ii into the residual stream of ii, and the ℓ\ell-th parent of T+1T+1 into the residual stream of TT.
    • The second attention layer uses block diagonal weights A(2)=diag(0d×d,β2IS,…,β2IS)A^{(2)} = \text{diag}(0_{d \times d}, \beta_2 I_S, \dots, \beta_2 I_S) to compute inner products β2∑ℓ=1k1(sp(i)ℓ=sp(T+1)ℓ)\beta_2 \sum_{\ell=1}^k \mathbf{1}(s_{p(i)_\ell} = s_{p(T+1)_\ell}), selecting tokens whose parent tuples match (sp(T+1)1,…,sp(T+1)k)(s_{p(T+1)_1}, \dots, s_{p(T+1)_k}).
  10. Knowl 10 — Assumptions on Transition Prior and Graph Mixing

    assumption

    The theoretical convergence guarantees for causal structure recovery and in-context learning require the following structural conditions on the prior PπP_\pi over Markov transition matrices π\pi on [S][S] and on the causal graph G=([T],E)G = ([T], E):

    1. Transition Lower Bound: Almost surely over π∼Pπ\pi \sim P_\pi, min⁡s,s′∈[S]π(s′∣s)>γ/S\min_{s, s' \in [S]} \pi(s' \mid s) > \gamma / S for a constant γ>0\gamma > 0, ensuring stationary measures satisfy min⁡sμπ(s)≥γ/S\min_s \mu_\pi(s) \ge \gamma / S.
    2. Non-Degeneracy: The Markov chain does not mix in a single step, satisfying ∑s∥π(⋅∣s)−μπ(⋅)∥22≥γ2/S\sum_s \|\pi(\cdot \mid s) - \mu_\pi(\cdot)\|_2^2 \ge \gamma^2 / S.
    3. Permutation Symmetry and Uniform Mean: For any permutation σ\sigma on [S][S], σ−1πσ=dπ\sigma^{-1} \pi \sigma \stackrel{d}{=} \pi, and Eπ[π]=1S1S1S⊤\mathbb{E}_\pi[\pi] = \frac{1}{S} \mathbf{1}_S \mathbf{1}_S^\top.
    4. Non-Vanishing Non-Root Fraction: The fraction of root nodes r:=∣R∣/Tr := |R|/T satisfies r≤1−γr \le 1 - \gamma.
    5. Effective Sequence Length: For graph decomposition G=⋃i=1kTiG = \bigcup_{i=1}^k T_i into disjoint trees where tree TiT_i has LiL_i leaves, the effective sequence length Texteff:=T/max⁡iLiT_{ ext{eff}} := T / \max_i L_i satisfies Texteff≥poly(γ−1,S)T_{ ext{eff}} \ge \text{poly}(\gamma^{-1}, S).

Coverage note — Omitted detailed finite-sample concentration proofs (Section I) and specific empirical heatmap figures, as their primary takeaways and sample complexity scaling ($N \gtrsim T_{\text{eff}} \log T$) are directly reflected in the theoretical and architectural knowls.

References

  1. 1.Ahn, K., Cheng, X., Daneshmand, H., and Sra, S. Transformers learn to implement preconditioned gradient descent for in-context learning. arXiv preprint arXiv:2306.00297, 2023.
  2. 2.Aky"urek, E., Schuurmans, D., Andreas, J., Ma, T., and Zhou, D. What learning algorithm is in-context learning? investigations with linear models. In The Eleventh International Conference on Learning Representations, 2023. URL https://openreview.net/forum?id=0g0X4H8yN4I.
  3. 3.Aky"urek, E., Wang, B., Kim, Y., and Andreas, J. In-context language learning: Arhitectures and algorithms. arXiv preprint arXiv:2401.12973, 2024.
  4. 4.Bai, Y., Chen, F., Wang, H., Xiong, C., and Mei, S. Transformers as statisticians: Provable in-context learning with in-context algorithm selection. arXiv preprint arXiv:2306.04637, 2023.
  5. 5.Bietti, A., Cabannes, V., Bouchacourt, D., Jegou, H., and Bottou, L. Birth of a transformer: A memory viewpoint. arXiv preprint arXiv:2306.00802, 2023.
  6. 6.Boix-Adsera, E., Littwin, E., Abbe, E., Bengio, S., and Susskind, J. Transformers learn through gradual rank increase. arXiv preprint arXiv:2306.07042, 2023.
  7. 7.Bradbury, J., Frostig, R., Hawkins, P., Johnson, M. J., Leary, C., Maclaurin, D., Necula, G., Paszke, A., VanderPlas, J., Wanderman-Milne, S., and Zhang, Q. JAX: composable transformations of Python+NumPy programs, 2018. URL http://github.com/google/jax.
  8. 8.Brown, T., Mann, B., Ryder, N., Subbiah, M., Kaplan, J. D., Dhariwal, P., Neelakantan, A., Shyam, P., Sastry, G., Askell, A., et al. Language models are few-shot learners. Advances in neural information processing systems, 33: 1877–1901, 2020.
  9. 9.Chen, L., Lu, K., Rajeswaran, A., Lee, K., Grover, A., Laskin, M., Abbeel, P., Srinivas, A., and Mordatch, I. Decision transformer: Reinforcement learning via sequence modeling. Advances in neural information processing systems, 34:15084–15097, 2021.
  10. 10.Chow, C. and Liu, C. Approximating discrete probability distributions with dependence trees. IEEE transactions on Information Theory, 14(3):462–467, 1968.
  11. 11.Cohen, J. E., Iwasa, Y., Rautu, G., Beth Ruskai, M., Seneta, E., and Zbaganu, G. Relative entropy under mappings by stochastic matrices. Linear Algebra and its Applications, 179:211–235, 1993. ISSN 0024-3795. doi: https://doi.org/10.1016/0024-3795(93)90331-H. URL https://www.sciencedirect.com/science/article/pii/002437959390331H.
  12. 12.Dosovitskiy, A., Beyer, L., Kolesnikov, A., Weissenborn, D., Zhai, X., Unterthiner, T., Dehghani, M., Minderer, M., Heigold, G., Gelly, S., et al. An image is worth 16x16 words: Transformers for image recognition at scale. arXiv preprint arXiv:2010.11929, 2020.
  13. 13.Edelman, B. L., Edelman, E., Goel, S., Malach, E., and Tsilivis, N. The evolution of statistical induction heads: In-context learning markov chains. arXiv preprint arXiv:2402.11004, 2024.
  14. 14.Elhage, N., Nanda, N., Olsson, C., Henighan, T., Joseph, N., Mann, B., Askell, A., Bai, Y., Chen, A., Conerly, T., DasSarma, N., Drain, D., Ganguli, D., Hatfield-Dodds, Z., Hernandez, D., Jones, A., Kernion, J., Lovitt, L., Ndousse, K., Amodei, D., Brown, T., Clark, J., Kaplan, J., McCandlish, S., and Olah, C. A mathematical framework for transformer circuits. Transformer Circuits Thread, 2021. https://transformer-circuits.pub/2021/framework/index.html.
  15. 15.Friedman, D., Wettig, A., and Chen, D. Learning transformer programs. arXiv preprint arXiv:2306.01128, 2023.
  16. 16.Fu, D., Chen, T.-Q., Jia, R., and Sharan, V. Transformers learn higher-order optimization methods for in-context learning: A study with linear models. arXiv preprint arXiv:2310.17086, 2023.
  17. 17.Garg, S., Tsipras, D., Liang, P. S., and Valiant, G. What can transformers learn in-context? a case study of simple function classes. Advances in Neural Information Processing Systems, 35:30583–30598, 2022.
  18. 18.Giannou, A., Rajput, S., Sohn, J.-y., Lee, K., Lee, J. D., and Papailiopoulos, D. Looped transformers as programmable computers. arXiv preprint arXiv:2301.13196, 2023.
  19. 19.Huang, Y., Cheng, Y., and Liang, Y. In-context convergence of transformers. arXiv preprint arXiv:2310.05249, 2023.
  20. 20.Jelassi, S., Sander, M., and Li, Y. Vision transformers provably learn spatial structure. Advances in Neural Information Processing Systems, 35:37822–37836, 2022.
  21. 21.Jumper, J., Evans, R., Pritzel, A., Green, T., Figurnov, M., Ronneberger, O., Tunyasuvunakool, K., Bates, R., Žídek, A., Potapenko, A., et al. Highly accurate protein structure prediction with alphafold. Nature, 596(7873):583–589, 2021.
  22. 22.Kingma, D. P. and Ba, J. Adam: A method for stochastic optimization, 2017.
  23. 23.Li, Y., Li, Y., and Risteski, A. How do transformers learn topic structure: Towards a mechanistic understanding. In International Conference on Machine Learning, pp. 19689–19729. PMLR, 2023.
  24. 24.Lu, H., Mao, Y., and Nayak, A. On the dynamics of training attention models. In International Conference on Learning Representations, 2021. URL https://openreview.net/forum?id=1OCTOShAmqB.
  25. 25.Mahankali, A., Hashimoto, T. B., and Ma, T. One step of gradient descent is provably the optimal in-context learner with one layer of linear self-attention. arXiv preprint arXiv:2307.03576, 2023.
  26. 26.Nguyen, T. and Grover, A. Transformer neural processes: Uncertainty-aware meta learning via sequence modeling. In International Conference on Machine Learning, pp. 16569–16594. PMLR, 2022.
  27. 27.Olsson, C., Elhage, N., Nanda, N., Joseph, N., DasSarma, N., Henighan, T., Mann, B., Askell, A., Bai, Y., Chen, A., Conerly, T., Drain, D., Ganguli, D., Hatfield-Dodds, Z., Hernandez, D., Johnston, S., Jones, A., Kernion, J., Lovitt, L., Ndousse, K., Amodei, D., Brown, T., Clark, J., Kaplan, J., McCandlish, S., and Olah, C. In-context learning and induction heads. Transformer Circuits Thread, 2022. https://transformer-circuits.pub/2022/in-context-learning-and-induction-heads/index.html.
  28. 28.Reddy, G. The mechanistic basis of data dependence and abrupt learning in an in-context classification task. arXiv preprint arXiv:2312.03002, 2023.
  29. 29.Snell, C., Zhong, R., Klein, D., and Steinhardt, J. Approximating how single head attention learns. arXiv preprint arXiv:2103.07601, 2021.
  30. 30.Tarzanagh, D. A., Li, Y., Thrampoulidis, C., and Oymak, S. Transformers as support vector machines. arXiv preprint arXiv:2308.16898, 2023a.
  31. 31.Tarzanagh, D. A., Li, Y., Zhang, X., and Oymak, S. Max-margin token selection in attention mechanism. In Thirty-seventh Conference on Neural Information Processing Systems, 2023b.
  32. 32.Tian, Y., Wang, Y., Chen, B., and Du, S. Scan and snap: Understanding training dynamics and token composition in 1-layer transformer. arXiv preprint arXiv:2305.16380, 2023a.
  33. 33.Tian, Y., Wang, Y., Zhang, Z., Chen, B., and Du, S. Joma: Demystifying multilayer transformers via joint dynamics of mlp and attention. arXiv preprint arXiv:2310.00535, 2023b.
  34. 34.Vaswani, A., Shazeer, N., Parmar, N., Uszkoreit, J., Jones, L., Gomez, A. N., Kaiser, Ł., and Polosukhin, I. Attention is all you need. Advances in neural information processing systems, 30, 2017.
  35. 35.Von Oswald, J., Niklasson, E., Randazzo, E., Sacramento, J., Mordvintsev, A., Zhmoginov, A., and Vladymyrov, M. Transformers learn in-context by gradient descent. In International Conference on Machine Learning, pp. 35151–35174. PMLR, 2023.
  36. 36.Xie, S. M., Raghunathan, A., Liang, P., and Ma, T. An explanation of in-context learning as implicit bayesian inference. arXiv preprint arXiv:2111.02080, 2021.
  37. 37.Zhang, R., Frei, S., and Bartlett, P. L. Trained transformers learn linear models in-context. arXiv preprint arXiv:2306.09927, 2023.

Citation

MLA
Nichani, E., et al. “How Transformers Learn Causal Structure with Gradient Descent”. arXiv, 2024, https://doi.org/10.48550/arxiv.2402.14735.
APA
Nichani, E., Damian, A., & Lee, J. D. (2024). How Transformers Learn Causal Structure with Gradient Descent. arXiv. https://doi.org/10.48550/arxiv.2402.14735
Chicago
Nichani, E., A. Damian, and J. D. Lee. 2024. “How Transformers Learn Causal Structure with Gradient Descent”. Preprint, ArXiv. https://doi.org/10.48550/arxiv.2402.14735.
Harvard
Nichani, E., Damian, A. and Lee, J.D. (2024) “How Transformers Learn Causal Structure with Gradient Descent”. arXiv. Available at: https://doi.org/10.48550/arxiv.2402.14735.
Vancouver
1. Nichani E, Damian A, Lee JD (2024) How Transformers Learn Causal Structure with Gradient Descent. https://doi.org/10.48550/arxiv.2402.14735

BibTeX

@misc{https://doi.org/10.48550/arxiv.2402.14735,
  doi = {10.48550/ARXIV.2402.14735},
  url = {https://arxiv.org/abs/2402.14735},
  author = {Nichani, Eshaan and Damian, Alex and Lee, Jason D.},
  keywords = {Machine Learning (cs.LG), Information Theory (cs.IT), Machine Learning (stat.ML), FOS: Computer and information sciences, FOS: Computer and information sciences},
  title = {How Transformers Learn Causal Structure with Gradient Descent},
  publisher = {arXiv},
  year = {2024},
  copyright = {arXiv.org perpetual, non-exclusive license}
}
Metadata:DOI registry

Source Code

This paper has an official code repository available. Click below to access the source code.

View Repository

Access the Paper

This paper is available from its original source. Click below to access the PDF.

Open PDF
License: https://creativecommons.org/licenses/by/4.0/