Rethinking Expressivity and Efficiency in Test-Time Training
Rethinking Expressivity and Efficiency in Test-Time Training
发表时间: 2026-08 · arXiv:2608.21308 (KIT, NUS, Fraunhofer IOSB, University of Bonn)
原文: https://arxiv.org/abs/2608.21308
Zeyun Zhong, Joya Chen, Manuel Martin, Frederik Diederichs, Juergen Gall, Juergen Beyerer
Karlsruhe Institute of Technology (KIT), National University of Singapore, Fraunhofer IOSB, Lamarr Institute for Machine Learning and Artificial Intelligence, University of Bonn
速读
一句话结论 本文提出了一种名为 E²-TTT 的测试时训练方法,通过推导闭式标量核,在保持分块并行训练高效性的同时,精确还原了逐 Token 更新的表达能力,在 8 倍训练长度的外推检索任务中保持了 90% 以上的准确率。
要解决什么问题 测试时训练(TTT)通过在推理阶段持续更新模型权重,把历史信息压缩进固定大小的神经网络状态中以处理长上下文。但现有的 TTT 方法面临表达能力与硬件效率的死结。传统的逐 Token 更新具有极强表达能力,能精细控制每个 Token 的重要性,但受限于严格的时序依赖,无法并行计算,导致训练极慢、硬件利用率极低。近期的分块更新方法(如 LaCT)将序列切成大块并行计算以提速,但代价是过度简化了更新规则,直接把块内的动量和衰减因子做平均。这种粗暴的平均抹除了块内不同 Token 的时序差异和重要性变化,导致模型在处理复杂检索或超出训练长度的序列时,效果大幅衰减甚至完全失效。
怎么做的 核心思路是推导出一个闭式的并行标量核,让分块更新能够精确等价于逐 Token 的时序更新。E²-TTT 沿用分块并行框架,但在块内不再粗暴平均,而是保留逐 Token 的学习率 $\eta_t$、动量因子 $\beta_t$ 和衰减因子 $\gamma_t$。标准的逐 Token 动态更新公式为:
$$ \mathbf{M}_t = \beta_t \mathbf{M}_{t-1} + \eta_t \mathbf{G}_t, \qquad \mathbf{W}_t = \gamma_t \mathbf{W}_{t-1} + \mathbf{M}_t $$为绕开串行计算卡点,作者发现矩阵状态 $\mathbf{W}$ 和 $\mathbf{M}$ 的计算虽昂贵,但动量和衰减因子仅为标量。通过展开递归公式并重排求和顺序,时序动态历史被压缩为标量系数,从而直接从块首状态($\mathbf{W}_0, \mathbf{M}_0$)一步跳跃计算出块尾状态($\mathbf{W}_C, \mathbf{M}_C$)。关键设计由两个标量核构成,通过定义后缀乘积 $\tilde{\beta}_t$、$\tilde{\gamma}_t$ 及累积比率和 $R_t$,块尾状态被精确重写为:
效果如何 作者在 FineWeb-Edu 上从头训练了 340M 和 1.3B 参数模型,数据量为 15B tokens,硬件为 H100 GPU。对比基线包括全注意力 Transformer++、线性注意力 DeltaNet 与 HQLT、状态空间 Mamba2,以及分块 TTT 路线的 LaCT。在上下文检索中,1.3B 的 E²-TTT 取得 43.6% 的平均准确率,显著超越 HQLT(35.5%)和 LaCT(36.7%)。最能说明问题的是长度外推能力。在“大海捞针”测试中,当长度扩展到 16K 上下文(8 倍训练长度)时,E²-TTT 仍保持 93.6% 准确率,而 LaCT 在 4K 后崩溃至接近 0,HQLT 衰减至 25.2%。在 LongBench 真实长文本基准上,E²-TTT 取得 14.1% 的最高平均分,几乎是 LaCT(7.7%)的两倍。代价方面,E²-TTT 的训练吞吐量与高效的 LaCT 基本持平(1.3B 下仅慢 3.3%)。局限在于,分块输出强制块内所有 Token 使用相同权重状态,导致局部因果依赖缺失,需依赖滑动窗口注意力弥补;且 SwiGLU 变体在大批次推理时显存占用略高。
主要贡献
当前的大型语言模型(LLMs)在推理、编码和多模态理解方面表现出色,但其权重在训练后保持冻结状态,仅能依赖累积的记忆(如键值缓存)来处理长上下文任务,这限制了它们在现实场景中从长周期任务中持续学习或跨任务泛化的能力。测试时训练(Test-Time Training, TTT)通过在推理期间进行连续的权重更新来实现长上下文处理,但现有方法在逐Token更新动态的表达能力与逐块(chunk-wise)近似的硬件效率之间难以取得平衡。
为了弥合这一差距,本文提出了 $\mathrm{E^2}$-TTT(Expressive and Efficient TTT)。在采用块起始权重计算梯度的标准近似下,本文推导出了一个闭式(closed-form)的状态转换,能够精确重现逐Token递归在块末尾(chunk-end)的快速权重和动量状态。这使得完全并行的块级训练成为可能,同时保留了先前逐块方法所抛弃的更新规则的时间结构。通过从头训练高达1.3B参数的模型,验证了 $\mathrm{E^2}$-TTT 的有效性。在语言建模方面,它与之前的TTT和混合注意力基准表现相当,而在上下文检索方面则超越了它们。其优势在长度外推方面最为显著:在标准的“大海捞针”密码测试中,在8倍于训练上下文长度的情况下,它仍保持超过 $90\%$ 的准确率。同时,$\mathrm{E^2}$-TTT 能够匹敌高效逐块方法的训练吞吐量,证明了它有效地调和了表达能力与效率。
背景知识
测试时训练(TTT)建模范式。与将历史压缩为激活向量或矩阵值状态的标准RNN或线性注意力模型不同,测试时训练(TTT)【57,Learning to (learn at test time): Rnns with expressive hidden states+2024+arXiv】将隐藏状态建模为序列相关非线性函数 $f_{\mathbf{W}_t}(\cdot) : \mathbb{R}^d \to \mathbb{R}^d$ 的可学习权重 $\mathbf{W}_t$。这些权重通常被称为快速权重(fast weights)【54,Linear transformers are secretly fast weight programmers+2021+ICML】,在训练和推理期间被快速调整以动态存储上下文,而慢速权重(即模型参数)在推理期间固定。TTT过程在每个步骤中分为更新阶段和输出阶段:首先,输入Token被投影为查询($\pmb{q}_t$)、键($\pmb{k}_t$)和值($\pmb{v}_t$)向量,键和值用于通过自监督目标更新快速权重;其次,更新后的快速权重应用于查询以生成输出Token $\pmb{o}_t$。
表达能力强的逐Token更新机制。在标准公式中,更新严格在Token级别发生。在每个时间步 $t$,通过最小化转换后的键 $f_{\mathbf{W}_{t-1}}(\pmb{k}_t)$ 和值 $\pmb{v}_t$ 之间的自监督重建损失 $\mathcal{L}$ 来更新快速权重 $\mathbf{W}_t \in \mathbb{R}^{d \times d}$。设 $\mathbf{G}_t$ 为负梯度,$\eta_t$ 为学习率:
损失函数通常是均方误差。更新后的快速权重 $\mathbf{W}_t$ 随后立即用于计算当前查询 $\pmb{q}_t$ 的输出向量 $\pmb{o}_t$:
带有Mini-Batch梯度的逐Token更新机制。虽然逐Token更新提供了高表达能力,但它强制执行顺序递归,阻碍了并行化,导致硬件利用率低。近期的工作引入了小批量(mini-batch)方法。序列被划分为索引为 $r$ 的小批量,每个批量大小为 $B$(如16)。在批量 $r$ 内,梯度在固定权重 $\mathbf{W}_B^{[r-1]}$ 处近似,从而解耦了前向传播的依赖关系:
这种近似允许并行计算所有Token的梯度 $\mathbf{G}_t^{[r]}$,随后可以通过并行关联扫描等方式计算快速权重。
高效的逐块(Chunk-wise)更新机制。尽管允许并行梯度计算,小批量大小和逐Token的依赖性仍然导致硬件利用率低。为了解决这个问题,LaCT【69,Test-time training done right+2025+arXiv】过渡到大块公式。序列被划分为大小为 $C \gg B$(如 $C=512$)的块。在第 $r$ 个块内,LaCT利用微分的线性来聚合优化目标,计算累积加权损失的梯度,为整个块产生单一更新方向 $\mathbf{G}^{[r]}$。当引入动量时,LaCT采用平均时间依赖动量因子 $\beta_t^{[r]}$ 的简化策略来维护块级动量缓冲区 $\mathbf{M}^{[r]}$:
这种公式允许维护块级状态,将物化的快速权重状态数量从 $C$ 减少到每块1个。因为新的权重 $\mathbf{W}^{[r]}$ 仅在处理完整个块后才可用,所以输出 $\mathbf{O}^{[r]}$ 是使用前一个块的权重计算的:
方法细节
解决表达能力与效率的二分困境。当前的TTT变体呈现出明显的权衡:逐Token更新提供高优化保真度但受制于顺序瓶颈,而逐块更新通过简化更新规则实现硬件效率,却模糊了块内Token重要性的时间变化。为了解决这种二分法,本文引入了 $\mathrm{E^2}$-TTT。核心贡献是一种新颖的理论公式,它将带有耦合动量和衰减的逐Token递归映射为闭式的、块并行的更新,在不近似逐Token动态的情况下重现其块末状态。这保留了每个梯度贡献的精确逐Token时间加权,同时匹配了粗粒度块处理的硬件效率。在下文中,省略了块索引 $r$,并将块内的步骤索引为 $t \in \{1, \ldots, C\}$,$t=0$ 表示从前一个块继承的状态。
原始公式:定义目标逐Token动态。与简化的块平均不同,本文旨在强制执行精确的、具有耦合L2衰减和动量的逐Token递归。设 $\mathbf{W}_t, \mathbf{M}_t, \mathbf{G}_t \in \mathbb{R}^{d \times d}$ 分别表示快速权重、动量和负梯度。设 $\gamma_t, \beta_t \in (0,1)$ 为标量衰减和动量因子,$\eta_t > 0$ 为学习率。目标顺序递归定义为:
$\gamma_t, \beta_t$ 和 $\eta_t$ 的具体参数化独立于 $(\mathbf{W}, \mathbf{M})$ 从输入中预测。与小批量设置一样,梯度 $\mathbf{G}_t$ 使用前一个块的固定权重进行评估,这使它们在 $t$ 上解耦,并允许在更新步骤之前并行计算。实际传播到下一个块的状态是块末状态 $(\mathbf{W}_C, \mathbf{M}_C)$;所有中间状态不被存储。
通过标量核实现闭式并行化。朴素地执行递归需要顺序迭代,且存储完整序列会产生 $O(C \cdot d^2)$ 的过高内存成本。为了在不简化动态的情况下实现高效训练,本文推导出一个闭式等价物,它在单步中从块起始 $t=0$“跳跃”到末尾 $t=C$。核心见解是,虽然矩阵状态昂贵,但衰减和动量因子是标量。通过展开递归并重新索引求和,可以将整个时间动态历史压缩为隔离每个梯度对最终状态贡献的标量系数:
梯度项中的内部求和封装了第 $t$ 个Token对最终块权重的累积影响。
定义高效后缀乘积与比例和。为了形式化上述过程,定义了高效的后缀乘积 $\widetilde{\beta}_t, \widetilde{\gamma}_t$ 和比例和 $R_t$:
约定 $\prod_{i=C+1}^C (\cdot) = 1$。这些项仅依赖于标量,因此可以通过对数空间累积和以可忽略的成本并行计算。
命题3.1:提取精确的块末状态。在梯度 $\mathbf{G}_t$ 在冻结的块起始 $\mathbf{W}_0$ 处评估且标量独立于状态的条件下,块末状态满足:
两次聚合共享单次反向传播。计算两个聚合通常需要两次反向传播。本文注意到这两个核仅在标量系数上有所不同,并且共享逐Token激活梯度 $\pmb{g}_t = -\partial \mathcal{L}_t / \partial f_{\mathbf{W}_0}(\pmb{k}_t)$。这些梯度作为LaCT反向传播的 $C \times d$ 中间产物出现。本文在两次聚合中保留它们,而不是在第一次聚合后丢弃。因此,相较于LaCT的边际成本仅为一次重新加权的聚合。算法1直接实例化了这一过程,保留了LaCT的 $O(C \cdot d^2)$ 渐近时间复杂度。输出使用前一个块的权重计算 $\mathbf{O}^{[r]} = f_{\mathbf{W}^{[r-1]}}(\mathbf{Q}^{[r]})$,确保完全并行的前向传播。
# Algorithm 1
Require: chunk-start state W0, M0; tokens (qt, kt, vt); scalars ηt, βt, γt
1: ot = f_W0(qt), Lt = L(f_W0(kt), vt) # Compute outputs and per-token losses
2: Compute scalar aggregates β_tilde_t, γ_tilde_t, Rt via log-space cumulative sums; form KWt, KMt
3: gt = -∂Lt / ∂f_W0(kt) # Per-token activation gradients (single backward)
4: ΔW = AGGREGATE({KWt}, {gt}, {kt}) # ≡ -Σ KWt ∇W Lt|W0
5: ΔM = AGGREGATE({KMt}, {gt}, {kt}) # ≡ -Σ KMt ∇W Lt|W0; reuses {gt} from line 3
6: WC = γ_tilde_0 W0 + β_tilde_0 R1 M0 + ΔW
7: MC = β_tilde_0 M0 + ΔM
8: return outputs {ot}, next-chunk state (WC, MC)
混合模型架构设计。块级处理的一个特征是输出步骤将固定状态应用于整个块,这使得TTT路径对当前块内的直接因果历史“失明”。为了解决这个问题,本文采用混合设计模式,将 $\mathrm{E^2}$-TTT 模块(用于建模非局部依赖)与标准的滑动窗口注意力(SWA)模块(用于捕获局部依赖)配对。宏观架构遵循标准Llama设计【59,Llama 2: Open foundation and fine-tuned chat models+2023+arXiv】。输入序列投影为共享查询、键和值,馈入两个并行分支。对于TTT路径,对查询和键应用SiLU激活和L2归一化。同时,网络直接从输入预测更新规则所需的标量。输出通过依赖于数据的门控进行元素级插值融合。
依赖于输入的动态参数化。为了允许模型在Token级别调节其学习动态,更新系数被参数化为输入 $\pmb{x}_t$ 的函数。学习率 $\eta_t$ 由基础学习率 $\eta_{\mathrm{base}}$ 按比例缩放的Sigmoid门控值参数化:
为了鼓励长期记忆保留,使用时间尺度因子 $\tau$ 将动量偏向接近1的值:
快速权重网络实例化。框架支持任何神经网络作为快速权重模块 $f$。本文使用两种标准MLP配置实例化:GELU MLP和SwiGLU MLP,均包含残差连接和层归一化(LN)。GELU MLP的输入前向转换为:
SwiGLU变体采用门控机制:
实验环境
- 数据集:从 HuggingFace FineWeb-Edu 数据集中提取 15B 个Token用于从头训练。
- 模型架构:训练了 340M 和 1.3B 可训练参数的模型(序列长度分别为 2048 和 2240)。所有模型具有 24 层,子词单元词汇表大小为 32K。LaCT 和本文模型均使用 512 的块大小。
- 硬件配置:实验在 H100 GPU 上进行,340M 和 1.3B 模型分别需要约 132 和 348 个 GPU 小时。
- 软件配置:在纯 PyTorch 中实现,未使用自定义 Triton 内核。
实验结果
通用语言建模能力评估。在 340M 和 1.3B 参数模型下,评估了语言建模困惑度和常识推理基准上的零样本准确率。结果显示,$\mathrm{E^2}$-TTT 在两个尺度上均获得了最低的困惑度,在 LAMBADA 上的优势最为明显(1.3B 下为 15.3,最强基线为 16.1)。在零样本准确率方面,$\mathrm{E^2}$-TTT$_{MLP}$ 在 1.3B 下达到 $54.5\%$ 的平均值,领先于最强基线的 $53.8\%$。
上下文检索能力评估。在 FDA、SWDE 和 SQuAD 任务上评估了模型的精确回忆能力。虽然全注意力模型(Transformer$^{++}$)由于明确的历史访问仍然保持优势,但 $\mathrm{E^2}$-TTT 显著缩小了这一差距。1.3B SwiGLU 模型平均达到 $43.6\%$,超越了 HQLT($35.5\%$)和 LaCT($36.7\%$)。SwiGLU 变体大幅优于 MLP 变体,验证了复杂门控机制有利于精确信息提取。
长度外推能力评估。
1. 语言建模外推:在 6 个长上下文数据集上可视化每个 Token 的损失指标。LaCT 在超出训练边界后损失立即爆炸,而 $\mathrm{E^2}$-TTT 变体在整个外推区域内表现出更低且稳定的损失曲线。
2. 长程关联回忆:在“大海捞针”(S-NIAH)基准测试中,测试序列扩展到 16K Token(8倍于训练范围)。在 S-NIAH-1 上,两个 $\mathrm{E^2}$-TTT 变体在 16K 处保持至少 $85\%$ 的准确率,而 LaCT 崩溃至接近零。
3. 真实世界长上下文理解:在 LongBench 的 14 个任务上,$\mathrm{E^2}$-TTT$_{SwiGLU}$ 获得最高平均分($14.1\%$),优于 HQLT($12.1\%$)和 Mamba2($10.3\%$),几乎是 LaCT($7.7\%$)的两倍。
控制变量消融实验。在相同的架构下分离更新规则。将 LaCT 的更新规则移植到本文框架(LaCT-Matched),并在本文模型中将逐 Token 因子折叠为块级标量(Chunk-averaged)。结果表明,折叠逐 Token 标量会导致 S-NIAH-1 在 8K 时的准确率从 $95.4\%$ 下降到 $6.8\%$。这归因于本文富有表达力的逐块更新规则在长上下文检索中的增益。
多模态视频理解评估。将 $\mathrm{E^2}$-TTT 集成到 Qwen3VL-2B-Instruct 中,作为与自注意力层融合的并行分支,仅在 LLaVA-Video-178K 的子集上训练 TTT 参数。结果表明,它在 VideoMMMU 和 LongVideoBench 上改进了冻结的基础模型,并匹配了在相同数据上训练的全微调基线的性能。
补充细节
相关工作概览。测试时训练(TTT)将 RNN 中的递归状态重新定义为在线适应非线性神经网络的可学习权重。现有方法通常采用自监督损失。虽然 TTT 开启了新颖递归模型架构的设计空间,但顺序瓶颈导致硬件利用率低。近期工作如 LaCT 试图通过从逐 Token 更新转向逐块更新来缓解这一问题,但往往诉诸于简化的聚合规则(如平坦平均),降低了优化保真度。本文的 $\mathrm{E^2}$-TTT 调和了分块的硬件效率与复杂逐 Token 更新动态的数学精度。此外,结合全局递归与局部注意力是高效长上下文建模的成熟策略,本文保留了这一架构模板,但用 $\mathrm{E^2}$-TTT 替换了递归模块。
附录细节
动量与权重的展开推导。通过归纳法,步骤 $t$ 的动量 $\mathbf{M}_t$ 是块起始衰减初始动量与梯度历史的总和:
类似地,展开权重更新得到:
累积核 $R_t$ 的稳定性分析。在接近单位保留机制($\beta_t \to 1, \gamma_t \to 1$)下,后缀乘积趋向于 1,$R_t \to C - t + 1 \leq C$。权重核受 $\eta_{\mathrm{base}} \cdot C$ 限制,保证了最坏情况下的边界。在激进衰减机制下,$\gamma_t$ 较小,比例呈几何衰减,$R_t$ 保持在 $\mathcal{O}(1)$。衰减参数化 $\gamma_t = 1 - \eta_t \alpha_t$ 强制执行了结构耦合:大更新自动增加遗忘,防止跨块的进位膨胀。
单次反向传播的链式法则分解。每个逐 Token 损失的形式为 $\mathcal{L}_t = \mathcal{L}(f_{\mathbf{W}_0}(\pmb{k}_t), \pmb{v}_t)$。应用链式法则:
雅可比矩阵仅依赖于 $(\mathbf{W}_0, \pmb{k}_t)$。因此,两个聚合简化为在相同的激活梯度集和输入上进行标量加权收缩。
块动态的两种等价视角。本文的公式承认两种互补的解读:优化风格视角(跨块边界自然解读,动量真正被携带)和核权重视角(块内聚合的自然解读,冻结的 $\mathbf{W}_0$ 约束使得“动量”和“衰减”与梯度贡献的学习逐 Token 加权相一致)。
组件消融分析。
- 块粒度:替换为平均因子的块级基线后,Wiki困惑度从 25.5 上升到 26.9。将块大小减小到 $C=1024$ 和 $C=512$ 逐步缩小了与理想逐 Token 更新的差距。
- 分支贡献:仅使用 SWA 分支在 8K 时仅达到 $6.8\%$,仅使用 TTT 分支则崩溃至 $0.0\%$,证明两者缺一不可。移除块后归一化或融合门也会损害性能。
- 内循环项:添加动量和权重衰减相对于无动量基线降低了困惑度。
- 基础超参数敏感性:性能对 $\alpha_{\mathrm{base}}$ 基本不变,对 $\eta_{\mathrm{base}} \in \{10^{-3}, 10^{-2}\}$ 稳定,仅激进的 $10^{-1}$ 会破坏训练。
学习系数的因果测试。测量显示 $\eta_t, \beta_t, \gamma_t$ 的变异系数很高。沿 Token 轴打乱系数(保留边缘分布但破坏与输入的对齐)表明:仅打乱 $\eta$ 会使 S-NIAH-1 准确率从 $93.6\%$ 降至 $6.0\%$,确认步长承载了内循环的大部分写入选择性。联合干预的成本基本等于个体干预的总和,证明这三个头提供的是独立贡献。
吞吐量与推理成本分析。在 H100 GPU 上,Titans 模型严重受限于吞吐量(在 1.3B 时 OOM)。相比之下,$\mathrm{E^2}$-TTT 变体保持接近高度优化的 LaCT 基线的吞吐量。精确内核在 340M 时的吞吐量成本为 $15.6\%$,在 1.3B 时仅为 $3.3\%$。在推理成本方面,所有递归模型在 2K 到 16K 之间的延迟和内存都是平坦的,而注意力模型成本随长度增长。
更大预算下的训练验证。在 4K 训练上下文下,以 100B Token(预算的 6.7 倍)训练 1.3B $\mathrm{E^2}$-TTT$_{SwiGLU}$。困惑度和常识准确率大幅提高,上下文检索从 $43.6\%$ 上升到 $58.1\%$,证明了该方法在更大预算下稳定训练的可行性。
结论
本文提出了 $\mathrm{E^2}$-TTT,该方法调和了逐块处理的硬件效率与测试时训练中逐 Token 更新的表达能力。通过推导闭式标量核,实现了逐 Token 递归动态的并行执行,而无需先前逐块方法的块级近似。经验表明,$\mathrm{E^2}$-TTT 在语言建模方面与强大的次二次基线相当,在检索方面优于它们,同时展示了卓越的长度外推能力,在 8 倍训练上下文长度的“大海捞针”测试中维持了 $>90\%$ 的准确率。这些结果表明,保留精确的时间动态对于长上下文泛化至关重要,使 $\mathrm{E^2}$-TTT 成为可扩展测试时训练的一个有前景的方向。未来的工作将探索高效的逐 Token 输出步骤,并将其扩展到更大的模型规模。
💬 评论讨论
欢迎在这里分享您的想法和见解!