普林斯顿只在中间两成层加跨步回路,改装1.7B后MATH500从12.8升到18.0

T^2MLR: Transformer with Temporal Middle-Layer Recurrence

Ziyang Cai, Xingyu Zhu, Yihe Dong, Yinghui He, Sanjeev Arora

cs.CL, cs.AI

2026-07-17

普林斯顿的 T²MLR 把上一 token 中间层表征注入当前更浅层。135M 只回路 20% 层,零样本均分 44.14 对 42.83;改装 1.7B 后 MATH500 从 12.8 升到 18.0,解码开销约 8%。

这篇在解决什么

自回归 Transformer 每步都要把高维隐状态压回离散 token,再当作下一步的唯一输入。中间层明明在做更抽象的计算,下一步的浅层却摸不到上一步已经算好的中间表征,只能靠注意力再走一遍深度,或指望那个被压扁的 embedding。Coconut 一类连续思维把回路放在末层或 embedding;Looped Transformer 靠同一块层反复前向来加深度,推理成本随圈数涨。

普林斯顿语言与智能实验室要的是第三条路:让中间层的抽象状态在时间上活下去,同时保持标准自回归接口和接近原模型的逐步解码成本。

方法

T²MLR 在标准 decoder-only 上加一条常尺寸循环缓存 R。指定起止层 ℓstart 到 ℓend。算当前 token 时,先把 R{t-1} 和 ℓstart 之前的表征用门控模块 Φ 融合,再走中间块;ℓend 之后用 RMSNorm(ht^{ℓend} + R{t-1}) 写成 Rt,留给下一步。Φ 里有两组可学习标量门和按位置的 sigmoid 调制,γ 可从零初始化,训练前期更稳。

训练不能直接做 token 并行,因为 R 依赖上一步。做法是 Jacobi 式定点迭代:先假设没有循环缓存跑一遍,把 ℓend 表征当作初值,再只在中间块上迭代 dforward=16 次逼近 R;反向用 dbackward=4,类似截断 BPTT。推理时逐步解码只多一个融合模块,论文测到逐步生成相对开销最多约 8%,随模型和长度增大还会被注意力盖住。

结果

合成任务 S5-Retrieval 同时要求置换群状态追踪和上下文检索。4 层 LSTM 和 4 层 Transformer 在精确匹配上垮掉;同样深度的 T²MLR 在训练长度内接近满分,超出训练长度仍有非零 token 准确率。训练步数还更少:T²MLR 15 万步,基线 40 万步。

135M、10B FineWeb-Edu、参数对齐(基线 hidden 576 调到 584,略占便宜)的零样本均分:

配置均分
Transformer42.83
T²MLR 全层 D=3043.36
T²MLR (13,18) D=644.14
Pause-token ×243.31
Full-looped ×242.99
Middle-looped ×342.68

只回路约 20% 中间层比全层回路更好。位置消融在固定宽度 D=6 和 D=14 时,中间块都压过同样宽度的最浅块和最深块。

下游多跳和小学数学上,中间层变体也高于全层回路。规模放大后推理任务相对涨幅更明显:361M 上 HotpotQA-Easy 24.43→28.28(+15.8%),GSM-Aug 自然语言 31.08→34.12;1B 上 HotpotQA 23.28→26.52(+14.0%)。预训练拉到 50B token,361M 零样本均分 52.83→54.78。

不必从零预训练。把循环通路加到现成的 SmolLM2-1.7B-Instruct 上,在 OpenMathReasoning 上继续微调一个 epoch:GSM8K 35.78→39.88,MATH500 12.80→18.00。

为什么重要

对想加强推理、又不想上循环深度或拉长思维链的人,这是一条可改装的路径:KV cache 还在,解码接口不变,逐步开销按论文测量不超过约 8%。中间层比全层更有效,说明「回路该长在哪」比「要不要回路」更关键。

但别把它读成训练算力胜利。Jacobi 近似让预训练墙钟大约慢 2 到 4 倍。同一 135M 设定下,把 Transformer 训到匹配的 2.24 epoch,零样本均分 45.30,已经超过 T²MLR 的 44.14。论文自己把主对比定位在参数、数据和推理算力对齐,不是训练墙钟对齐。

局限与存疑

训练贵是作者承认的主代价,未来工作也指向减少迭代或在 on-policy RL 里直接复用精确循环状态。主实验仍偏小:从零预训练到 1B,改装到 1.7B,没有多种子方差。S5-Retrieval 的精确数字只在图上。Coconut 式 embedding 回路没有直接对比,理由是那些方法缺可扩展的稠密 teacher forcing。下游若干任务用了加难版(变量赋值深度 5、ProsQA 平均 60 节点),和原论文数字不能横比。

术语

原文与代码

社区讨论

相关论文

全部论文解读