心智模型:Transformer 如何通过上下文学习因果结构
1. 这是什么类型的问题?
因果结构推断 (Causal Structure Inference) — 这是一个关于 Transformer 如何从上下文样例中推断序列元素之间潜在依赖关系的核心问题。
更具体地说,这是一个结构不确定性下的贝叶斯推断问题:
- 输入是一组共享同一潜在因果图 G 的序列(L 个演示样例 + 1 个目标序列)
- 每个 token x_h 只依赖于它的一个”父 token” x_{pa(h)}(前驱中的一个)
- Transformer 必须从演示样例中推断出每个位置的父节点是谁,然后基于推断结果预测目标序列的下一个 token
- 关键难点:父关系是潜在的、在每个上下文之间变化的,模型必须动态适应
这不同于传统的固定依赖结构假设(如 bigram),更接近现实世界中语言句法结构变化的场景。
2. 前提知识
马尔可夫链 (Markov Chain)
- 论文使用马尔可夫链生成序列:每个 token x_h 从转移核 π(·|x_{pa(h)}) 中采样
- 前提:理解马尔可夫性质(给定父节点后,token 条件独立于更早的历史)
- 有两个变体:离散马尔可夫链(有限词表 V)和连续线性动力系统(稠密向量)
贝叶斯推断 (Bayesian Inference)
- BMA (Bayesian Model Averaging):在假设空间上计算后验概率的加权平均,而非选择单一最优假设
- 公式:P(pa(h)=h′ | x¹:ᴸ₁:ᴴ) ∝ P(x¹:ᴸ₁:ᴴ | pa(h)=h′) · P(pa(h)=h′)
- 在均匀先验下简化为 softmax over log-likelihood 累加:σ(Σˡ log π(xˡ_h | xˡ_{h′}))
- BMA 是因果结构推断的统计最优算法
信息论 (Information Theory)
- 互信息 I(X; Y) 衡量两个随机变量之间的依赖性
- χ²-互信息 I_χ² 是互信息的一种推广,可在 f-散度框架下统一处理
- 数据处理不等式 (DPI):若 x→y→z 形成马尔可夫链,则 I(x; z) ≤ I(y; z)
- 强 DPI:在特定条件下,非父节点与当前 token 的互信息严格小于父节点与当前 token 的互信息
Transformer 架构基础
- 自注意力机制:Q·K^T 计算注意力分数,softmax 归一化后加权求和 V
- 因果掩码 (Causal Mask):限制注意力只能关注当前位置及之前的 token
- 相对位置编码 (RPE):注意力分数仅依赖于查询和键之间的相对距离
- 解缠 Transformer:将残差流中每层的输出拼接而非相加,便于参数级分析
3. 与研究图景的关系
ICL (In-Context Learning) → 贝叶斯推理
- ICL 的早期理论工作(Xie et al., 2022; Zhang et al., 2025)将其理解为隐式贝叶斯推断
- 本论文将这一连接具体化:不仅证明 Transformer 可以实现 BMA,还揭示了如何实现
贝叶斯推理 → 因果结构学习
- 大多数 ICL 理论研究假设固定的依赖结构(如 i.i.d. tokens 或固定 bigram)
- 本论文突破这一限制:因果结构本身是潜在变量,需要在上下文中推断
- 这是从”学习固定映射”到”学习如何学习结构”的关键跃迁
与相关工作的关系
| 工作 | 关系 |
|---|---|
| Nichani et al. (2024) | 证明 Transformer 可通过梯度下降编码固定因果结构。本论文扩展为可变结构 |
| Edelman et al. (2024) | 研究马尔可夫链上的 Induction Heads。本论文关注的是父节点选择而非单纯复制 |
| D’Angelo et al. (2025) | 也研究上下文因果学习。本论文提供更通用的 BMA 框架 |
| Von Oswald et al. (2023) | ICL 作为梯度下降。本论文提供不同的机制:BMA 而非 GD |
4. 关键概念映射
问题形式化
- 因果图 G = {pa(h)}_{h∈[H]}:每个位置的父节点索引,形成一个有向树
- 上下文演示:L 个序列 {x¹:ᴸ₁:ᴴ},共享同一潜在因果图 G
- 目标:预测第 L+1 个序列的每个 token,基于演示和历史
Transformer 机制分解
第一层 RPE 注意力(K 头,K = L):
- 每个头学习一个”示例索引”偏好(通过 w_L 参数)
- 每个头从特定的演示样例中”复制”对应的 token x_h
- 输出:每个演示样例中同一位置 h 的 token 被提取出来
解缠残差连接:
- 将第一层的输出拼接:v_h = [u¹_h, …, u
- 将原始 token 也拼接:z_h = [x^L+1_h, v_h]
第二层注意力(单头):
- 查询向量 v_h 与候选父节点的 v_{h′} 进行双线性交互
- 注意力分数 = Σˡ (xˡ_{h′})^T W xˡ_h
- 当 W = log π 时,这正是 BMA 的 log-likelihood 累加
- Softmax 输出即为父节点的后验概率分布
信息论保证
- 强 DPI(引理 3):在转移核有下界的条件下,非父节点与当前 token 的互信息严格小于父节点
- 期望 log-likelihood 不等式(引理 4):E[log π(x_h|x_{pa(h)})] > E[log π(x_h|x_{h′})]
- 一致性(定理 2):当 L→∞ 时,BMA 的注意力权重收敛到独热向量 e_{pa(h)}
- 梯度动力学(定理 3 非正式):梯度初始化时 χ²-互信息驱动模型早期发现因果结构
连续系统的局限
- 在连续线性动力系统中,BMA 的 logit 包含二次项 ||xˡ_{h′}||²
- Transformer 的双线性形式无法表示此项
- 因此存在表示局限:Transformer 无法在连续设置中精确实现 BMA
5. 心智模型图示
输入: [例1: x₁¹,x₂¹,...,xᴴ¹] [例2: x₁²,x₂²,...,xᴴ²] ... [目标: x₁ᴸ⁺¹,x₂ᴸ⁺¹,...,xᴴᴸ⁺¹]
↓
第一层 RPE 注意力 (K 头) ←─────────── 每个头聚焦一个演示样例
↓ 提取对应位置的 token
解缠拼接: [x_hᴸ⁺¹, x_h¹, x_h², ..., x_hᴸ]
↓
第二层注意力 ←─────────── 双线性分数 = Σ xˡ_{h′}ᵀ W xˡ_h
↓ 当 W=log π → BMA
Softmax → 父节点后验分布 P(pa(h)=h′|数据)
↓
预测: π(·|x_{pa(h)}) → 下一个 token 分布
这个心智模型的核心洞察是:Transformer 通过第一层的”复制-拼接”和第二层的”双线性评分”,实现了对因果结构的贝叶斯后验推断。这正是 ICL 的机制解释——模型并没有显式地”学习”参数来编码结构,而是在前向传播中动态地执行统计推断算法。