论文

通用测试时训练

Universal Test-Time Training

模型架构递归架构

摘要

最近的测试时训练 (TTT) 架构将上下文压缩为快速权重,这些权重在线更新并作为记忆进行查询。现有的 TTT 设计使记忆对每一层都是私有的:它仅随着时间的推移而重复,并且深度仅索引 L 个单独的记忆。我们认为记忆所有权不必与深度绑定,并引入通用测试时训练 (uTTT),其中所有层都读写一个共享的记忆,同时保留特定于层的骨干模型参数。因此,共享的记忆在时间和深度两个维度上重复出现,以块和层为单位:一个块中深层的写入可以由下一个块中的浅层读取。我们将这个想法实例化为 uTTT-MoE 和 uTTT-Dense。 uTTT-MoE 将每个token头路由到所有层共享的池中的几个专家; uTTT-Dense 在每一层应用整个共享的记忆,无需路由。在语言建模中,uTTT-MoE 在 124M 和 760M 下达到 15.5 和 27.9 RULER 精度,比同等状态和主动计算下的层私有对应精度高出 2.6 和 2.1 个点,是经过测试的有界状态模型中最高的,每个token损失匹配或超过完全注意力。在新颖的视图合成中,以固定的每层计算共享在路由模型中的视图 23 对象 PSNR 中获得 0.92 dB,在密集模型中获得 0.76 dB。

通用测试时训练的原论文方法或结果图
图 4:状态读取的向后和向前依赖性。 LLM诊断使用 124M 模型和各个 Books3 文档的第一个 32K token。 (a) Squared-梯度共享来自第 1 层、块 8 读取输出雅可比矩阵探针,按文档标准化,平均超过 300 个文档。这探针是外部任务丢失反向传播的局部因素。颜色是对数的;灰色为零。装箱计数要求每个文档中的梯度非零;圆圈标记已读内容。 (b) 在 1,024 个文档中,相对于每个模型对块 2-8 的正常预测,条形图显示冻结状态写入或仅读取自己的写入后损失增加;点显示位置 $\geq 30{,}000$ 处未修改的尾部 PTL,是在这些未打包的文档上测量的,因此不在表的范围内。 Full 使用完整的 uTTT-MoE 反向传播; Cut在训练期间停止跨层读写梯度; TTT 是 TTT-MoE。 (c) 在 uTTT-MoE-e64-p1 中删除一条写入器-读取器路径后,PSNR 下降,平均超过 1,019 个 GSO 对象,每个对象有 23 个目标视图。包括对角线路线。颜色是线性的,下降到低于 $0.01$ dB 白色;轮廓标记了最大值。面板(b,c)使用固定训练的检查点。