论文概要: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 点)
- 理论构造: 证明两层 Transformer + 相对位置编码 (RPE) 可实现贝叶斯模型平均 (BMA)——因果结构推断的最优统计算法。
- 实验验证: 训练后的 Transformer 在参数级别逼近 BMA,注意力权重直接编码因果后验概率。
- 信息论保证: 利用 χ²-互信息与数据处理不等式 (DPI),建立因果结构可辨识性条件;梯度初始化即可恢复因果结构。
- 连续系统扩展: 分析线性动力系统 (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−ρ²) η_h | MSE |
模型架构
- 第 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 |