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

Full-bandwidth Transformer

3,750 字约 14 分钟

原文: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 压缩为了单个符号,每步传回开头的信息最多只有log2|V|log_2|V|。模型中间的 latent 信息存储在 KV Cache 中并没有丢失,但处于一种 “深度冻结” 的状态:第ll层产生的状态只能被ll之后的层度去,永远不会被从头开始重新处理。模型在这种预设下面临一种两难:要么花 Token 把中间状态显示地 “念出来” 从而传回开头(CoT),要么在每个位置从头开始计算此状态。 作者提出通过把模型的 output latent 通过 GLU 融合之后反馈到开头,Token 仍然照常采样和输出(所以监督信号不变)。主要的收益来自于:

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

2. Background

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

Rstd={(t′,l′):t′<t,l′<l}R_\text{std}=\{(t', l'): t'<t,l'<l\}

规模𝒪(lt)\mathcal{O}(lt)。对于浅层的新 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

主要做法是

𝐡tL=fθ(et⊗𝐡t−1L;C),et⊗ht−1=WUht−1⊙σ(WGet)\textbf h^L_t=f_\theta(e_t\otimes \textbf h^L_{t-1}; C), \quad e_t\otimes h_{t-1}=W^Uh_{t-1}\odot \sigma(W^Ge_t)

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

3.2. Latent feedback decoding vs. standard CoT

作者把 CoT 形式化为 “状态 = token 序列” 的 MDP,跨步传递的只有离散动作。理论上问题求解状态是历史动作的确定函数,但从 token 历史恢复它本身就是状态跟踪问题,而固定深度 transformer 每次前向只有有界的串行计算,CoT 需要靠把中间状态外化为语言解决这个问题。相比之下,本文提出的 latent feedback decoding 的状态则是 st=(a1:t,zts_t = (a_{1:t}, z_t,token 轨迹加最近 latent,历史 latent 则已经在 KV Cache 中。 作者还指出,由于zt+1z_{t+1}是x1:t+1x_{1:t+1}的确定函数,它不含上下文没有提到的消息(“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 陡升,‖hk−hk−1‖‖h^{k}−h^{k−1}‖ 震荡。混入 3% 三遍批次 (75/22/3),验证损失在 30 步反馈内保持平坦,状态变化量衰减到小平台,即学到的反馈映射变成朝固定点的压缩映射。 训练的时候每个 token 都能拿到 latent,但是推理的时候对 prompt 的 prefill 是不处理这些 latent 的(为了效率),形成了错配。本文提出 Prefix Mixing,每遍随机采样前缀长度pp,把t≤pt\le p的位置还原为纯 embedding,只 fuse 后缀的部分,即混入一部分随机的不融合的部分模拟 prefill。也支持了在 prefill 阶段跑两次实现每个 token 都有 latent mixing 的做法。 为了稳定训练还是用了:

  1. Depth Scaling11 , 在每个 residual 上乘一个随深度衰减的系数,防止 latent 不停灌回开头导致不断放大然后爆炸;
  2. embedding 和 readout 权重绑定;
  3. 对 latent 加一个 Jitter 噪声(𝒩(−σ,σ)D,σ=0.02\mathcal N(-\sigma, \sigma)^D, \sigma=0.02),让整个流程更加 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 计算得到了2×2\times的数据效率;

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。22被运送的摘要本身就是残缺的,那单纯地增加带宽对模型的性能也不会有什么帮助。而相应地,如果从一开始就是完全串行、全部递归传输信息,状态是增量维护的,每个位置的状态有点像 RNN 的 Hidden State 一样被维护,对这个信息就更好。 这个结果某种意义上和前面的 Task 相互印证,数学题的状态应该是在生成过程中积累的更多,增大带宽降低了模型重复思考相似问题的难度(此时 soft 推理就足够),而代码题的难度需要从 prompt 里面反复获取遵守(此时 fused 更好)。

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?)上的实验根本没法进行。

脚注

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

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