Back to Blog Untitled

Untitled

Paper

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 学到后者的那种表征。

Model comparison
图:不同预测机制的对比。MTP/JTP 只在 token 空间做监督,隐状态是隐式的。NextLat 显式训练模型预测下一个隐状态 ht+1,再从 ht+1 预测 token。

NextLat 方法:极简但有理论深度

核心目标

NextLat 在标准 next-token loss 之外,加了一个 latent dynamics model pψ(一个简单的 MLP),让它从当前隐状态 ht 和下一个 token Xt+1 预测下一个隐状态 ht+1。训练目标由三部分组成:

NextLat = ℒnext-token + λnext-h · ℒnext-h + λKL · ℒKL
其中 ℒnext-h 用 Smooth L1 回归对齐隐状态,ℒKL 用 KL 散度对齐 token 预测分布。推理时 pψ 完全不使用,零额外开销。

关键设计:stop-gradient 应用在目标隐状态上(防止表征坍缩)。多步预测 d>1 只是为了提供更丰富的梯度信号,理论上 d=1 就足以保证 belief state 收敛。

理论保证

论文证明了一个核心定理(Theorem 3.2):在 NextLat 训练下,模型的隐状态会收敛到 belief states——即足以预测未来的最小充分历史摘要。这个收敛性和预测步长 d 无关,只需要单步预测最优。

Figure 1: Illustration of NextLat
图 1:NextLat 的核心思想。训练时显式预测下一个隐状态,鼓励形成一致的状态转移动态。

Self-Speculative Decoding:附赠 3.3x 加速

NextLat 的隐状态动态天然支持递归展开。即使只训练了 d=1(单步预测),推理时可以链式推进:

θ→ Xt+1 —pψ→ ht+1 —pθ→ Xt+2 —pψ→ ht+2 → ...

这种递归可组合性使得 NextLat 能在推理时使用任意长度的 speculative draft,远超训练视界 d,而 MTP/JTP 的 draft 长度被训练时的 d 严格限制。
Speculative decoding comparison
图:Self-speculative decoding 对比。MTP 受限于固定 draft 长度 d,NextLat 可以使用可变长度 draft,大幅减少验证循环。

实验:全面碾压

World Modeling:Manhattan 出租车

这是最直观的实验。在曼哈顿随机游走数据集上,GPT 虽然 next-token 准确率 100%,但内部表征的有效秩高达 160.1,序列压缩率只有 0.65。NextLat 的有效秩降到 52.7(压缩了 3 倍),序列压缩率升至 0.71(更接近真实世界模型的 1.0)。

方法Next-Token (%)有效轨迹 (%)序列压缩 ↑有效秩 ↓绕行鲁棒 ↑
GPT10097.00.65160.185.0%
MTP10098.10.6457.795.0%
JTP10097.10.32215.887.0%
NextLat10098.70.7152.795.0%
真实世界模型1001001.00100%
Manhattan NHS visualization
图:曼哈顿随机游走实验中模型学到的内部"地图"(NHS 可视化)。NextLat 学到的地图结构更接近真实曼哈顿路网。

推理与规划

在 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% 接受率),证明隐状态动态在远超训练视界后仍然连贯。

FineWeb speculative decoding results
图:FineWeb-Edu 上 speculative decoding 的加速和接受率随 draft 长度的变化。NextLat 在 draft 长度 10 时仍保持高接受率。

方法对比与定位

GPTBSTMTPJTPNextLat
训练参数 (d=2)1.32B2.57B1.42B1.34B1.40B
推理参数1.32B1.32B / 2.57B1.32B1.34B1.32B
训练速度 (it/s)3.090.892.582.922.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 训练范式创新里最值得跟踪的方向之一。

Tags: #Paper