跳到主要内容
起居室老虎
返回

Full-bandwidth Transformer

更新于:

原文arXiv:2608.08888 · Wang et al., 2026

摘要: 自回归 transformer 沿两条轴计算:水平方向跨越已生成的 token,垂直方向穿过模型深度。稠密注意力赋予每个 token 对历史的宽阔水平访问,但解码步之间的垂直反馈通道仍然狭窄:只有被采样的 token 返回栈底,顶层隐状态则被丢弃。我们提出全带宽 transformer,用latent feedback拓宽这一通道:每个解码步中,将上一步的顶层隐状态与采样 token 的嵌入经一个GLU融合后,作为下一步输入送回。潜反馈让未被言语化的计算带着“焕新的深度预算”重新进入网络栈,同时完整保留标准 transformer 架构、KV cache 与语言建模目标。为了在不牺牲并行 teacher forcing 的前提下训练它,我们采用按计划调度的multi-pass目标:在预训练后期才引入潜反馈,并混入少量更深的反馈遍数以稳定训练。我们训练了 1B 参数、至多 400B token 的全带宽 transformer,发现潜反馈改进了验证损失、5-shot LM 评测、数学与代码生成、以及指令微调后的性能。在每 token 解码开销可忽略的前提下,全带宽 transformer 达到或接近用约 1.5× 更多 token 训练的标准 transformer,并能在同等或更好准确率下产出更短的推理轨迹。

1. Intro

由于高质量数据获取越来越困难,开始讨论“能否通过给每个 token 分配更多计算来榨取更多学习信号”。作者认为的一个关键前提:额外的FLOPs只有在转换为训练期更丰富的表示,或是推理期更有效的计算时才有意义。 Tokenizer在Auto Regression的过程中,将模型输出的D维latent压缩为了单个符号,每步传回开头的信息最多只有。模型中间的latent信息存储在KV Cache中并没有丢失,但处于一种“深度冻结”的状态:第层产生的状态只能被之后的层度去,永远不会被从头开始重新处理。模型在这种预设下面临一种两难:要么花Token把中间状态显示地“念出来”从而传回开头(CoT),要么在每个位置从头开始计算此状态。 作者提出通过把模型的output latent通过GLU融合之后反馈到开头,Token仍然照常采样和输出(所以监督信号不变)。主要的收益来自于:

  1. latent中含有的没有显式变成文字的reasoning的状态,能够通过这种方式被模型继续加工处理;
  2. 每一层,即使是最浅层,也能够看到前后所有层的所有状态。

2. Background

标准Transformer在位置,层的可达状态集合是

规模。对于浅层的新Token,只能看到被Tokenizer压缩过的信息,也有可能更深层的latent中已经蕴涵了本次需要的信息。这种深度方向的依赖约束正是 transformer 能跨位置并行训练的原因。直接使用Latent作为输入会导致模型训练时退化回RNN的BPTT,直接无法大规模并行训练了。 然而,这个约束只在训练时提供了便利,在推理的时候模型本来就是逐个Token生成的(AR不会并行推理生成Token),如果在推理的时候能够以某种方式让模型以latent而非token作为输入,对推理本身的开销来讲并不算很大。训练并行性和解码带宽的约束来源相同,但只在训练期是必要的。

3. Widening the bandiwidth with latent feedback decoding

3.1. Latent feedback decoding

主要做法是

其中是前序的KV Cache,是第层的第个latent,是纯token。作者认为,加性融合存在捷径,模型有可能把学小而回到纯token输入,变回普通预训练,尤其是如果不在预训练一开始就引入这个latent通道的情况下。而本文使用的GLU融合则没有这个问题。 额外推理开销与上下文长度和模型深度无关,每 token 低于 1%。 本来就要算(用于采样),额外工作只有两次 乘;维度保持 不变,所以架构、KV cache 布局全部不动,并且借用 MTP 实现的 buffer 机制做到了 vLLM 兼容。

3.2. Latent feedback decoding vs. standard CoT

作者把 CoT 形式化为“状态 = token 序列”的 MDP,跨步传递的只有离散动作。理论上问题求解状态是历史动作的确定函数,但从 token 历史恢复它本身就是状态跟踪问题,而固定深度 transformer 每次前向只有有界的串行计算,CoT 需要靠把中间状态外化为语言解决这个问题。相比之下,本文提出的latent feedback decoding的状态则是 ,token 轨迹加最近latent,历史latent则已经在KV Cache中。 作者还指出,由于的确定函数,它不含上下文没有提到的消息(“the gain is computational, not informational“),并且过去的状态不会像RNN一样被覆写而是保留在KV Cache中、串行的深度也不发生改变,作者认为这与Loop Transformer等增加等效深度的方法有本质区别。

3.3. Parallel training for latent feedback decoding

前面也提到过,直接训练latent feedback的做法会导致模型无法像Transformer一样并行地teacher forcing,本文提出的方法是通过多次forward来处理。首先进行一次普通的forward,然后把上一遍的output全部记录,“整体移动一位”,与token融合,然后再全位置并行地重跑,作者称之为temporal parallelism。这个过程可以做k次,得到的是跨k个Token的信息传递,训练的开销是标准训练的k倍。loss在每遍上都是标准的NTP,梯度也不detach,后面的梯度回反穿进入前面的latent。第一遍也要做NTP的训练是因为prefill的时候是没有latent的,需要保留模型在没有latent的时候也能工作的能力。 用 75% 单遍 + 25% 双遍训练的模型在训练深度内表现良好,但外推立刻崩溃,val loss陡升, 震荡。混入 3% 三遍批次(75/22/3),验证损失在 30 步反馈内保持平坦,状态变化量衰减到小平台,即学到的反馈映射变成朝固定点的压缩映射。 训练的时候每个token都能拿到latent,但是推理的时候对prompt的prefill是不处理这些latent的(为了效率),形成了错配。本文提出Prefix Mixing,每遍随机采样前缀长度,把的位置还原为纯embedding,只fuse后缀的部分,即混入一部分随机的不融合的部分模拟prefill。也支持了在prefill阶段跑两次实现每个token都有latent mixing的做法。 为了稳定训练还是用了:

  1. Depth Scaling1 , 在每个residual上乘一个随深度衰减的系数,防止latent不停灌回开头导致不断放大然后爆炸;
  2. embedding和readout权重绑定;
  3. 对latent加一个Jitter噪声(),让整个流程更加robust。

3.4. Latent-feedback training improves pre-training data efficiency

额外发现,这样训练的模型,即使在推理的时候完全不引入latent,这个训练过程本身也增强了模型的能力。

4. Experiments

setup:

4.1. Fused prefilling improves non-generative performance

主要有三个发现:

  1. 主要的增益出现在第一次prefill融合的过程,后续收益递减,符合“增加了prompt的有效深度”的解释;
  2. 完全不做feedback也几乎不掉进度;
  3. k=2的情况下,100B的训练能力接近200B的Baseline,200B的接近400B的baseline,即用过少量的prefill计算得到了的数据效率;

4.2. Latent feedback decoding improves decoding performance

对于一个训练好的模型有三种推理模式:Standard,Soft(只在decode阶段有latent feedback),Fused(prefill额外跑一遍融合latent + decode阶段latent feedback)。 Soft在所有任务、所有规模上都比Standard好。任务偏好上,数学问题soft更好,代码问题fused更好,作者认为是因为代码奖励更深的prompt表示。Pass@3和Pass@1都有提升,证明不是靠采样多样性塌缩产生的能力。

4.3. Latent feedback enables more concise reasoning

在Math500上,Median Reasoning Length Standard > Soft > Fused。注意到此效应在做了Instruction Finetuning之后消失。 作者认为微调数据相对潜反馈解码是 off-policy 的,目标轨迹由标准逐 token 推理产生并模仿其冗长风格,拟合它们就重新强加了全言语化风格。on-policy 后训练可能保住简洁性,留作未来工作。

4.4. Full-bandwidth transformer carries richer information in shallow-layer residuals

前面提到了一个可达集的问题,这里设计了一个额外的实验讨论此问题。 具体来说,实验构造了一批合成输入,每条输入包含一个二元的真值标签(由输入内容决定),所有的输入以同一个token”:“结尾。在这个结尾token的位置,这个benchmark记录每个深度的residual latent,depth=0的定义为送进第一层之前的输入向量。 然后在每个深度上,独立训练一个L2 Regularized Linear分类器,从这个residual latent预测标签,绘制Depth vs. Accuracy的深度图。acc越高,意味着此处的latent“被一个Linear解码为信息”的能力更强。 合成数据有三种不同的Task,Task 1 Completion Tracking:

reiquired = 4
completed = 9
scratch = 7
scratch += 0
...
Status :

requried和completed需要全都相等则statue为done,否则为more。模型需要读两个位置、做比较,然后算出关系状态。 Task 2 Delayed Memory也类似,先赋值一个初始状态,然后添加一串跟标签无关的scrach异或/加法,末尾 # final state: 的标签对应的就是最开始的初始状态。测试的是模型跨过干扰段阅读信息的能力。 Task 3 Multi-Register Latest Write比较复杂:

r4 = 0
r4 = 1
r0 = 1
r7 = 0
  ...(10 条赋值)
r7 = 1
r1 = 0
scratch = 7
  ...(7 条空转)
query = r0
Value:

模型在读到query之前不知道要问的是哪个,在序列处理过程中需要维护所有寄存器的最新状态,然后把qeury绑定到正确的分量上,把任务升级为“维护一个m dim的状态向量 + 覆写 + 查询”。 对比对象包括

  1. 标准prefill
  2. 一步递归prefill
  3. k步递归prefill
  4. 完全串行的全部递归prefill 结果上,对于Task 1,2只要一步递归就能把Layer0的acc拉到接近100%,对所有干扰长度全部成立,而标准prefill需要经过比较多的layer才能解决此问题。这个结果想反映的问题是,即使是标准的prefill,模型其实已经把结果算出来了,但是这个信息很有可能被冻结在较深的层,新产生的token看不到这些信息,又要经过前几层重新做一轮计算,而一次递归fuse相当于把深层的信息解冻带入到了前面的层来。 对于Task 3,增加递归步数有帮助但没有Task 1, 2那么明显,在深层/更复杂的write操作下衰减显著,完全串行的全部递归prefill效果最好。一步递归运送的是前一位置的顶层状态,而状态本身是在标准 prefill 下算出来的,它对“8 个寄存器的覆写历史”的聚合质量受限于单次前向 L 层深度内能完成的状态跟踪。覆写越多,“解析最新值”越接近困难的顺序状态跟踪 regime。2被运送的摘要本身就是残缺的,那单纯地增加带宽对模型的性能也不会有什么帮助。而相应地,如果从一开始就是完全串行、全部递归传输信息,状态是增量维护的,每个位置的状态有点像RNN的Hidden State一样被维护,对这个信息就更好。 这个结果某种意义上和前面的Task相互印证,数学题的状态应该是在生成过程中积累的更多,增大带宽降低了模型重复思考相似问题的难度(此时soft推理就足够),而代码题的难度需要从prompt里面反复获取遵守(此时fused更好)。

5. Related work

Feedback Transformer,各层表示混合之后给未来的attention使用,但训练沿着token顺序难以扩展;同期的 T²MLR(晚中层注入早中层)和 Latent Recurrent Transformer(固定源层状态经额外 K/V 投影注入)在精神上相近。相比之下,本文的注入点“外置”于模型,不需要大幅度修改模型(只改了输入的方法),额外参数最少。 对于Latent Reasoning(Coconut、Soft thinking),本文主要聚焦于pretraining,用latent增强而不是替代token,监督训练更容易,代价是token效率的提升应该没有直接用latent的高。 对于Looped Transformer前面也提到了,本文实际没有增加模型的等效深度。

文章里面比较不可思议的地方在于,训练只覆盖了k非常小的情况,而decode动辄几百步,模型居然还能正常推理,我认为起作用的大概率是Token输入 + Latent是GLU混合这个逻辑,在Latent完全不可用的时候,保留下来的Token输入还让模型能够正常运行。这是一个比较有趣的现象。这个训练的逻辑有点像一个fwd和bwd都做Sliding Window的BPTT,是否有可能将相似的逻辑应用到SNN训练上?如果timestep之间的监督太复杂,是不是可以用在Video Task上做Diff Encode Aware的训练?能够通过这个逻辑让模型学到1-2frame之间的差别代替整个clip上的完整训练,同时又让模型学到多个frame之间的状态转移,就像前面提到的只训练小的就泛化到大的上。 单纯讨论这个工作本身还有比较遗憾的就是这个方法和现在常用的MTP方法应该都是完全不兼容的,训练的时候不detach梯度也导致在更大的模型(比如7B + K=4?)上的实验根本没法进行。

Footnotes

  1. Yang, Yu, Zhu, Hayou, “Tensor Programs VI: Feature Learning in Infinite Depth Neural Networks”, ICLR 2024

  2. 这个问题在固定深度Transformer的能力上限相关的论文亦有相关讨论。


分享这篇文章:

下一篇
SeeDNorm: Self-Rescaled Dynamic Normalization