学习长度可外推的循环模型
Learning Length-Extrapolatable Recurrent Models
论文信息
标题: Learning Length-Extrapolatable Recurrent Models
作者: Hanwen Jiang
发布日期: 2026-09-08
arXiv ID: 2609.09157v1
PDF 链接: 下载 PDF
3 分钟速览
- 研究问题:循环模型天然适合长序列,但用短序列做 BPTT 训练后,往往在训练长度之外失效;论文研究的是 “未来损失如何把信用传回早期循环状态” 这一问题。
- 核心方法:提出 Credit Stabilization through Time(CST),在反向传播的块边界对状态信用做局部正标量重缩放,不旋转该分量方向,也不改变前向计算。
- 关键结果:在 Meta-FSA 上训练长度 1k、评估到 128k 时,Event-CST 在 16k–128k 的全部 16 个深度–长度单元中平均比 BPTT 高 6.58 个百分点(图 1)。
- 主要局限:CST 只调整信号尺度,不能恢复缺失方向或消除时间贡献间干扰;真实语言模型证据来自单一 123M 模型和单一训练运行,GovReport 上的远距离上下文使用控制未统计解析。
- 适合读者:关注循环网络、状态空间模型、线性注意力、BPTT 训练与长上下文外推的研究者和工程师。
论文背景和研究动机
长上下文建模是当前序列模型的核心能力之一,但具有极端长依赖的训练数据往往稀缺。论文采用 “短训练、长部署” 的设定:从较短的可用序列中学习,再在更长的推理序列上运行。Transformers 在这种设定下面临键值缓存增长和位置编码外推问题;循环模型则保持固定大小状态、天然适合长部署。然而,架构上的适合并不能自动保证学习到的循环规则能在训练长度外继续可靠工作。
论文首先质疑传统观点:梯度沿时间路径消失或爆炸并不足以解释学习失败。实验中,Fixed-FSA 的状态信用从后期损失向早期传播时几乎完全收缩,但由于每个位置都有稠密监督,共享转移规则仍能训练。因此,作者把诊断从 “参数梯度” 移到上游的 “状态信用”:
它表示未来损失如何评价一个更早的循环状态。论文的核心动机是:未来损失先经过 BPTT 形成状态信用,再经过局部参数 Jacobian 映射并聚合为参数梯度;传统诊断主要关注后者,而论文关注前者。
核心方法和技术细节
论文先用受控任务诊断何时需要远距离信用。把序列按边界分为若干块,令完整 BPTT 的循环参数梯度为 ,在每个边界 detach 后得到仅保留本地路径的 ,定义边界穿越分量 。用投影量 衡量该分量与完整梯度的对齐程度。在序列长度 1024、边界间距 128 下,Fixed-FSA、Meta-FSA、MQAR 的 分别为 0.021、0.089、0.259(表 2)。同时,用固定后期损失窗口测量状态信用范数传播比 ,三者分别为约 0、1.05、0.82。也就是说,参数梯度层面的远距离依赖与状态信用层面的传播强度并不等价。
MQAR 的消融进一步显示:只用 训练,准确率停留在 0.6%;完整 BPTT 达到 96.3%。这表明该任务中局部监督无法替代跨边界信用。作者据此得出的不是 “衰减必然失败”,而是 “远距离学习需要两个条件”:任务确实需要跨边界信用,并且该信用以可用尺度和方向到达早期状态。
CST 在块反向传播上实现。块 将状态从 映射到 ,包含损失 。普通 BPTT 的边界状态信用满足:
其中 是当前块损失对边界状态的贡献, 是块级状态转移 Jacobian。CST 在内部边界以正标量 对进入前一块的信用做局部重缩放:
由于 ,它只改变范数,不旋转方向。截断 BPTT 等价于 ,而 CST 保持图连接、仅重缩放信号。
受控合成任务中信用收缩占主导,因此采用单边 Event-CST:维护参考范数 ,当 时触发修正,增益为 。EMA replay 则每 步做一次 dense 探测,用 EMA 维护边界收缩轮廓,普通步骤只重放预测的事件位置,但增益仍从当前 minibatch 重算。真实文本中信用变化两边都有,因此采用按层、按头的对称控制器:将头内状态信用范数取对数得到 ,与参考值 的偏差 决定增益:
论文在语言模型实验中使用 、、、。
创新点和贡献
这项工作的主要贡献是把长度外推失败问题从参数梯度层面移动到状态信用层面。论文没有提出新的前向架构,而是干预反向传播中更上游的信号,并保持前向计算不变。这个视角使方法可以应用于多种具有循环形式的模型。
第二个贡献是任务特异化。作者没有给所有数据用一个控制器,而是先通过探针区分 “收缩主导” 与 “双向不规则” 的信用动态,再分别设计 Event-CST 和对称 CST。EMA replay 则进一步通过稀疏化边界测量来恢复吞吐,但仍保留当前批次的信号计算,避免直接复用陈旧梯度。
第三个贡献是证据结构。受控任务分离了状态跟踪、上下文检索和混合任务,明确展示 “何时需要远距离信用” 和 “信用是否传到” 之间的区别。真实语言模型部分不仅报告困惑度变化,还使用 LongPPL 风格的 key-token 标注和 context-removal control,试图检验模型是否真的使用了超过训练长度的远处上下文。
实验结果分析
在 Meta-FSA 混合任务上,Event-CST 在训练设置长度 1k、 时比 BPTT 低 3.55 个百分点,但在 16k–128k 的全部 16 个深度–长度单元中平均提升 6.58 个百分点(图 1)。最长的 、128k 设置下,BPTT 为 21.16%,Event-CST 为 28.14%,提升 6.98 个百分点;未见过的深度 也有 4.77、4.09、3.87 个百分点的提升。论文同时指出,48 个种子–单元配对中有 29 个为正,说明收益方向一致但幅度因训练运行而异。
吞吐方面,BPTT 处理 0.955M tokens/s;eager dense Event-CST 降至 0.119M tokens/s,CUDA Graph 可恢复到 0.526M tokens/s,是 BPTT 的 55.1%;EMA replay 投影约 0.808M tokens/s,达到 BPTT 的 84.6%,但平均长上下文中准确率为 23.2%,比 dense Event-CST 的 25.1% 低 1.91 个百分点(图 2)。
泛化实验结果与任务结构一致:Fixed-FSA 上收益基本不随长度增长;MQAR 上随保留时间延长的收益更明显,2k–128k 平均提升 5.66 个百分点,128k 下从 25.56% 提升到 41.66%。在更大的 8 层模型中,较难任务 S64/A8/L2k 上 Event-CST 在所有深度都有正收益,平均 7.14 个百分点;但在标准任务 S32/A4/L1k 上,训练深度之外的迁移为负(图 4)。论文没有把这一反差归因于单一因素,因为两组比较同时改变了多个条件。
真实语言模型实验中,CST 在全部 14 个数据集–长度评估上都降低了 all-token NLL,平均降低 0.00724。最长评估长度上,LongData 32K、Books3 128K、GovReport 32K 的 all-token 增益分别为 +0.006843、+0.005949、+0.009726(表 3)。关键 token 上的增益更大:在最长上下文中分别是 all-token 增益的 3.50、2.48、4.54 倍。LongData 的 context-removal control 显示,CST 相对 BPTT 更好地利用了 4K 之外的上下文,增益增加 0.01470,并在 16K、32K 上统计解析;GovReport 的对应值只有 0.00268,未统计解析(表 5、表 6)。
实践建议
如果要在循环或状态空间模型上做 “短训练、长部署”,可先按论文诊断是否存在远距离信用需求:用类似 或远距离参数梯度分离,确认任务不能只靠局部监督学会。不要只看梯度范数是否衰减;论文在 Fixed-FSA 上观察到严重收缩但局部监督仍可训练。
对于结构较规整、收缩事件可识别的任务,可以尝试 Event-CST 或 Adjacent-CST。论文在 Meta-FSA 上的默认设置为 、、,但最好在目标外推长度上做小规模搜索;需要注意训练长度内性能可能略有下降。若关注训练吞吐,CUDA Graph 可显著恢复 dense Event-CST 速度,EMA replay 可进一步降低边界测量开销,但会牺牲部分长上下文中收益。
在真实文本或依赖结构不规律的数据上,建议使用温和的对称头级控制器,而不是单边强放大。论文采用 、、,并报告更激进的 平均表现更差。该配置是在论文报告的评估套件上选择出来的,直接迁移到其他模型、词表或数据域前需要重新评估。
最后,评估时不应只看全序列 NLL。可结合类似 LongPPL 的 key-token 标注和 context-removal control,区分 “预测得更好” 和 “真的使用了更远上下文”。CST 的修正信号也不是原始目标的精确梯度,无法恢复方向缺失,因此更适合作为训练诊断后的干预,而非无条件替换 BPTT。