← 返回 2026-08-14

全带宽Transformer:潜在反馈解码 Full-bandwidth transformer

Xi Wang, Ziyang Cai, Zheng Zhan, Harry Dong, Ying Fan, Gustavo de Rosa, Tim Pearce, John Langford 📅 2026-08-09 👍 21 2026-08-19 18:30
LLM预训练 数据效率 测试时计算 潜在反馈 递归计算

将顶层隐状态门控融合为下一步输入,以约1%推理开销换取约2倍训练数据量的效果

前置知识

KV缓存与自回归解码

decoder-only模型逐token生成时,把之前所有位置的Key/Value缓存下来复用,避免对前缀重复计算。与RNN把历史压缩进固定大小状态不同,稠密attention保留全部过去token的显式表示,新token的每一层都能直接attend到缓存里的全部历史。

本文的核心分析就是把信息通路分成'水平轴'(跨位置,靠KV缓存,已是全带宽)和'垂直轴'(跨深度,被限制),并指出KV缓存让旧计算'深度冻结'——这是理解潜在反馈动机的前提。

残差流与语言模型头

transformer每层的输出都累加在一个D维残差流上,顶层隐状态 $h^L$ 经过线性投影 $W_{head}\in\mathbb{R}^{|V|\times D}$ 得到下一token的分布。标准解码中 $h^L$ 只被用来采样一个离散token,之后就被丢弃。

全带宽transformer的创新正是回收这个被丢弃的 $h^L$(而非只用log2|V|比特的token),把它反馈回栈底输入,读懂这点需要知道残差流和LM头各自的角色。

门控线性单元(GLU)

一种双线性风格的融合算子:$e\otimes h=(hW_U)\odot\sigma(eW_G)$,其中 $\odot$ 是逐元素乘,$\sigma$ 是sigmoid。一个输入经线性变换作value通路,另一个经线性变换加sigmoid作乘性门,两路都是 $D\times D$ 矩阵。

本文用非对称GLU融合token嵌入与上一步顶层隐状态:隐状态走value、token只当门。这个设计封死了'退化为普通token输入'的捷径,是方法成立的关键。

并行教师强制(teacher forcing)

训练时把完整目标序列一次性喂入模型,所有位置的前向计算跨token完全并行,损失只在输出端计算。这是transformer比RNN高效的核心原因:串行性只存在于层间,不存在于token间。

潜在反馈的精确递推在第 $t$ 步依赖第 $t-1$ 步的完整前向结果,直接训练会毁掉并行性;本文的多pass训练方案就是为了保住teacher forcing,这是全文工程难点所在。

Jacobi迭代与时间并行

求解不动点方程 $x=f(x)$ 时,用上一轮迭代值 $x^{(k-1)}$ 同步更新所有分量得到 $x^{(k)}$。把这种思路搬到训练:每遍前向都用上一遍的隐状态(右移一位后融合)并行更新全部位置,序列性由遍数k支付而非由序列长度支付。

本文的多pass训练就是Jacobi式更新,理解它才能明白为何k遍训练只覆盖k-1步反馈地平线,以及为何需要讨论'训练深度之外的外推稳定性'。

研究动机

大模型预训练长期依赖“更多参数+更多token”的Scaling定律,但高质量独特数据正日益成为硬约束,作者由此追问:能否对每个token榨出更多学习信号?切入点是一个结构性缺陷:自回归transformer在水平轴(跨位置)上靠稠密attention做到全带宽——第 $\ell$ 层状态能读到缓存中所有更早位置的同层表示——但在垂直轴(跨深度)上通道极窄。形式化地,标准模型在位置 $t$、第 $\ell$ 层的可达集为 $R_{std}(t,\ell)=\{(t',\ell'): t'<t,\ \ell'<\ell\}$,规模仅 $\Theta(T\ell)$:浅层永远读不到更深层的过去状态,而“处理得最彻底”的顶层输出 $h^L$ 甚至从不进缓存。解码步间唯一能回流的信息,是被压成单个符号、至多携带 $\log_2|V|$ 比特的采样token。于是非言语化的中间计算(不确定性、部分结果、计划)要么被深度冻结在缓存里只能被其上的层读取,要么逼模型用chain-of-thought逐token叙述,每一步都为“把内心状态说出口”支付token预算。

本文的目标是本文的目标是把解码步之间的垂直反馈通道从'一个token宽'拓宽到'整个隐状态宽',让上一步的顶层隐状态 $h^L_{t-1}$ 以潜在形式重新进入栈底、获得全新的深度预算继续被加工,而不必先被言语化。同时目标附带三条强约束:其一,推理开销必须可忽略——不能像loop transformer那样每个递归步都重跑整个栈;其二,架构、KV缓存布局和serving栈必须原封不动,只改输入构造;其三,训练不能牺牲并行teacher forcing,否则在400B token规模上根本跑不起。理想情况下,这套机制还应带来可量化的收益:用更少的训练token达到标准transformer用更多token才能达到的验证损失、多选题准确率和生成任务成绩,并且在数学/代码等推理任务上产生可观测的行为差异(例如更短的推理链)。

与已有工作不同的是,在'给解码加递归'这个方向上已有多个分支,本文的独特切入在于注入点和训练方式的组合。Feedback Transformer(Fan et al., 2020)逐token顺序训练、需把各层表示聚合成记忆改变attention,不可扩展;T2MLR和Latent Recurrent Transformer在模型内部注入隐状态,需要额外MLP($5D^2$)或逐层投影($LD^2$)参数;PonderLM-2用交错嵌入/隐状态作输入,但输入长度和KV缓存都翻倍;Coconut、Soft Thinking等潜在推理工作用隐状态替代离散token,难以用标准监督训练。本文选择把注入放在模型外部——只修改输入的构造方式:上一步顶层状态与当前token嵌入经维度保持的门控融合后当作输入,因此架构零改动、仅新增两个 $D\times D$ 投影($2D^2$ 参数)、KV缓存与vLLM服务栈完全兼容。训练侧沿用时间并行(Jacobi式多pass),并把递归调度的经验规律(多少pass、何时引入)作为一等研究对象,最终在1B参数、400B token的尺度上完成了该方向此前缺失的大规模实证验证。

核心方法

直觉上,解码本就逐token串行,“浅层读不到深层过去状态”的约束在推理时毫无收益——深层状态早已算完,却无人送回栈底。方法分推理与训练两半。推理端引入潜在反馈解码:$h^L_t=f_\theta(e_t\otimes h^L_{t-1};\,C)$,融合算子 $e_t\otimes h=(hW_U)\odot\sigma(e_tW_G)$,KV缓存不变,每token额外开销仅两次 $D\times D$ 矩阵乘,实测低于1%。训练端是主要障碍:递推沿位置轴串行,直接展开会毁掉并行teacher forcing,故采用多pass时间并行:第1遍标准前向;第 $k$ 遍把上一遍状态右移一位、与嵌入融合后在全部位置并行重跑整个栈,所有pass共享NTP损失($\lambda=1$)。配料有三:pass调度(主体单pass,中后期引入两pass与3%三pass);prefix mixin(随机回退前缀为纯嵌入,模拟prompt与生成段异构);稳定性配方(深度缩放使 $\|h^L\|\sim O(1)$、RMSNorm、权重绑定、$\sigma=0.02$ 抖动)。

核心创新是非对称门控融合加外部注入这两个决定。融合算子 $e_t\otimes h_{t-1}=(h_{t-1}W_U)\odot\sigma(e_tW_G)$ 刻意不对称:隐状态占value通路,token嵌入只以乘性门的方式进入。作者明确解释了为什么不用对称融合:若用 $e_t+Wh_{t-1}$ 这类加性形式,模型可以学会压制状态通路、还原出普通token输入、从而在从标准预训练checkpoint热启动时轻松复现原有低损失,让宽通道形同虚设;而式(4)中丢弃 $h_{t-1}$ 就等于丢弃输入本身,token身份只存活于它施加在状态上的 $D$ 维门控模式里,读状态成为强制。第二个决定是注入点选在模型外部:相比Feedback Transformer改attention聚合、T2MLR/LRT在层间穿插投影,本文只改'输入是什么',因此预训练架构、KV缓存、解码循环(只改两行)和serving栈全部保留,新增参数仅 $2D^2$。与Coconut类潜在推理的本质区别是:隐状态是'增强'生成而非'替代'离散token,模型仍输出普通文本,可以用任何标准语言建模损失监督。

方法步骤详情

训练(Fig. 2左)四步。①对token序列取嵌入 $e$,标准前向得第1遍顶层状态并累计NTP损失,以保留无反馈模式处理prompt。②对其余 $k-1$ 遍,先把上一遍状态右移一位,再与嵌入做门控融合 $(hW_U)\odot\sigma(eW_G)$ 作为输入;梯度不detach,后续pass损失回传早期隐状态,充当辅助目标但增加显存。③prefix mixin:随机采样前缀长度 $p$,把 $t\le p$ 的位置回退为纯嵌入、只融合后缀,覆盖推理时的切换结构;也可对prompt整体多做一遍融合prefill。④前向得 $h^{(k)}$ 累加损失,总损失为标准NTP加上各后续遍损失均值($\lambda=1$);200B/400B运行采用75%单pass、22%两pass、3%三pass。推理(Fig. 2右):prefill后循环采样 $tok$,把输入从纯 $embed(tok)$ 换成 $(hW_U)\odot\sigma(embed(tok)W_G)$ 再前向更新KV缓存;prompt先加一遍融合prefill即FUSED——与标准解码仅一行之差。

技术新颖性

新颖性有四层。机制上:把垂直带宽形式化为可达集 $R_{lf}(t,\ell)=\{(t',\ell'): t'<t,\ 0\le\ell'\le L\}$,将标准的 $\Theta(T\ell)$ 提升到 $\Theta(TL)$,并论证收益是计算性而非信息性的——$z_{t+1}$ 本就是 $x_{1:t+1}$ 的确定函数,增益来自把全局聚合信息送进浅层的捷径。参数效率上:注入仅新增 $2D^2$ 参数,对比T2MLR的 $5D^2$ 与LRT的 $LD^2$,架构与KV缓存零改动,是同方向工作中注入最轻的方案。训练上:时间并行把长度 $T$ 的递归展开压缩为 $k$ 次位置并行评估(约 $k\times$ teacher forcing成本),并首次系统研究pass调度——发现75%单pass+25%两pass的映射在训练深度外发散,掺入仅3%三pass批次即使其变为向不动点的收缩,30步内验证损失平坦、$k=1000$ 次迭代仍稳定。副产物上:反馈训练即便不使用反馈解码也能提升表征(标准解码下LM Eval与生成任务更好),提供了一条用额外训练FLOPs换数据效率的新路径。

Standard decoding vs. latent feedback decoding.
Fig. 1: Standard decoding vs. latent feedback decoding.
Latent feedback in pseudo-code.
Fig. 2: Latent feedback in pseudo-code.
A small fraction of three-pass batches stabilizes long-horizon latent feedback.
Fig. 3: A small fraction of three-pass batches stabilizes long-horizon latent feedback.

实验结果

1B模型(Phi-4数据)预训练至多400B token。非生成任务(Fig. 4):融合prefill pass同提验证损失与10任务LM Eval;两次pass令100B≈200B、200B≈400B标准(2倍数据效率);不用反馈时仅小损验证损失且LM Eval更高。生成任务(Fig. 5):SOFT数学最强,Math500 200B 0.27→0.37超1T;FUSED代码最强,HumanEval 0.31→0.34、MBPP 0.38→0.40,GSM8K/HumanEval近1T基线。指令微调后(Table 1):GSM8K 400B FUSED 71.80(标准400B 68.39、1T 70.13);MATH-500 48.40 vs 1T 47.40;HumanEval 200B FUSED 45.92 vs 标准37.16;MBPP 41.70近1T 41.93。机制上:3%三pass令反馈映射在地平线外稳定(Fig. 3);一步递归prefill使layer-0探针达99.6%/100%(Fig. 7);SOFT推理链更短且准确率不降(Fig. 6/8)。

Latent-feedback gains carry over through instruction tuning.
Table 1: Latent-feedback gains carry over through instruction tuning.
0-shot LM Eval comparison with models of similar parameter scale (Appendix B).
Table 2: 0-shot LM Eval comparison with models of similar parameter scale (Appendix B).
Feedback passes during prefilling improve non-generative performance.
Fig. 4: Feedback passes during prefilling improve non-generative performance.
We compare the three decoding regimes defined at the start of Sec. 4.2: STANDARD, SOFT, and FUSED.
Fig. 5: We compare the three decoding regimes defined at the start of Sec. 4.2: STANDARD, SOFT, and FUSED.
Reasoning length and accuracy on Math500 from the 200B run (green line in Fig. 5).
Fig. 6: Reasoning length and accuracy on Math500 from the 200B run (green line in Fig. 5).
Full-bandwidth transformer exposes global state to shallow layers.
Fig. 7: Full-bandwidth transformer exposes global state to shallow layers.
查看结构化数据
任务指标本文基线提升
LM Eval 10任务平均(RTE/ARC/BoolQ/PIQA/WinoGrande/MMLU等,5-shot) 平均准确率(prefill施加2次反馈pass) 100B全带宽≈200B标准基线;200B全带宽≈400B标准基线 同token数标准transformer(无反馈) 约2倍预训练数据效率,增益集中于第1次pass
GSM8K(指令微调后0-shot) Pass@1 400B SOFT 71.00 / FUSED 71.80 400B标准 68.39;1T标准 70.13 +3.41 pt vs 同token标准,超过1T基线
MATH-500(基模型,0-shot) Pass@1 200B SOFT 0.37 200B标准 0.27;1T标准基线 +10 pt,反超1T token训练的标准基线
HumanEval(基模型,0-shot) Pass@3(10次rollout估计) 200B FUSED 0.34 200B标准 0.31 +3 pt;指令微调后45.92 vs 标准200B的37.16
MBPP(指令微调后0-shot) Pass@3 400B FUSED 41.70 1T标准 41.93;400B标准 40.28 以约40%的token等效计算逼近1T基线
Math500推理简洁性(基模型) 中位推理长度 + Pass@1 SOFT约450 token,准确率相同或更高 标准解码约500+ token 更短推理链、等或更高准确率(指令微调后消失)
合成状态追踪(completion/delayed memory) layer-0线性探针准确率 99.6% / 100%(一步递归prefill) 标准prefill接近随机 全局状态在输入层即完全可解码

局限与改进

作者承认两条硬限制:实验规模停留在1B参数,未验证更大(尤其更深)模型上的表现,尽管他们推测更深的模型顶层隐状态信息更丰富、收益可能更大;反馈pass调度基于启发式而非原理,未来需要更严格的消融和诸如Jacobi迭代收敛诊断之类的原则性方法。此外还有几点值得注意:3%三pass批次带来稳定性的现象只有经验证据和'收缩映射'的类比解释,缺乏理论保证;多pass训练中梯度不detach地回传早期隐状态,显存占用随pass数增长,论文未给出具体开销数字;指令微调使用标准推理产生的off-policy数据,导致基模型上'更短推理链'的优势在SFT后完全消失,说明当前后训练管道与潜在反馈解码不匹配;FUSED模式把prefill成本翻倍,对超长prompt场景的收益成本比未讨论;作者自己也强调探针可解码性只说明信息存在、不等于模型因果地使用它;另外Pass@3由每题仅10次rollout估计、温度逐方法网格搜索,统计功效有限。

独立分析的弱点

第一,稳定性缺理论:3%三pass把发散变收缩是全文最戏剧性的发现,但为什么是三pass、阈值随规模如何变化只有一张图支撑,换数据混合或深度可能失效,应引入不动点迭代的收敛分析(如谱半径估计)自适应决定pass混合。第二,监督单一:所有pass只挂NTP损失,作者承认MTP/JTP/next-latent等隐状态监督兼容但未尝试,显式监督可能进一步放大反馈通路价值。第三,后训练脱节:SFT数据由标准解码产生,使简洁推理优势在SFT后消失,应在潜在反馈解码下做on-policy RL或DPO,让偏好数据与推理时分布一致。第四,对照不完全:与T2MLR/LRT的比较停留在参数量与评测范围层面(作者也承认无资源做同尺度对比),缩放只覆盖1B/400B token,与1T基线实为不同总算力的比较。第五,工程细节缺口:长上下文下两次 $D\times D$ 融合的延迟占比、与投机解码等加速技术的兼容性未评估,顶层状态缓存仅有附录级描述。

未来方向

作者明确提出的方向包括:把方法扩展到1B以上、更深的模型(顶层状态信息更丰富,收益可能更大);用更原则化的手段取代启发式调度,例如Zeng et al. (2025)的Jacobi迭代收敛诊断来决定反馈步数与训练阶段长度;以及在潜在反馈下做on-policy后训练以保留简洁推理。在此基础上可以延伸出多条线:其一,隐状态显式监督——在多pass框架上叠加MTP/JTP/next-latent预测目标,让中间遍的状态直接被训练成'可复用的输入';其二,prefill时测试时缩放的系统化——论文已展示k次融合prefill pass单调改善困惑度与准确率,可与CoT、自一致性、束搜索等TTS手段组合出新的推理时预算分配策略;其三,与反馈解码原生兼容的数据管道——让指令数据由SOFT/FUSED解码自举生成,形成'宽通道训练-宽通道采样'的闭环;其四,机制研究——用Fig. 7式的探针方法追踪哪些任务真的把计算从token轴搬到了深度轴,回答'信息存在'与'信息被使用'之间的鸿沟;其五,把递归调度思想推广为'训练后期才引入昂贵辅助目标'的通用配方,验证其在反馈之外的辅助损失上是否同样可行。

复现评估

论文为arXiv预印本(JHU/Princeton/Microsoft),正文未给出开源代码或权重链接,也未承诺发布,这是复现的最大不确定性。有利面:方法描述极完整——训练与推理都有逐行伪代码(Fig. 2,含全部技巧的完整版在附录Fig. 9),超参全部公开:NorMuon矩阵参数学习率 $1\times10^{-2}$、权重衰减0.01,其余Adam $5\times10^{-4}$;WSD调度200步warmup、25%冷却;冷却期z-loss系数 $1\times10^{-5}$;抖动噪声 $\sigma=0.02$;全局batch 300K token;并说明vLLM兼容实现(附录D)。不利面:数据用Phi-4混合、不可公开直接下载;复现400B token主结果约需512B token等效计算、数百张高端GPU训练数周,超出多数实验室预算。务实路径:先在1B/200B token规模复现Fig. 3稳定性对比与Fig. 4 prefill反馈增益,这两个核心claim在单节点到几十卡即可验证。综合:方法可复现性高、资源门槛中高、小规模验证难度中等。