论文
测试时学习
Learning to (Learn at Test Time): RNNs with Expressive Hidden States
把 RNN 隐状态变成一个在测试时持续自监督训练的小模型(TTT 层)
Yu Sun, Xinhao Li, Karan Dalal, et al. · Stanford / UC San Diego · arXiv · 2024-07 · 被引 340
一句话Stanford/UCSD 提出 TTT 层:隐状态本身是一个模型(线性或两层 MLP),每读一个 token 就对自监督重建损失做一步梯度下降,推理时也在"训练"。125M-1.3B 规模上,TTT-Linear/TTT-MLP 在 8k-32k 长上下文 perplexity 上超过 Mamba,且能像 Transformer 一样随上下文变长持续降低 perplexity(Mamba 在 16k 后失效);ICML 2025 接收,引用 340+,是 Titans 等"test-time memorization"路线的直接前驱。
这是什么
线性复杂度的 RNN(Mamba、RWKV 等)在长上下文里有一个根本瓶颈:固定大小的隐状态存不下越来越多的历史信息。论文的观察是,把上百万 token 压缩进固定容量、又能保留内在结构,这件事机器学习本身就在做——参数化学习就是把训练集压缩进模型权重。于是他们提出:让隐状态 s_t 直接等于一个小模型 f 的权重 W_t,update rule 是对自监督损失 ℓ 做一步梯度下降 W_t = W_{t-1} − η∇ℓ(W_{t-1}; x_t),output rule 是 z_t = f(x_t; W_t)。因为这个内层训练在测试序列上也照常发生,所以叫 Test-Time Training(TTT)层。
论文给出两个实例:TTT-Linear(f 是线性模型 + LayerNorm + 残差)和 TTT-MLP(f 是两层 MLP)。TTT 层接口与 self-attention/RNN 层完全一致,可以直接替换进 Transformer 或 Mamba backbone;外层网络照常用 next-token prediction 训练,内层的自监督任务(低秩投影出的 K/V/Q 三个 view)本身也是外层学出来的。这就形成一个干净的双层结构:外层是 meta-learning(学怎么学),内层是 test time 的在线学习。
对课题组来说,这篇的意义在于它把"学习发生在测试时"从一种 trick(经典 TTT 用于分布偏移)升级成了一种序列建模的架构原语:遗忘/记忆由梯度大小自然决定(梯度大的输入被记住),上下文越长内层模型学得越好。Google 的 Titans(2025-01)明确沿此路线扩展(加动量和遗忘门),后续还有 TTT 做一分钟视频生成等工作。
TTT 层的核心机制:隐状态就是模型权重 W_t,update rule 是对自监督损失的一步梯度下降 W_t = W_{t-1} − η∇ℓ(W_{t-1}; x_t),output rule 是用最新权重做预测 z_t = f(x_t; W_t)。整个内层训练过程被编进序列层的 forward pass,推理时照常发生。机制与做法
内层自监督任务:可学习的多视角重建
最朴素的内层损失是去噪重建 ℓ(W; x_t) = ‖f(x̃_t; W) − x_t‖²,但论文不手工设计损坏方式,而是把任务本身做成外层参数:训练视角 θ_K x_t(低秩投影当作输入)、标签视角 θ_V x_t(重建目标)、测试视角 θ_Q x_t(输出时用)。即 ℓ(W; x_t) = ‖f(θ_K x_t; W) − θ_V x_t‖²,z_t = f(θ_Q x_t; W_t)。θ_K/θ_V/θ_Q 与 self-attention 的 K/V/Q 参数完全对应——外层训练相当于在一族重建任务里挑一个最利于 next-token prediction 的任务。
稳定性技巧:W_0 也做成可学参数 θ_init(显著改善训练稳定性);内层学习率 η 做成 token 依赖的门控 η(x) = η_base·σ(θ_lr·x)(TTT-Linear 的 η_base=1,TTT-MLP 是 0.1);f 里固定带 LN 和残差。外层反向传播要穿过内层的 ∇ℓ,即 meta-learning 里的 gradient-of-gradient。
Books 数据集上 FLOPs vs perplexity 的规模曲线(125M-1.3B)。左:2k 上下文,各方法基本重叠,Mamba 略优;右:32k 上下文,TTT-Linear(M) 和 TTT-MLP(M) 全规模优于 Mamba,而 Transformer 因二次复杂度 FLOPs 代价大幅右移。上下文越长 TTT 优势越明显,是全文主结果。mini-batch TTT + dual form:让内层训练跑得动
逐 token 的 online GD 无法并行(W_t 依赖 W_{t-1})。论文改用 mini-batch GD:把序列切成大小 b 的段,段内所有梯度都对上一段末尾的 W 求(G_t = ∇ℓ(W_{t'}; x_t)),这样段内 b 个梯度可并行,再用 cumsum 拼出各时刻的 W_t。b 控制质量-速度权衡(b=1 是 online GD 最准但最慢,b=T 是 batch GD 最快但有效搜索空间小),全文取 b=16。
但并行还不够——朴素实现里每个 G_t 是 d×d 的外积,matmul 太少喂不饱 TensorCore,显存 I/O 也重。dual form 的关键观察是:不需要显式物化中间的 G 和 W,只需要段末的 W_b 和输出 z_1..z_b。对 TTT-Linear 可以推出 W_b = W_0 − 2η(W_0X − X)Xᵀ、Z = W_0X − 2ηΔ,其中 Δ 用一个带上三角 mask 的 (W_0X−X)·mask(XᵀX) 算出——形式上和 attention 矩阵很像,全是 matmul。非线性 f(MLP)也有对应的 dual form。这一段是论文工程含量最高的部分,也是它能在 wall-clock 上和 Mamba 打平的原因。
实验设置与主结果
严格沿用 Mamba 论文的评测协议:125M/350M/760M/1.3B 四个规模、Chinchilla 训练配方、Pile(2k/8k)+ Books3(1k 到 32k)。baseline 是 Llama 架构 Transformer 和 Mamba,并验证了自己的 baseline 能复现 Mamba 论文数字。TTT 层默认用 Mamba backbone(带时间卷积),(T)/(M) 标记 backbone 消融。
结论分层很清楚:2k 上下文三家基本打平(Books 2k 上 Mamba 还略好);8k 起 TTT-Linear/TTT-MLP 明显超过 Mamba,且上下文越长优势越大;最关键的图是 perplexity 随 token 位置的变化——Mamba 在 16k 之后不再从更多上下文获益,TTT 和 Transformer 一样能持续下降。速度上:TPU v5e-256 训练迭代 TTT-Linear 0.27s vs Transformer 0.30s;A100 推理 kernel 下 TTT-Linear 的 prefill/decode 每 token 时延不随上下文增长、且低于 Mamba,但 TTT-MLP 因内层状态大、内存 I/O 重,每 token 时延约为 Mamba 的 1.3 倍(论文自己承认这是未解决的问题)。
关键结果
- 长上下文是分水岭:Pile 8k 和 Books 32k 上 TTT-Linear(M)/TTT-MLP(M) 全模型规模优于 Mamba;而 2k 上下文三者打平(Books 2k 上 Mamba 还略优)。上下文越长 TTT 相对 Mamba 的优势越大。
- Mamba 的 perplexity 在超过 16k 上下文后不再随 token 位置下降,TTT-Linear/TTT-MLP 与 Transformer 一样能持续从更长上下文获益——这直接验证了"固定隐状态容量不足、可学习隐状态更能压缩长历史"的核心论点。
- TTT mini-batch size b=16 是质量-速度折中点;b 越小(越接近 online GD)perplexity 越好,1.3B TTT-Linear 在 Pile 2k 的 perplexity 为 11.09。
- 速度:TPU 上训练迭代比 Transformer 快 10%(0.27s vs 0.30s,2k 上下文);A100 推理时 TTT-Linear prefill 每 token 约 1.4e-5s,低于 Mamba 的 2.0e-5s;但 TTT-MLP 约 2.6e-5s,比 Mamba 慢约 30%,内存 I/O 是瓶颈。
- backbone 有讲究:Mamba backbone(带时间卷积)对 TTT-Linear 帮助大;Transformer backbone 下 TTT-MLP 明显好于 TTT-Linear,作者推测隐状态表达力越弱越依赖卷积补局部信息。
- 内层损失确实在降:测试序列上一步 GD 就能把 ℓ(W_{t-1};x_t) 降到 ℓ(W_t;x_t),且 t 越靠后 ℓ(W_t;x_t) 相对 ℓ(W_0;x_t) 改善越大——推理时的"学习"是真实发生的,不只是一种比喻。
实证核查
扎实代码(PyTorch 教学版 + JAX 训练版 + 推理 kernel)和 checkpoint 全部公开,评测协议严格对齐 Mamba 论文且论文对自身短板(TTT-MLP 的 I/O、2k 无优势)写得很诚实;ICML 2025 接收,340+ 引用,被 Titans 等后续工作实质性沿用。主要保留意见是 baseline 只有 Transformer 和 Mamba,没比同期的 DeltaNet/GLA 一系。
论文称"所有实验可用公开代码和数据复现"。
确实开了三个仓库:ttt-lm-pytorch(1.4k star,HF 接口教学实现,README 明说"不建议用它训练")、ttt-lm-jax(463 star,训练代码,2025-11 仍有维护)、ttt-lm-kernels(推理速度基准)。checkpoint 发在 HF Test-Time-Training org(如 ttt-mlp-1.3b-pile-8k,见 pytorch 仓库 issue #31/#22)。issue #16 'Replicating the experiments' 被作者答复后关闭,issues 里没有出现复现不出主结果的投诉,多为理解性提问(#24 关于 W 的含义有 13 条讨论)。
摘要称 TTT 层"线性复杂度 + 表达力强的隐状态",给人两全其美的印象。
论文正文和数据其实更收敛:wall-clock 上只有 TTT-Linear 全面快于 Mamba;TTT-MLP 的 A100 prefill 每 token 时延约 2.6e-5s,比 Mamba(约 2.0e-5s)慢 30% 左右(论文 Figure 15,本地 fig12),作者明确承认 memory I/O 是未解决的挑战。且训练 kernel 没写(只写了推理 kernel),TPU 上 10% 的训练加速是无系统优化的 JAX 实现对比。质量方面 2k 短上下文无优势(Books 2k Mamba 略好)。声称与证据一致,论文没有掩饰这些短板。
论文只与 Transformer 和 Mamba 对比,并称 TTT-Linear 是更优的线性复杂度层。
同期线性注意力工作(DeltaNet、GLA 等)未进入主实验。OpenReview(forum eifW0W0xgt)审稿讨论中有人指出 TTT-Linear 与 DeltaNet 的数学联系,作者在修改稿中补充了对比,并说明带 LayerNorm 的 TTT-Linear 是非线性 RNN、不再等价于 DeltaNet。所以"击败所有线性复杂度方案"不能从本文得出,本文严格证明的是"优于 Mamba 于 8k+ 上下文"。
收录理由称其为 Titans 等后续工作的直接前驱、'学习发生在测试时'的架构化表述。
成立:S2 引用 340+,ICML 2025 poster 接收;Google Titans(arXiv 2501.00663)的 neural memory 模块就是在 TTT 的梯度更新记忆上加动量与权重衰减(遗忘),同团队后续还做了 TTT 生成一分钟视频(arXiv 2504.05298)。这条路线(test-time regression/memorization)已成为 linear attention 之后的活跃分支。
与我们方向的关系
对 continual-learning 方向,这篇给出了一个重要的概念转换:把"在测试分布上持续学习"内化为架构的 forward pass,而不是靠外部的 fine-tune 循环。内外双层的划分很干净——外层(meta)学自监督任务本身(θ_K/θ_V/θ_Q、W_0、η 门控),内层在推理时对每条序列在线做 GD。课题组如果做 test-time adaptation 或 memory 机制,TTT 层的"梯度大 = 值得记住"这一隐式遗忘准则,和 Titans 在其上显式加遗忘门的演化路径,是一条值得对照的设计谱系。
工程上可借鉴的是 mini-batch TTT + dual form 的并行化思路:任何"推理时做梯度更新"的方案都会撞上串行依赖和 matmul 利用率两堵墙,本文的解法(段内对同一 W 求梯度 + 不物化中间权重、用 mask(XᵀX) 直接算输出)是目前最干净的参考实现。复现建议直接用 ttt-lm-jax,PyTorch 版只适合读懂机制。
阅读笔记
读代码注意三个仓库分工:pytorch 版是教学实现(纯 PyTorch,训练慢,官方不推荐);jax 版才是论文实验代码;kernels 版复现速度基准。TTT-Linear 去掉 LN/残差后与 DeltaNet 数学上同源,加了 LN 就是非线性 RNN——看后续对比工作时留意这一点。论文附录 B 有非线性 f 的 dual form 推导,附录 C 把 self-attention 解释为 TTT 框架下的 Nadaraya-Watson 估计器,理论口味的同学可以看。
材料清单
TeX 源码已存档:Raw/learning-to-learn-at-test-time/source/
同类条目