心智模型: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 的机制解释——模型并没有显式地”学习”参数来编码结构,而是在前向传播中动态地执行统计推断算法。