TL;DR:给标准 next-token prediction 加一个辅助目标——预测模型的下一个隐状态。理论上证明隐状态收敛到 belief states,实验上全面领先 GPT/MTP/JTP,推理速度 3.3x。不改架构、不改推理流程,只改训练目标。
Transformer 的"记忆膨胀"问题
Transformer 的自注意力机制允许它回溯整段历史来预测下一个 token。这赋予了极强的表达能力,但带来了一个被忽视的副作用:模型没有动力将历史压缩成紧凑的内部状态。
标准 next-token prediction 只在 token 空间做监督,隐状态 ht 的形成是隐式的、不受约束的。结果就是模型学到的是"记住所有细节"的策略,而不是"理解底层规律"的策略。Vafa et al. (2024) 的 Manhattan 出租车实验证明:即使 next-token 准确率 100%,模型内部的"地图"仍然是混乱的——有不可能的街道方向,有立交桥。
论文用了一个精妙的类比:托勒密的地心说能预测观测(next-token 准确),但结构混乱、难以外推;哥白尼的日心说更紧凑、更能推广。NextLat 的目标就是让 Transformer 学到后者的那种表征。
NextLat 方法:极简但有理论深度
核心目标
NextLat 在标准 next-token loss 之外,加了一个 latent dynamics model pψ(一个简单的 MLP),让它从当前隐状态 ht 和下一个 token Xt+1 预测下一个隐状态 ht+1。训练目标由三部分组成:
其中 ℒnext-h 用 Smooth L1 回归对齐隐状态,ℒKL 用 KL 散度对齐 token 预测分布。推理时 pψ 完全不使用,零额外开销。
关键设计:stop-gradient 应用在目标隐状态上(防止表征坍缩)。多步预测 d>1 只是为了提供更丰富的梯度信号,理论上 d=1 就足以保证 belief state 收敛。
理论保证
论文证明了一个核心定理(Theorem 3.2):在 NextLat 训练下,模型的隐状态会收敛到 belief states——即足以预测未来的最小充分历史摘要。这个收敛性和预测步长 d 无关,只需要单步预测最优。
Self-Speculative Decoding:附赠 3.3x 加速
NextLat 的隐状态动态天然支持递归展开。即使只训练了 d=1(单步预测),推理时可以链式推进:
这种递归可组合性使得 NextLat 能在推理时使用任意长度的 speculative draft,远超训练视界 d,而 MTP/JTP 的 draft 长度被训练时的 d 严格限制。
实验:全面碾压
World Modeling:Manhattan 出租车
这是最直观的实验。在曼哈顿随机游走数据集上,GPT 虽然 next-token 准确率 100%,但内部表征的有效秩高达 160.1,序列压缩率只有 0.65。NextLat 的有效秩降到 52.7(压缩了 3 倍),序列压缩率升至 0.71(更接近真实世界模型的 1.0)。
| 方法 | Next-Token (%) | 有效轨迹 (%) | 序列压缩 ↑ | 有效秩 ↓ | 绕行鲁棒 ↑ |
|---|---|---|---|---|---|
| GPT | 100 | 97.0 | 0.65 | 160.1 | 85.0% |
| MTP | 100 | 98.1 | 0.64 | 57.7 | 95.0% |
| JTP | 100 | 97.1 | 0.32 | 215.8 | 87.0% |
| NextLat | 100 | 98.7 | 0.71 | 52.7 | 95.0% |
| 真实世界模型 | 100 | 100 | 1.00 | — | 100% |
推理与规划
在 Countdown(数字组合推理)和 Path-Star(图遍历规划)上,NextLat 全面领先。Countdown 的准确率提升显著,Path-Star 上在复杂度增加时优势更明显。
语言建模 + Speculative Decoding
在 FineWeb-Edu、Wikipedia、Books、Code、Math 五个域上,NextLat 更好地保持了 next-token perplexity(优于 MTP/JTP),同时推理速度远超所有基线。
| 方法 | Wikipedia 加速 | Books 加速 | Code 加速 | Math 加速 |
|---|---|---|---|---|
| JTP (d=2) | 1.90× | 1.90× | 1.88× | 1.89× |
| MTP (d=2) | 1.68× | 1.72× | 1.75× | 1.72× |
| NextLat (d=1) | 2.72× | 2.72× | 2.29× | 2.30× |
| NextLat (d=2) | 3.32× | 3.32× | 2.38× | 2.87× |
注意:NextLat d=1 就远超 MTP/JTP d=2 的加速效果。而且 draft 在长度 10 时仍然完全有效(100% 接受率),证明隐状态动态在远超训练视界后仍然连贯。
方法对比与定位
| GPT | BST | MTP | JTP | NextLat | |
|---|---|---|---|---|---|
| 训练参数 (d=2) | 1.32B | 2.57B | 1.42B | 1.34B | 1.40B |
| 推理参数 | 1.32B | 1.32B / 2.57B | 1.32B | 1.34B | 1.32B |
| 训练速度 (it/s) | 3.09 | 0.89 | 2.58 | 2.92 | 2.79 |
| 梯度信号 | O(T) | O(T²) | O(Td) | O(Td) | O(Td) |
| Belief State 保证 | ✗ | ✓ | ✗ | 有条件 | ✓ (任意 d) |
BST 有理论保证但训练极慢(O(T²) 梯度、两倍参数)。MTP/JTP 训练快但缺乏 belief state 保证(JTP 需要预测步长 d≥k,k 为可观性水平,实际中往往非常大)。NextLat 用最少额外开销获得了最强保证:d=1 就够了。
局限与未来
论文承认了几个局限:latent dynamics 模型只用了简单 MLP(更复杂的架构可能更好)、小规模消融的设计选择(stop-gradient、KL loss 等)缺乏大规模验证、没有完全探索 adaptive-length speculative decoding。
最值得关注的未来方向有两个:一是用 NextLat 做 post-hoc 微调(不从头训练就改善现有模型的推理/规划能力);二是 NextLat 的递归隐状态结构是否更适合 RL post-training(类 Bellman 的 latent structure 有利于 value estimation)。
总结
NextLat 的核心洞察很优雅:Transformer 的表达能力是够的,缺的是正确的归纳偏置。不需要改架构,只需要在训练时告诉它"你的隐状态应该能预测你自己的下一个隐状态",它就会自发地压缩历史、形成一致的世界模型。
方法简单、有理论保证、效果全面、推理零额外开销、还送 3.3× 加速。这种"加一个辅助目标就改变游戏规则"的工作,是 transformer 训练范式创新里最值得跟踪的方向之一。