[ICLR 2025] 边际优先:解析 Transformer 学习歧义消除的阶梯式演进
Marginals Before Conditionals
本文提出了“边际优先于条件”(Marginals Before Conditionals)的学习范式,通过构建一个受控的 K 重歧义消除任务,揭示了 Transformer 在学习条件概率前会长时间停留在边际概率的损失平台(Loss Plateau)。研究发现该转换是一个集体同步的“突变”过程,且受梯度噪声的熵力稳定。
TL;DR
为什么大模型有时明明拥有所有必要的信息,却依然无法给出正确答案?本文通过一个极简的受控实验发现:Transformer 在学习“条件概率”之前,会固执地先学习“边际概率”。模型会精准地卡在 的损失平台上( 为歧义度),直到某种“集体电路”形成。研究最震撼的结论是:梯度噪声不仅无法帮助模型跳出这个平台,反而像一种“熵力”拉住了模型,使其更难向条件预测转化。
1. 痛点:为什么模型“视而不见”?
在自然语言处理中,模型常面临歧义。例如,“苹果”在没有上下文时可能是水果,也可能是科技公司。这就构成了 的边际分布。当我们给出选择器(Selector)(如“乔布斯”),分布应坍缩为确定的条件概率 。
前人研究发现模型存在“逆向诅咒”(Reversal Curse),即模型知道“A 的父亲是 B”,却答不出“B 的儿子是谁”。本文作者认为,这种不对称性本质上是模型在从**边际分布(一对多)向条件分布(一对一)**过渡时的动态阻碍。
2. 核心实验:受控的“风洞”任务
作者设计了一个简洁的任务:
- 输入:字符串 B(6个字符)
- 选择器:z(2个字符,指明 B 对应的 K 个候选者中的哪一个)
- 目标:预测确定的 A。
- 信息论指标:如果模型忽略 z,损失函数将严格等于 (边际损失);如果模型学会利用 z,损失将降至 0。

3. 深度直觉:哪些因素在控制“等待时间”?
3.1 规模 D 决定时长,歧义 K 决定高度
作者发现了一个违反直觉的现象:平台的持续时间 并不随着任务的复杂程度 增加而线性增长。通过控制变量,作者证明了 。这意味着,模型转化的快慢主要取决于它看过了多少样本数据(D),而不是每个样本有多纠结(K)。
3.2 梯度噪声的“稳定效应”
通常认为,增加随机性(如提高学习率 或减小 Batch Size)有助于模型逃离局部最优。但在本任务中,结论截然相反:
- 学习率越高,平台期越长(固定吞吐量下,延迟高达 3.6 倍)。
- Batch Size 越小,逃离越晚。
物理直觉:边际解(损失为 处)在优化景观中是一个极度各向异性的“鞍点”。由于 个方向的梯度在抵消,这里的梯度模长极小。梯度噪声产生的“熵力”倾向于让模型留在这种低梯度的平坦区域,就像在平原上随机漫步的人很难发现极细的下坡小径。

4. 机制分析:同步突变与内部级联
通过 Mechanistic Interpretability 工具(TransformerLens),作者解码了模型内部的变化:
- 突变性:这种从 到 0 的下降不是缓慢发生的,而是所有样本组(Base Groups)几乎同时在某个时刻“顿悟”(Collective Snap)。
- 内部级联:在损失下降前的 50% 时间点,一个被称为“选择器路由头”(Selector-routing head)的特定注意力头已经开始悄悄形成,模型正在后台构建处理 z 的电路。
5. 局限性与行业启示
- 逆向诅咒的解释:研究表明 A→B(无歧义)的训练速度远慢于 (B,z)→A。因为后者可以复用已经学习到的组结构,而前者纯粹靠死记硬背。
- 局限性:该结论目前是在 4 层小 Transformer 上得出的,且任务是合成数据。在超大规模模型和自然语言分布下,这种“边际优先”的现象是否会被更强的归约偏差抵消,仍需进一步验证。
总结:
本论文为我们提供了一个看待模型训练的新视角:模型不是在学习任务,而是在对数据进行降维和消歧。 很多时候,性能停滞不前并不是因为模型“不行”,而是因为它正在熵力的牵引下,在某个 平台上等待足够多的数据来触发电路的集体坍缩。
