发表机构
Johns Hopkins University; Princeton University; Microsoft(约翰斯·霍普金斯大学; 普林斯顿大学; 微软公司)
机构由 AI 辅助整理,请以论文原文为准。AI 中文总结
该研究提出全带宽Transformer,通过潜在反馈拓宽解码步骤间的垂直反馈通道,经训练验证其在多项任务上性能提升,且解码开销可忽略。
AI 中文摘要
自回归Transformer沿两个轴计算:水平方向是生成的token,垂直方向是模型深度。密集注意力让每个token能广泛访问过去的内容,但解码步骤间的垂直反馈通道仍较狭窄:仅采样得到的token返回至模型栈底部,而顶层隐藏状态会被丢弃。我们提出全带宽Transformer,通过潜在反馈拓宽该通道:在每个解码步骤,将前一步的顶层隐藏状态与采样得到的token嵌入通过门控线性单元融合,作为下一步输入反馈回模型栈。潜在反馈让非言语计算能在保留标准Transformer架构、KV缓存及语言建模目标的同时,以更新后的深度预算重新进入模型栈。为在训练全带宽Transformer时不丢失并行教师强迫,我们采用调度多通目标,在预训练后期引入潜在反馈,并混入小比例的更深反馈通以保证稳定性。我们训练参数规模达10亿的全带宽Transformer,训练数据量最多达4000亿token,发现潜在反馈能提升验证损失、5次射击语言建模评估、数学与代码生成及指令调优性能。在每token解码开销可忽略的情况下,全带宽Transformer的表现与用约1.5倍token训练的标准Transformer相当或接近,且能在准确率相当或更优的情况下生成更短的推理轨迹。
英文摘要
Autoregressive transformers compute along two axes: horizontally across generated tokens, and vertically through model depth. Dense attention gives each token broad horizontal access to the past, but the vertical feedback channel between decoding steps remains narrow: only the sampled token returns to the bottom of the stack, while the top-layer hidden state is discarded. We introduce the full-bandwidth transformer, which widens this channel with latent feedback: at each decoding step, the previous top-layer hidden state is fused with the sampled token embedding through a gated linear unit and fed back as the next input. Latent feedback lets non-verbalized computation re-enter the stack with a renewed depth budget, while preserving the standard transformer architecture, KV cache, and language-modeling objective. To train full-bandwidth transformers without losing parallel teacher forcing, we use a scheduled multi-pass objective that introduces latent feedback late in pretraining and mixes a small fraction of deeper feedback passes for stability. We train 1B-parameter full-bandwidth transformers on up to 400B tokens and find that latent feedback improves validation loss, 5-shot language-model evaluation, math and coding generation, and instruction-tuned performance. With negligible per-token decoding overhead, full-bandwidth transformers match or approach standard transformers trained with roughly 1.5x more tokens, and manage to produce shorter reasoning when no off-policy templates are provided.