← 返回 2026-08-10

模块化测试时训练:将内部学习器重构为可组合模块 Modular TTT: Rethinking Test-Time Training as Composable Modules

Bohao Tang, Zhen Qin, Yuqi Pan, Zheng Li, Pengfei Liu, Ya Zhang 📅 2026-08-07 👍 8 2026-08-15 18:30
序列建模 快速权重 测试时训练 消融研究 线性注意力

把测试时训练的内部学习器表示为有向无环图,系统消融发现浅层小学习率方案最优

前置知识

测试时训练(Test-Time Training, TTT)与快速权重

TTT 把序列建模重新表述为在线学习问题:模型的隐藏状态不再是固定向量,而是一个可学习网络(通常是线性层或 MLP)的“快速权重”。模型一边处理序列,一边用自监督损失对这些权重做一两步梯度更新,使状态能够“学习”当前上下文,需要回答时再用查询读取这块权重。

这是本文的全部前提——只有先理解“状态即权重”,才能理解作者为什么要把内部学习器抽象成可组合的计算图。

Train-view 与 Query-view(因果对偶形式)

TTT 的计算天然分成两个“视角”。Train-view 用输入 K 做前向和反向,得到预测 $\hat{V}$ 和参数更新 $\Delta W$;Query-view 用查询 Q 读取更新后的权重,并通过下三角算子 $\mathrm{Tril}(\cdot)$ 保证只看到过去。Modular TTT 把它整理成统一的三趟计算。

这是 Modular TTT 自动组合的核心抽象:每个原语只要注册这两个视角(加上反向)的规则,整个图就能跑通。

分块递归计算(Chunkwise Recurrence)

把逐 token 的递归改写成按 chunk(本文取 256)的分块矩阵运算,利用矩阵乘法结合律先聚合 K-V 对,避免逐 token 更新无法利用 GPU 并行。它在线性注意力、SSM 和 TTT 中都广泛使用,是高效长序列训练的关键工程手段。

实验中的 chunk size=256、吞吐对比,以及深层记忆失败与 chunk 内更新稳定性的分析都建立在分块计算之上。

内部学习规则与学习率参数化

TTT 的“内部循环”用一步或多步梯度下降更新快速权重。本文把学习率参数化为 $\eta_t = 2\,\mathrm{Sigmoid}(\beta_t + b)$,其中 $\beta_t = x_t^\top w$ 由外部 token mixer 预测,$b$ 控制初始量级,只在快速节点注入。

small-lr init 是本文最重要的发现之一,理解学习率如何影响更新矩阵 $I - \hat{K}^\top \mathrm{diag}(\eta)K$ 的特征值,才能看懂为什么 $\eta_0\approx 1$ 会不稳定。

权重衰减(scalar / vector decay)

对过去写入快速权重的贡献乘以一个遗忘因子:scalar decay 是全局标量收缩,vector decay 是逐特征的衰减矩阵。它相当于线性注意力里的 decay 机制,让模型能够遗忘陈旧上下文。

decay 消融是核心结论之一——scalar decay 以近乎零开销拿到大部分增益,直接决定了后续大模型实验的配置选择。

研究动机

TTT(测试时训练)把序列建模重新表述为在线学习:模型的隐藏状态本身就是一个可学习网络的“快速权重”,随序列推进被内部学习规则不断更新。这条路线已经催生了 TTT-Linear、TTT-MLP、LaCT、Titans 等多个变体。但问题在于,这些变体几乎都是各自手工硬编码(hard-code)的前向/反向计算。这带来两个实际困难:其一,开发新变体很困难——改动一个变体往往会同时牵动快速权重网络、损失函数、学习率、权重衰减、归一化等多个组件,无法“只换一个零件”;其二,各组件的作用被掩盖——当多个设计选择同时变化时,很难判断到底是哪一项带来了性能差异。例如官方 TTT、LaCT、Titans 各自捆绑了不同的拓扑、损失和归一化方案,至今没有人在统一框架下把快速权重网络、损失函数、学习率初始化、衰减、归一化作为独立维度逐项隔离地分析过,导致 TTT 的设计空间整体上是“一团迷雾”。此外,朴素的 token 级递归实现无法利用 GPU 并行,现有的分块实现也各自为政,缺乏可复用的工程底座。

本文的目标是本文有两个层次的目标。框架层面,作者希望构建一个统一的模块化框架 Modular TTT,能够把 TTT 的内部学习器表示为一个由可注册原语(Linear、Gate、Norm、Act、Add、Mul 等)组成的有向无环图(DAG),并自动组合出完整的图级 TTT 计算(含快速权重状态转移),从而免去为每个新拓扑手工推导全局更新规则。研究层面,作者希望在骨干网络、数据、训练预算、实现、评测设置全部对齐的条件下,系统性地隔离并量化 TTT 各主要组件(损失函数、学习率初始化、衰减、非线性、深度、残差/门控)的贡献,最终用得到的最佳配置在 410M 和 1.45B 参数规模、100B tokens 上训练,与强基线 Gated DeltaNet(GDN)、LLaMA、LaCT 做可比性验证,并实现比官方 TTT 实现更高的训练吞吐量。

与已有工作不同的是,以往工作的切入点都是“再发明一个具体变体”,把精力放在手工推导某种新拓扑的前向/反向公式上。本文的独特视角是:与其继续堆叠变体,不如先把 TTT 的内部学习器抽象成计算图,把“学习器做什么”和“用什么规则更新”彻底解耦。关键洞见是:对训练视角损失做自动微分就能得到局部反向信号,而因果查询读出和快速权重状态转移只需在原语级别单独定义——因此只要给每个原语注册 train-view forward、train-view backward、query-view forward 三条规则,整个图的 TTT 计算就能自动组合。这种“原语级三视角注册 + 图级自动组合”的视角,把 TTT 设计从一个个定制实现变成可系统性探索和分析的模块化空间,是此前 TTT-Linear、TTT-MLP、LaCT 等工作都没有提供的能力。

核心方法

直觉上,可以把 TTT 的内部学习器想象成一块“可塑的记忆橡皮泥”:每来一段上下文,就用自监督损失在上面捏几下(更新权重),需要回答时再用查询去读这块橡皮泥。Modular TTT 的做法是把这块橡皮泥的结构显式画成一张计算图:节点是原语操作(带快速权重的 Linear、门控 Gate、归一化 Norm、激活 Act、加法 Add、乘法 Mul),边是张量依赖。整张图执行标准的“三趟计算”:第一趟 train-view forward,按拓扑序跑前向得到预测 $\\hat{V}$;第二趟 train-view backward,算自监督损失 $L(\\hat{V}, V)$ 并反向,得到各节点的伴随梯度和局部参数更新 $\\Delta\\theta_j$;第三趟 query-view forward,用查询 Q 跑一遍前向,但每个节点的计算还要额外依赖 train-view 的激活、梯度和参数更新(这才是真正的 TTT 前向)。这三趟由 Algorithm 1 统一编排,自动微分负责 train-view 的局部反向,因果查询读出和快速权重状态转移由各原语注册的规则负责,节点数和拓扑可任意重组。

核心创新在于把 TTT 的设计从“逐变体手工推导”提升为“原语级规则注册 + 图级自动组合”。最本质的区别是:以前每提出一个 TTT 变体(比如从 Linear 到 MLP),作者都得重新手工推导一整套新的全局前向、反向和分块双形式更新公式;而 Modular TTT 只需定义单个原语在 train-view forward、train-view backward、query-view forward 三个视角下各自做什么(见表 1),框架就按拓扑序自动把它们组合成任意 DAG 的完整 TTT 计算,包括快速权重的因果状态转移 $W = W - dW$。关键洞见是:train-view 的反向信号可以由自动微分直接给出(它作用在已确定的学习器图上),只有因果查询读出和快速权重状态转移需要原语级专门定义。这套抽象让“换损失”“换衰减”“加一层非线性”“换成 MLP”都变成图里增删节点的事,从而第一次把 TTT 的设计空间变成可逐项控制、可独立消融的模块化空间,而不只是又一篇文章里又一个新拓扑。作者特别强调 Modular TTT 不应被视为又一个具体的 TTT 变体,而是一个统一框架。

方法步骤详情

完整流程按 Algorithm 1 的三趟执行,输入为 key K、query Q、target V、拓扑序 $\tau$、输出节点 o、损失函数 L,以及每个节点 j 的三类算子 $\{\phi^{\mathrm{train}}_j, \phi^{\mathrm{train\text{-}bwd}}_j, \phi^{\mathrm{query}}_j\}$ 和参数 $\theta_j$。Pass 1(train-view forward):从 $h^{\mathrm{train}}_{\mathrm{input}}\leftarrow K$ 起,按拓扑序计算 $h^{\mathrm{train}}_j=\phi^{\mathrm{train}}_j(\{h^{\mathrm{train}}_i\};\theta_j)$,得到 $\hat{V}=h^{\mathrm{train}}_o$。Pass 2(train-view backward):算 $l=L(\hat{V},V)$,令伴随梯度 $\bar{h}^{\mathrm{train}}_o=\partial l/\partial\hat{V}$,按逆拓扑序调用 $\phi^{\mathrm{train\text{-}bwd}}_j$ 把回传梯度累加到父节点并产出参数更新 $\Delta\theta_j$。Pass 3(query-view forward):用 Q 跑前向,每个节点计算 $h^{\mathrm{query}}_j=\phi^{\mathrm{query}}_j(\{h^{\mathrm{query}}_i\},\{h^{\mathrm{train}}_i\},\bar{h}^{\mathrm{train}}_j,\theta_j,\Delta\theta_j)$,即同时用到查询输入、train-view 激活、梯度和更新,最后输出 $\hat{O}=h^{\mathrm{query}}_o$。对单个 Linear 原语,query-view 规则即分块双形式 $O=Q W_{\mathrm{start}}-\mathrm{Tril}(Q\hat{K}^\top)d\hat{V}$、$W_{\mathrm{end}}=W_{\mathrm{start}}-\hat{K}^\top d\hat{V}$,其中 $\hat{K}=\eta\odot K$ 是注入学习率与衰减的缩放键,chunk size 取 256 平衡效率与性能。

技术新颖性

和已有技术的区别体现在三点。第一,相对于 TTT-Linear、TTT-MLP、LaCT、Titans 这些“单一硬编码变体”,Modular TTT 不是又一个变体,而是一个能表达、实现和分析整个 TTT 设计空间的统一框架——作者明确强调“不应把它看作又一个具体的 TTT 变体”。第二,相对于自动微分本身,AD 只负责 train-view 损失的局部反向,因果查询读出 $O=Q W_{\\mathrm{start}} - \\mathrm{Tril}(Q\\hat{K}^\\top)d\\hat{V}$ 和快速权重状态转移仍需原语级专门注册,这是把 AD 与 TTT 因果双形式结合的新机制。第三,相对于线性注意力的分块实现,本文把快速权重网络、损失、学习率、衰减、归一化都提升为一等公民的显式设计维度,并配备了手写解析反向算子(Linear 快 1.65×、Norm 快 2.62×),使组合出的图能进入 torch.compile 的编译图,端到端实现 2.2×–3.3× 的吞吐提升。

Linear, MLP, and gated learners viewed as composable TTT memory forms; training throughput vs official TTT at 160M; ablations over key design choices.
Figure 1: Linear, MLP, and gated learners viewed as composable TTT memory forms; training throughput vs official TTT at 160M; ablations over key design choices.
Modular TTT: graph-structured memory with a shared three-pass computation.
Figure 2: Modular TTT: graph-structured memory with a shared three-pass computation.

实验结果

消融在 160M/410M、10B tokens 上展开。损失函数:MSE 与内积是仅有的两个有竞争力选择,L1 与 RMSE 显著更差——160M 上 MSE 3.0380,而 L1 高达 3.2665、RMSE 3.0658;因为 MSE 保留残差幅度、内积直接用目标作写信号,而 L1 只留符号、RMSE 归一化残差,削弱了写入($\Delta W=\hat{K}^\top d\hat{V}$)。学习率初始化:TTT 极敏感,small-lr init($\eta_0\approx 10^{-3}$)相比标准 init($\eta_0\approx 1$)大幅降损,MSE+无衰减 160M 从 3.6012 降到 3.2005,410M 从 3.4036 降到 2.9343,用更新矩阵 $I-\hat{K}^\top\mathrm{diag}(\eta)K$ 的特征值解释稳定性。衰减:scalar 以近乎零开销恢复 vector 大部分收益(160M MSE:none 3.2005/scalar 3.0380/vector 3.0038),vector 要付约 25% 吞吐下降和约 3GB 更高显存。非线性:Linear 后加 GELU/SiLU 一致小幅增益(3.0380→3.0205),Norm 不稳定。深度:深层不超越浅层前沿,Linear-SiLU(3.0205) 强于所有两层及更深配置,含 Norm 深层变体甚至发散,作者用因子化耦合论证同一有效权重 $W=W^{(1)}W^{(2)}$ 的不同分解诱导差异巨大的更新方向。效率:解析算子让 Linear/Norm 反向快 1.65×/2.62×,端到端比官方 TTT 快 2.2×–3.3×。放大到 1.45B、100B tokens,内积线性变体损失 2.3150、多项选择均值 42.07,与 GDN(2.3042/41.30) 相当;但 containment 与 RULER 检索明显弱于 LLaMA,8k niah_single 约 21% vs LLaMA 92.13%。

Primitive operators and loss functions used in Modular TTT.
Table 1: Primitive operators and loss functions used in Modular TTT.
Experimental results for TTT loss functions.
Table 2: Experimental results for TTT loss functions.
Experimental results for TTT decay ablation.
Table 4: Experimental results for TTT decay ablation.
Deep TTT variants at 160M.
Table 6: Deep TTT variants at 160M.
Throughput validation in two representative 160M norm-containing settings.
Table 8: Throughput validation in two representative 160M norm-containing settings.
Large-scale downstream evaluation at 410M and 1.45B.
Table 9: Large-scale downstream evaluation at 410M and 1.45B.
Family-level niah_single RULER results at 410M and 1.45B.
Table 31: Family-level niah_single RULER results at 410M and 1.45B.
查看结构化数据
任务指标本文基线提升
损失函数消融(160M,small-lr init + scalar decay) 验证损失 MSE 3.0380 / 内积 3.0383 L1 3.2665 / RMSE 3.0658 MSE、内积显著优于 L1/RMSE,两者几乎打平
学习率初始化(160M,MSE+无衰减) 验证损失 small-lr init 3.2005 标准 init(η0≈1)3.6012 降损 0.40
衰减消融(160M,MSE) 验证损失 / 吞吐 scalar 3.0380(tgs 105118) vector 3.0038(tgs 79024)/ none 3.2005 scalar 以近乎零开销拿到大部分增益
深层快速权重网络(160M) 验证损失 浅层 Linear-SiLU 3.0205 两层及以上 3.11–3.17,含 Norm 发散 浅层即前沿,加深无益
训练吞吐(160M) tokens/s Modular TTT 93366.9(TTT-Linear)/ 71101.1(TTT-MLP) 官方 TTT 43397.5 / 21336.9 2.2×–3.3× 加速
1.45B 大规模下游(100B tokens) 训练损失 / 多项选择均值 内积线性 2.3150 / 42.07 GDN 2.3042 / 41.30;LLaMA 2.3041 / 40.43 与 GDN 相当,多项选择略优
长上下文精确检索(RULER niah_single 8k) 准确率 1.45B MSE 变体约 21% LLaMA 1.45B 92.13% 仍明显落后(局限)

局限与改进

作者明确承认几点。第一,研究范围仅限自回归语言建模,未穷尽所有内部学习器或优化调度,更深的、门控的学习器在不同更新规则、chunk size、优化器或混合注意力设计下可能表现不同。第二,当前的高效融合图实现无法容纳 LaCT 的 Muon 风格更新归一化和动量变体,因而不在框架内;作者检查发现这些精修未带来明确质量提升。第三,也是最关键的实证局限:选出的浅层 Modular TTT 变体在 containment 类任务(SWDE/SQuAD/FDA)和显式长上下文检索(RULER NIAH)上仍明显弱于 LLaMA,尤其 8k 上下文,说明固定状态的 TTT 在精确召回上仍有硬伤。我自己的观察是:所有消融都基于一步(one-step)内部更新,且评测上下文(2k–4k)相对短,很难断言结论是否能外推到真正的长上下文(如 32k+)或多步更新的 LaCT 设定;同时 containment 任务的全面落后暗示“快速权重压缩记忆”范式在需要精确逐字复现时存在结构性瓶颈,这正是 TTT 这条路线相对全注意力的根本短板。

独立分析的弱点

独立分析几个弱点并给改进方向。第一,“深即不好”的结论是在一步更新设定下得到的,深层记忆的因子化耦合困难可能被多步更新或专门的内循环优化器(如 LaCT 的 Muon)缓解——改进方向是把更新调度本身也做成可注册原语。第二,所有评测上下文最长 8k,而真实长上下文场景常需 32k–128k,论文的衰减与召回结论无法直接外推——应补充 32k+ 的 needle-in-a-haystack 与长文档 QA 评测。第三,长上下文精确召回明显落后 LLaMA(8k niah_single 约 21% vs 92%),说明纯快速权重压缩不足以精确存储——改进方向是引入混合架构(TTT 层 + 少量稀疏注意力层)或可寻址的外部缓存。第四,当前融合实现排除了动量/Muon 类更新精修,限制了与 LaCT 等最新工作的公平对比——应扩展融合内核支持更丰富的更新族。第五,结论仅在英文自回归 LM 上验证,跨模态(视觉、语音)和双向模型的迁移性未知,框架虽支持更多拓扑但尚未实证。

未来方向

作者明确提出的方向包括:研究更丰富的更新调度、更好的检索导向记忆机制,以及在语言建模之外的更广评测。基于本文成果可延伸的方向:把更新调度(多步内循环、Muon/动量)作为新的可注册原语纳入图,从而把 LaCT 等也统一进来;探索 TTT 层与稀疏注意力层的混合,专门针对精确召回短板;把 Modular TTT 的模块化思路扩展到视觉(论文提到 TTT 已有视觉变体)和语音等模态;利用框架的低实验成本特性,做更大规模的设计空间搜索(如神经架构搜索自动找最优图拓扑);进一步把框架与硬件内核协同设计,争取在更长上下文上同时保住吞吐与召回。

复现评估

复现友好度较高。作者已公开代码仓库 github.com/ByteDance-Seed/Modular-TTT,基于 PyTorch 实现、Flame 训练框架、lm-eval-harness 评测,用 OpenAI GPT-2 BPE 分词器(词表 50257)。论文给出完整的架构与训练超参表(Table 11–13):160M/410M/1.45B 的 dmodel、层数、头数、dffn,chunk size=256,scalar 衰减初值约 0.99,small-lr 初值约 10⁻³,快速权重高斯初始化 std 0.02。消融阶段只需 160M/410M、10B tokens,单卡或小集群即可复现主结论,门槛较低;放大到 1.45B/100B tokens 则需较大算力。主要不确定性在于:分块融合内核的具体实现细节、英文预训练语料的成分未完全公开,以及 LaCT 的 Muon 类更新未纳入对比。总体复现难度中等偏低,框架本身对二次开发非常友好。