方法详解:How Transformers Learn Causal Structures In-Context
论文: How Transformers Learn Causal Structures In-Context: Explainable Mechanism Meets Theoretical Guarantee
会议: ICLR 2026 (Oral)
作者: Jianzhe Wei, Siyu Chen, Jianliang He, Zhuoran Yang (Georgia Tech & Yale)
1. 任务设定
1.1 马尔可夫链 × 随机因果依赖
考虑一个具有潜在因果结构的序列生成过程。每条序列 中,第 个 token 仅依赖于它的一个前驱 token,称为父节点 。依赖关系可以表示为一棵有向树 ,其中:
即每个位置的父索引从所有前序位置中均匀随机采样。数据的生成过程为:
其中 是马尔可夫转移核(随机矩阵), 是词表大小。tokens 为 one-hot 向量 。
1.2 In-Context 学习任务
给定 条来自相同因果图 的序列 :
- 前 条是 in-context 示例(demonstrations),模型从中推断隐式的因果结构 ;
- 最后一条是预测目标(target),需要对 逐个进行 next-token prediction,条件为所有上下文 和已观测的 。
核心挑战: 在不同的上下文中是不同的(跨序列变化),模型必须在上下文中自适应地推断出父节点是谁。
1.3 连续动力系统扩展
作为离散马尔可夫链的对照,论文还研究了连续线性动力系统:
其中 , 控制序列的稳定性。
2. 贝叶斯模型平均 (BMA)
2.1 BMA 公式
给定 条观测序列,贝叶斯模型平均通过贝叶斯公式计算父节点 的后验分布:
2.2 Lemma 1:BMA 的闭式表达
在父节点先验均匀()且观测条件独立的假设下:
其中 ,。
这恰好是一个 softmax 形式,为 Transformer 的注意力机制提供了天然的桥梁。
3. 模型架构
3.1 总体架构:两层 Decoder-Only Transformer
论文使用一个简化但可分析的两层 Transformer,核心设计如下:
- 第一层: 头 RPE(相对位置编码)自注意力——作为 “copyer”
- 第二层:单头标准自注意力——实现 BMA 父节点选择
- 解耦残差连接(Disentangled Residual):将每头输出 concat 而非加和
3.2 位置编码
论文采用两种独立的相对位置编码:
- :跨示例(cross-example)的位置,表示 个 in-context 示例间的顺序;
- :示例内(inner-example)的位置,表示 个 token 间的相对偏移。
对第 个示例的第 个 token,相对注意力分数为:
这种分离设计将参数量从 降至 ,使参数级分析可行。
3.3 第一层:RPE Copyer
第一层第 头的输出为:
关键直觉:当 RPE 参数满足 且 时,头 只关注第 个示例中同一位置的 token,从而将 “复制”到残差流中。 个头的输出拼起来就恢复了全部 个历史观测 。
3.4 第二层:单头 BMA 注意力
第二层以解耦残差为输入:
然后计算注意力:
其中 ,,且满足块稀疏结构:
这意味着只有来自第一层的特征 (而非原始 token )参与第二层的 query-key 交互。
4. 构造定理 (Theorem 1)
4.1 参数假设
在以下参数设定下:
Theorem 1:当 时,第二层注意力权重在 时逼近 BMA 后验:
且预测收敛到真实条件分布:
4.2 构造直觉
- 第一层(Copyer):每个头 专门复制第 个示例中同位置 token 。拼合后得到 。
- 第二层注意力分数:对候选父节点 的注意力分数为: 当 时,这恰是 BMA 的 。
- 因果掩码 + Softmax: 正好是 BMA 后验分布。
- OV 矩阵: 将选中的父 token 映射到正确的预测分布。
4.3 算法伪代码
输入: L 个上下文序列 {x^{(l)}_{1:H}}_{l=1..L}, 目标序列前缀 x^{(L+1)}_{1:h-1}
参数: RPE {w^H_k, w^L_k}_{k=1..K}, W_KQ, W_OV
# 第一阶段: Copyer (第一层)
for k = 1..K:
for h = 1..H:
# 用 RPE 注意力从第 k 个示例复制同位置 token
u^k_h = softmax(w^H_k(h,:) + w^L_k(L+1,:)) · x_{1:T}
end
end
# 第二阶段: 解耦拼接
for h = 1..H:
v_h = [u^1_h, ..., u^K_h] # 第一层输出
z_h = [x^{L+1}_h, v_h] # 解耦残差 (含原始 token)
end
# 第三阶段: BMA 父节点选择 (第二层)
for h = 2..H:
# 注意力分数 = 所有候选父节点的 bilinear score
score[h'] = v^⊤_{h'} W_KQ v_h # 对所有 h' < h
attn = softmax(score[1:h-1]) # BMA 后验
# 预测 next token
pred = attn^⊤ · x^{L+1}_{1:h-1} · W_OV
# 或等价于:
# pred = Σ_{h' < h} attn[h'] · π(·|x^{L+1}_{h'})
end
输出: 预测序列 {pred_h}_{h=2..H}
5. 参数验证 (Proposition 1)
训练后学到的权重 与 之间存在一个列-wise 平移不变性:
Proposition 1:如果 的各列与 相差一个列平移:
则注意力的输出与 BMA 完全相同:
验证方法:不应检查行级 softmax ,而应检查列级 softmax ,因为列平移不影响列 softmax。实验显示列 softmax 误差 。
6. 信息论保证
6.1 数据处理不等式 (DPI)
Lemma 2:如果 构成马尔可夫链,则 。
在我们的设定中, 形成马尔可夫链,因此 DPI 成立。
6.2 严格 DPI(Strong DPI)
Lemma 3:若转移核满足 (均匀 minorization),且边际分布 ,则存在 使得:
核心含义:非父节点与 的互信息严格小于真实父节点的互信息,保证因果结构可识别。
6.3 期望对数似然的 DPI
Lemma 4:在 Lemma 3 条件下,若 ,则:
6.4 Theorem 2:因果父节点的可识别性
在 Lemma 4 条件下,实现 BMA 的 Transformer(Theorem 1)在 时注意力权重收敛到真实父节点的 one-hot 指标:
7. 训练动力学 (Theorem 3)
7.1 初始化梯度的因果结构恢复
Theorem 3(非正式):考虑 Theorem 1 构造的 Transformer, 的对角块 可训练,使用交叉熵损失训练。初始化为 (此时模型对所有输入输出平稳分布 )。在初始化处的梯度满足:
即对真实父节点的梯度(幅度)最大。这些梯度项与 -互信息高度相关,通过 DPI 建立了梯度的序关系。
直觉:Transformer 并非将因果结构硬编码到参数中,而是通过中间参数 (注意力分数)从数据中提取结构信息。梯度信号在初始化时就暴露了真实的父节点,使得早期训练就能快速学习因果结构。
8. 连续系统扩展 (Proposition 2)
8.1 Transformer 的表示局限
Proposition 2:在连续动力系统设定下,Transformer(满足 Eq.(7))的 logits 为双线性形式:
而 BMA 的 logits 为:
由于 BMA 包含二次项 (来自高斯分布的归一化常数),而双线性形式无法表示此项,因此不存在 使得 Transformer 对所有样本都能精确匹配 BMA。
8.2 实践中的应对
尽管存在理论上的表示局限,实验表明 Transformer 在 足够大时仍能达到与 BMA 可比的父节点选择性能,但在小样本( 较小时)存在明显的性能差距。
9. 实现注意事项
| 方面 | 说明 |
|---|---|
| 第一层头数应等于 in-context 示例数 ,每头专门复制一个示例 | |
| RPE 初始化 | 中间位置(偏移 0)初始值应显著大于其他位置; 的对角偏好可随机初始化 |
| 对角结构 | 实验中 自然学到块对角结构,非对角项趋近于 0 |
| 的列级 softmax 逼近转移矩阵 | |
| 列平移检验 | 验证 时使用列级 softmax 而非行级 |
| 连续系统 | Transformer 无法精确实现 BMA,但大的 可缩小差距 |
参考文献
- Wei et al., ICLR 2026. How Transformers Learn Causal Structures In-Context
- Nichani et al., ICML 2024. How Transformers Learn Causal Structure with Gradient Descent
- Edelman et al., NeurIPS 2024. The Evolution of Statistical Induction Heads
- Von Oswald et al., ICML 2023. Transformers Learn In-Context by Gradient Descent
- Polyanskiy & Wu, 2023. Information Theory: From Coding to Learning