Layer-Feedback Transformer:用更少参数跑更多层计算的新架构

0 阅读

传统Transformer的单向信息流局限

大多数Transformer语言模型采用严格的前馈结构:输入依次经过L1、L2……直到最后一层,每层只处理一次。这意味着一旦信息离开某一层,该层就再也无法接触后续层生成的更深层表示。这种设计简单高效,但也可能限制了模型对信息的精细加工能力。

为突破这一限制,研究人员提出了Layer-Feedback Transformer(LFT)架构。它的核心思想很直接:不让信息单向流动,而是让相邻的Transformer层反复“对话”。比如在一个5层模型中,信息流不再是简单的L1→L2→L3→L4→L5,而是变成:

L1 → L2 → L1 → L2 → L3 → L2 → L3 → L4 → L3 → L4 → L5

这样,一个5层LFT模型虽然只有5个独特的层参数,却执行了11次层运算。关键在于,没有新增任何参数,只是改变了计算路径。

LFT的执行机制

LFT的实现并不复杂。以相邻两层Li和Li+1为例,标准流程是:

h1 = Li(h)
h2 = Li+1(h1)

而LFT在此基础上增加两次反馈:

h1 = Li(h)
h2 = Li+1(h1)
h3 = Li(h2)  // 让Li再处理一次
h4 = Li+1(h3) // Li+1再次处理

这样,较早的层Li有机会处理已经被更深一层Li+1加工过的表示,从而获得更丰富的上下文信息。

对于N层模型,完整的执行序列可概括为:

x = self.layers[0](x)  # 先过第一层
for i in range(1, len(self.layers)):
    x = self.layers[i](x)  # 进入新层
    if i < len(self.layers) - 1:  # 如果不是最后一层
        x = self.layers[i - 1](x)  # 回退到前一层
        x = self.layers[i](x)      # 再次进入当前层

值得注意的是,Transformer层本身无需任何修改。实验中使用的仍是标准组件:pre-normalization、RMSNorm、因果自注意力、RoPE位置编码、SwiGLU激活函数、残差连接,以及共享的输入输出嵌入。

参数深度 vs 执行深度

传统Transformer中,“参数深度”(独特层数)和“执行深度”(实际层调用次数)是相等的。例如5层模型就是5次执行。

LFT打破了这种耦合。同样是5层:

  • 参数深度:5(存储5个独特层)
  • 执行深度:11(实际调用11次)

这种设计让模型在不增加参数的情况下,获得了更多的表示变换机会。对于N层LFT,执行次数E(N) = 3N − 4。这意味着:

  • 6层模型:14次执行(2.33倍计算)
  • 11层模型:29次执行(2.64倍计算)
  • 12层模型:32次执行(2.67倍计算)

随着层数增加,计算开销趋近于标准模型的3倍。

控制变量实验设置

为了公平比较,研究人员在三个参数规模(2.5M、10M、25M)下训练了配对模型。每对Standard和LFT模型共享以下所有配置:

  • 参数数量
  • 独特层数
  • 隐藏层维度
  • 注意力头数
  • FFN大小
  • 初始化种子
  • 分词器(3072-token byte-level BPE)
  • 上下文长度(768 tokens)
  • 训练数据(FineWeb-Edu)
  • 优化器与学习率调度
  • 总训练token数(5亿)

唯一区别就是层的执行路径。这种设计确保了任何性能差异只能归因于LFT的反馈机制,而非其他因素。

基准测试结果

小模型(2.5M参数):标准模型胜出

在最小规模下,LFT反而表现更差:

  • Base Bench 1.1整体准确率:Standard 32.57% vs LFT 30.86%
  • PIQA、ARC-Easy、HellaSwag三项LM-Eval指标也略低

这说明在参数极度受限时,额外的计算开销并未带来收益,可能因为模型容量不足以支撑复杂的反馈交互。

中等模型(10M参数):LFT显著领先

当参数增至10M,LFT开始展现优势:

  • Base Bench准确率从32.57%提升至36.86%(+4.29个百分点)
  • 正确回答题数:114 → 129(共350题)

提升最明显的几个类别:

  • 上下文跟踪:22.29% → 33.44%(+11.15%)
  • 定量任务:19.87% → 28.08%(+8.21%)
  • 逻辑推理:28.84% → 36.68%(+7.84%)

这些任务通常需要模型在长上下文中保持连贯性或进行多步推理,LFT的反馈机制可能帮助模型更好地整合信息。

大模型(25M参数):LFT小幅领先

在25M参数规模,LFT继续保持优势,但差距缩小:

  • Base Bench准确率:35.43% → 36.57%(+1.14%)
  • LM-Eval三项平均:38.64% → 38.84%(+0.20%)

有趣的是,代码补全任务在25M规模下反超:LFT达到26.43%,而Standard仅21.66%。这可能因为代码具有严格的结构依赖,反馈机制有助于捕捉跨行的语义关联。

训练量的关键作用

LFT的效果高度依赖训练充分性。研究人员对比了同一10M模型在不同训练量下的表现:

  • 2亿token训练:Standard 32.00% vs LFT 29.43%(LFT落后2.57%)
  • 5亿token训练:Standard 32.57% vs LFT 36.86%(LFT领先4.29%)

这说明LFT需要更多训练才能“学会”如何有效利用反馈路径。在训练不足时,额外的计算反而成为负担;只有当模型充分收敛后,反馈机制的优势才显现出来。

计算成本与公平性讨论

必须强调:LFT的性能提升部分源于更高的计算开销。实验中虽然参数量和训练token数相同,但FLOPs(浮点运算次数)并不相等。例如10M模型的LFT版本计算量是标准版的2.64倍。

因此,LFT的收益可能来自两方面:

  1. 架构优势:反馈机制确实提升了表示质量
  2. 计算红利:更多层调用相当于隐式增加了模型深度

要区分这两者,未来需要进行计算量匹配的实验(即让标准模型也运行更多步)。不过即便如此,LFT提供了一种在固定参数预算下换取更高计算效率的思路——对于部署场景中参数量受限但计算资源充裕的情况,这可能很有价值。

潜在应用场景

LFT的设计特别适合以下场景:

  • 边缘设备部署:参数量直接影响内存占用,而LFT能在不增加参数的前提下提升性能
  • 长上下文任务:如文档摘要、代码生成,需要模型在长距离依赖中保持一致性
  • 推理密集型应用:如数学证明、逻辑谜题,多轮反馈可能帮助模型逐步修正中间表示

当然,如果计算资源极其紧张(如手机端实时推理),标准Transformer可能仍是更优选择,因为LFT的多次层调用会显著增加延迟。

总结

Layer-Feedback Transformer通过重构计算路径,在不增加参数的情况下实现了更深层次的信息交互。实验证明,这种设计在中等以上模型规模、充分训练条件下能带来实质性提升,尤其在需要复杂推理的任务上。但它并非万能药——小模型或训练不足时反而有害。这提醒我们:架构创新的价值往往取决于与训练规模、任务特性的匹配度,而非绝对优劣。