论文概要:How Transformers Learn Causal Structures In-Context

作者: Jianzhe Wei, Siyu Chen, Jianliang He, Zhuoran Yang (Georgia Tech & Yale) 会议: ICLR 2026 Oral 论文标签: in-context-learning causal-structure transformer-mechanism bma iclr2026


背景问题

传统假设现实挑战
固定依赖结构(如 bigram、i.i.d. tokens)序列间因果结构动态变化
Nichani et al. (2024): 训练中编码固定父-子依赖同一模型需在上下文中推断不同因果图
理论分析仅适用于 n-gram 等刚性模式语言中句法结构随文档变化,股票市场资产关系随时间漂移

核心问题 (⋆): Transformer 能否在上下文中推断并适应可变因果结构


核心贡献(4 点)

  1. 理论构造: 证明两层 Transformer + 相对位置编码 (RPE) 可实现贝叶斯模型平均 (BMA)——因果结构推断的最优统计算法。
  2. 实验验证: 训练后的 Transformer 在参数级别逼近 BMA,注意力权重直接编码因果后验概率。
  3. 信息论保证: 利用 χ²-互信息与数据处理不等式 (DPI),建立因果结构可辨识性条件;梯度初始化即可恢复因果结构。
  4. 连续系统扩展: 分析线性动力系统 (DS) 下离散 vs 连续因果推断的表征能力差异,证明 Transformer 在 DS 下无法精确实现 BMA。

方法

任务设置

  • 数据: 长度 H 的 Markov 链,每个 token x_h 依赖唯一前驱 pa(h)~Unif(1,…,h−1)
  • 上下文: L 个同因果图示例 + 1 个待预测序列
  • 目标: 下一个 token 预测(需先推断 pa(h))

两种数据形式

类型生成方式损失函数
离散 Markov 链 (MC)x_h ~ π(·x_{pa(h)}), 有限词表 V,
连续动力系统 (DS)x_h = ρAᵀx_{pa(h)} + √(1−ρ²) η_hMSE

模型架构

  • 第 1 层: K 头 RPE 注意力 → 各头拷贝对应历史 token x_h
  • 残差拼接: 将 K 头输出拼接至 token 特征
  • 第 2 层: 单头注意力,WKQ 实现 BMA 分数计算,WOV 逼近 π(MC)或映射(DS)

关键结果

1. Transformer 实现了 BMA(Theorem 1)

在假设下(式 7),当 β→∞, L→∞ 时:

  • 第二层注意力权重 A^(2)(h→h’) = σ(Σ_l log π(x_h^l | x{h’}^l)) = BMA 后验
  • 预测收敛到真实条件分布 π(·|x_{pa(h)})

2. 参数级验证(Fig. 6)

指标数值
行 softmax 误差 ||σ_row(Wtf) − π||₁~0.35
列 softmax 误差 ||σ_col(Wtf) − σ_col(log π)||₁< 0.05

训练后的 Wtf 满足 Wtf = log π + 1aᵀ(列偏移不变性,Proposition 1),等价于 BMA。

3. 父节点选择损失对比(Fig. 5)

  • L=10 训练,测试 L’ ∈ [1,20] 时 Transformer 的父选择损失逼近 BMA
  • 小 L 训练的模型泛化更好;L’ 增大时损失趋近于 0

4. 信息论保证(Theorem 2)

  • DPI + 转移核下界条件 → 严格 DPI(Lemma 3)
  • 期望对数似然: E[log π(x_h|x_{pa(h)})] > E[log π(x_h|x_{h’})] (h’ ≠ pa(h))
  • L→∞ 时注意力权重收敛到 one-hot 真实父节点

5. 训练动力学(Theorem 3,informal)

梯度初始化时:∂ℓ(θ₀)/∂p̂_{pa(h)} ≥ ∂ℓ(θ₀)/∂p̂_{h’},即梯度自然偏向真实父节点,该关系由 χ²-互信息驱动。

6. 连续 DS 的表征局限(Proposition 2)

BMA 在 DS 下包含 x 的二范数项 d·Σ_l ||x_{h’}||²,而 Transformer 的 bilinear 形式 Wtf 无法产生此项,因此 Transformer 无法精确实现 BMA(与 MC 情形形成对比)。


实验设置

  • 训练配置: 学习率 0.0001,Adam 优化器,批量大小 64,训练步数 1024~2048
  • 参数规模: d ∈ {10, 20, 30, 50}, H ∈ {10, 15, 50}, L ∈ {1,…,20}
  • 消融实验: 标准 disentangled Transformer + FFN 也收敛到相同模式(Appendix G-I)
  • 可视化: 注意力权重 A² 直接显示因果选择(Fig. 2);参数 WKQ 呈块对角,w^H 的 0 位置主导(Fig. 3)

与 Prior Work 对比

工作因果结构方法局限
Nichani et al. (2024)固定结构(bigram)训练编码无法适应上下文变化的因果图
D’Angelo et al. (2025)可变结构选择性归纳头任务不同,本文框架可覆盖
本文可变结构BMA 实现 + 信息论保证DS 下无法精确 BMA