← 返回 2026-07-31

多头注意力残差:把深度路由查询分裂为逐子空间头 Multi-Head Attention Residuals

Cheng Luo, Zefan Cai, Junjie Hu 📅 2026-07-22 👍 10 2026-08-05 19:06
Transformer架构 Triton内核 多头注意力 持续预训练 残差连接 注意力残差 深度路由 缩放实验

把注意力残差的单查询路由多头化,零参数零额外计算即提升跨层信息读取

前置知识

残差流与 pre-norm 加法更新 (Residual Stream)

标准 pre-norm Transformer 把信息沿深度通过单一加法流传递:第 $\ell$ 层只读最近状态并写出 $h_\ell = h_{\ell-1} + f_\ell(\text{LN}(h_{\ell-1}))$。因为是运行求和,每个子层输入被耦合到紧邻前一子层输出,许多层之前算出的有用特征虽仍存在于求和中,却已无法被单独寻址——网络无法在选择性地重新读取(比如)第 3 层输出时不让它先熬过中间每一次加法更新。

这是 MHAR 要改造的对象。理解「单流、加法、只读最近状态」才能看懂注意力残差为什么把历史做成可寻址记忆、以及为何单查询会成为瓶颈。

注意力残差 (Attention Residuals, Kimi 2025)

Kimi 2025 把所有先前子层输出当作有序记忆 $S=(s_0,\dots,s_{N-1})$($s_0$ 为嵌入,后续每个为某 attention/MLP 子层原始输出),用单个学习伪查询 $q\in\mathbb{R}^d$ 对经 RMSNorm 的源打分,softmax 得深度分布 $\alpha_i=\text{softmax}(q^\top\text{RMSNorm}(s_i))$,路由输入 $\tilde{h}=\sum_i \alpha_i s_i$;子层输出再 append 为新源(replacement routing)。这是「对深度而非对 token 的注意力」,用一个共享查询路由整段历史。

MHAR 是它的严格泛化($H=1$ 即还原)。不理解单查询如何把整条深度历史压缩成一个分布、又为何对所有 $d$ 维共用,就无法理解「被迫妥协」这个核心动机。

多头注意力与「一个查询无法服务所有子空间」

自注意力之所以被拆成多头,正是因为一个查询无法同时服务每个特征子空间:不同子空间偏好关注不同 token。本文指出深度路由在结构上就是一次自注意力查询,于是继承了同样限制——单查询把所有 $d$ 个坐标的深度偏好压进一个分布 $\alpha$,代价随子空间分歧增大而增大,而分歧又随模型宽度增大而增大。

这是全文论证的主轴:既然 token 注意力要多头,深度注意力(注意力残差)也理应多头。把握这个类比才能理解零参数 reshape 为何是最自然的改法。

超连接 (Hyper-Connections, Zhu et al. 2024)

超连接用 $n$ 条并行残差流取代单流,通过逐层学习「深度连接」与「宽度连接」把流混合,是当前与注意力残差并列的「学习型连接拓扑」代表。论文把它作为最强对照之一($n=4$),在 100M/350M/1B 分别取得 $-0.032/-0.066/-0.086$ 的验证损失提升。

它是 MHAR 的主要对比方法,理解它「改连接拓扑但仍是加法式多流」与 MHAR「单流 + 多头深度路由」的差异,才能读懂 Table 1 中 MHAR 为何在每尺度都更优且开销更小。

等参数/等计算/等墙钟时间对照

评估一种架构改动是否真有效,需控制变量。MHAR 与单头注意力残差在参数($H$ 个 $d/H$ 查询合起来仍 $d$ 参数)、FLOP(只多了 $H$ 个长度 $\leq 2L+1$ 的小 softmax)、乃至墙钟时间(同样的内核结构)上都严格一致,因而「MHAR vs 单头」是完美隔离多头机制贡献的对照;而注意力残差家族相对纯基线并非等计算(物化与路由历史源有机制开销),故论文以等步数汇报。

这是论文实验设计严谨性的核心。不懂这套对照逻辑,就会误把「机制提升」与「参数/计算变多」混为一谈,从而高估或低估 MHAR 的真实贡献。

研究动机

标准 pre-norm Transformer 通过单一加法残差流传信息:每个子层只读最近状态并写出 $h_\ell = h_{\ell-1} + f_\ell(\text{LN}(h_{\ell-1}))$,有用特征可能早在很多层前就已算出,却淹没在运行求和中、不再可单独寻址。注意力残差(Kimi 2025)放松这一约束,把先前子层输出当作记忆库 $S=(s_0,\dots,s_{N-1})$,用单个学习伪查询 $q\in\mathbb{R}^d$ 做跨深度 softmax 路由 $\alpha_i=\text{softmax}(q^\top\text{RMSNorm}(s_i))$、$\tilde{h}=\sum_i \alpha_i s_i$。然而该读用一个跨整个宽度的共享查询,所有 $d$ 维子空间必须共用同一深度分布 $\alpha$。论文称之为「被迫妥协」(forced compromise):代价随子空间对「该读哪些层」分歧增大而增大,而分歧随模型宽度增大而增大——token 维度早已因同理被拆成多头,深度维度却仍是单头。

本文的目标是本文目标是把深度路由的「单头限制」彻底去掉,让每个特征子空间拥有自己独立的深度分布,同时不引入任何额外参数或计算开销,并以严格、可规模化的控制实验证明提升纯粹来自机制本身。具体地,作者希望:(1) 在等参数、等前向、等墙钟时间条件下证明多头路由相对单头的提升;(2) 在 100M/350M/1B 三个尺度从头训练,验证收益随宽度增长;(3) 提供融合 Triton 内核让深度路由在系统层面也可行(吞吐、峰值显存接近基线);(4) 通过身份保持转换把 MHAR 接到 8B 中训上验证真实大模型有效性;(5) 探明头数 $H$ 的最优取值并给出实践默认值(最终采用 $H=8$)。作者还希望证明:单查询在跨数据分布上并不稳健(web 语料上会从有益退化为有害),而多头 reshape 是让它稳健的「零成本保险」。

与已有工作不同的是,独特切入角度是「多头注意力应同时作用于 token 维度与深度维度」。已有改进残差连接的工作分两类:超连接(hyper-connections)把单流换成 $n$ 条并行流并用逐层学习连接混合,参数与显存开销大;注意力残差只在深度上加注意力但仍是单头。本文指出深度读本质是一次自注意力查询,理应和 token 注意力一样多头——而把单查询 $(d,)$ 重塑成 $(H,d/H)$ 恰好零参数、零 FLOP,且 $H=1$ 严格还原原注意力残差。这让「MHAR vs 单头」成为完美的等参数、等前向、等墙钟时间对照(连墙钟时间都一致),从而干净地隔离出多头机制本身的贡献,是改连接拓扑的工作做不到的实验纯度。配合融合内核与身份保持转换,论文把「单查询是瓶颈」从理论代价一路推到大规模可落地。

核心方法

直觉先于技术路线:注意力残差的深度路由在结构上就是一次自注意力查询,自注意力之所以要多头,正因为一个查询无法同时服务所有子空间;那么深度读也应当多头。技术路线保持 Kimi 2025 的前向流程不变——每个子层先把先前所有源(嵌入 $s_0$ 与每个 attention/MLP 子层原始输出)经 RMSNorm 后用伪查询打分、softmax 加权求和得到路由输入,再把当前子层输出作为新源 append(replacement routing)。唯一改动在 route 内核:把单查询 $q\in\mathbb{R}^d$ reshape 成 $(H,d/H)$ 的 $H$ 个头,源也按头切分成 $H$ 片,每头用各自 $q_h$ 仅对自己那一片源打分并独立 softmax,再把 $H$ 个路由结果拼回宽度 $d$。重塑零参数、零额外 FLOP,$H=1$ 即原始单头。系统层面则提供融合 Triton 前向/反向内核,解决「每子层重读全部历史」带来的 $O(L^2BTd)$ 显存带宽瓶颈。

核心创新是「多头深度路由」,与已有方法的本质区别在于:它不是加宽($H$ 个 $d/H$ 查询合起来仍 $d$ 参数,$H$ 只是 reshape)、不是加流(仍是单残差流)、不改前向流程(forward 与 Kimi 2025 逐字相同)。公式:对头 $h$,$\alpha^{(h)}_i=\text{softmax}_i(q_h^\top\text{RMSNorm}(s_i)[h])$,$\tilde{h}[h]=\sum_i \alpha^{(h)}_i s_{i,[h]}$,再把 $H$ 个切片拼回 $d$。每头只对自己那片源打分与混合,所以 $H$ 个 softmax 完全独立,不同子空间可关注不同层,把读从「全 $d$ 维共用一个秩-1 分布」变成「块对角的多分布」。三性质:等参数($H$ 个查询合共 $d$ 参数)、等 FLOP(仅多了 $H$ 个长度 $\leq 2L+1$ 的小 softmax,可忽略)、严格泛化($H=1$ 还原原始注意力残差)。作者还把「被迫妥协」量化为「代价随子空间分歧增长、分歧随宽度增长」,在损失与训练查询探针上双重验证。

方法步骤详情

从头训练流程(Figure 3 伪代码):① 初始化源列表 $S=[s_0]$(token 嵌入);② 对每个块的 attention 子层,调用 route 得路由输入——堆叠源为 $V\in\mathbb{R}^{N\times B\times T\times D}$,投影权重 view 成 $(H,D/H)$,源按头切分,einsum 算 logits、对深度轴 softmax(0) 得 $H$ 个独立分布,再混合并 reshape 回 $(B,T,D)$;③ $h$ 经 LN 送入 attention,输出 append 进 $S$;④ MLP 子层重复 route+LN+mlp,输出也 append;⑤ 训练用 AdamW、20K 步、cosine 调度、bf16,伪查询零初始化使深度 softmax 起始为源的均匀平均。8B 中训用 delta 变体($h=\text{partial}+\alpha\cdot\text{routed}$,源按块聚合),输出门零初始化使第 0 步与基线数值完全一致路由在中训中自然开启而无 loss 尖峰。

技术新颖性

技术新颖性四点:(i) 把「被迫妥协」形式化为可量化代价——单查询损失取决于子空间分歧,而分歧随宽度增长,作者在损失(§3)与训练查询探针(Appendix D,头间偏差近不相关、复制于不相交文本 $r=0.77$)上双重验证。(ii) 零成本多头化——reshape $(d,)\to(H,d/H)$ 是唯一改动,使 MHAR-vs-单头成为等参数、等 FLOP、等墙钟时间的完美对照,单头「给更多时间/参数也弥合不了差距」。(iii) 融合 Triton 内核——深度路由是显存带宽瓶颈(每子层重读全部历史,$O(L^2BTd)$ 流量/步),融合确定性前向/反向内核把每微批路由时间相对 torch.compile 缩短约 2×、相对 eager 约 6.4×(8×H100 端到端中训吞吐约 +30%)。(iv) 身份保持转换用 delta 形式可无缝嫁接到预训练模型做中训,第 0 步数值恒等、整个 9500 步每步差距保持在批噪声的 1/36 以下,使「在已有 8B 上续训而不掉点」成为可复现工程路径。

The depth read is an attention, so it should be multi-head.
Figure 1: The depth read is an attention, so it should be multi-head.
What each method lets the current sublayer read across depth.
Figure 2: What each method lets the current sublayer read across depth.
Multi-Head Attention Residuals pseudocode.
Figure 3: Multi-Head Attention Residuals pseudocode.

实验结果

核心发现按实验分析。(1) 从头预训练(Table 1,统一 LR 5e-4):100M/350M/1B 上 MHAR 相对 baseline 的 Δ 为 -0.061/-0.149/-0.140,在每尺度都是四种对照方法中最佳,收益随宽度持续增长。(2) 头数消融(Figure 5):1B 上 H=1→3.270、H=4→3.132、H=8→3.129(持平)、H=16→3.173,U 形平底盆地;H=16 过分裂还回约三分之一/一半收益,证实头数是真实设计轴而非免费旋钮。(3) 下游(Table 2):WikiText-2 PPL 与 LAMBADA 三尺度均改善(LAMBADA 随尺度 +2.4/+3.8/+7.2 点)。(4) 8B 中训(Table 3,Marin-8B,日程匹配):+3.2 GSM8K(p=0.004)、+3.1 GPQA(p=0.038)。(5) 融合内核让 MHAR 训练吞吐达 0.88×/0.71×/0.55× 基线、显存近似基线,单头在 web 语料 350M/1B 反劣于基线,而 MHAR 两语料各尺度都改善,证实多头是零成本保险。

From-scratch validation loss after a unified 20K-step schedule on anneal_pt_v3.
Table 1: From-scratch validation loss after a unified 20K-step schedule on anneal_pt_v3.
Zero-shot downstream evaluation at 100M, 350M, and 1B.
Table 2: Zero-shot downstream evaluation at 100M, 350M, and 1B.
Downstream accuracy after 8B mid-training on anneal_pt_v3 (final checkpoints, EMA).
Table 3: Downstream accuracy after 8B mid-training on anneal_pt_v3 (final checkpoints, EMA).
Training speed and memory.
Table 4: Training speed and memory.
Routing-operation speedup, isolating the kernels.
Table 5: Routing-operation speedup, isolating the kernels.
The H heads carry genuinely different depth-links.
Figure 4: The H heads carry genuinely different depth-links.
Head count at 1B: a flat H=4–8 optimum.
Figure 5: Head count at 1B: a flat H=4–8 optimum.
查看结构化数据
任务指标本文基线提升
从头预训练验证损失(anneal_pt_v3 语料,20K 步,统一 LR 5e-4) 验证损失 (↓, 尾部均值) MHAR(H=8):100M 2.969、350M 2.848、1B 2.754 标准 Transformer:100M 3.031、350M 2.997、1B 2.894 Δ=-0.061/-0.149/-0.140;收益从 100M 到更大尺度持续增长,每尺度居四种方法之首
零样本下游评测(LM Evaluation Harness) WikiText-2 PPL ↓ / LAMBADA ↑ 1B:PPL 45.4、LAMBADA 16.0% 1B baseline:PPL 56.4、LAMBADA 8.8% PPL 相对降幅随尺度扩大(100M -9%、350M -20%、1B -19%);LAMBADA 单调 +2.4/+3.8/+7.2 点
8B 持续预训练下游准确率(Marin-8B 基座,约 10B token,日程匹配对照) GSM8K / GPQA 准确率 (↑) GSM8K 0.502、GPQA 0.346 plain-CPT 对照:GSM8K 0.470、GPQA 0.315 +3.2 GSM8K(配对 p=0.004)、+3.1 GPQA(p=0.038);MMLU/MATH/代码统计上不变
训练吞吐与峰值显存(融合 Triton 内核) 相对基线吞吐 / 峰值显存 MHAR+融合内核:100M 0.88×、350M 0.71×、1B 0.55×;显存 42.0/20.0/20.1 GB 标准 Transformer:1.00×;显存 41.5/19.4/19.0 GB 路由操作单独加速 2.0-5.3×;MHAR 仅 +0.02% 参数,相对 torch.compile 路由快约 2×

局限与改进

作者承认的局限:(1) 注意力残差家族相对纯基线并非等计算——物化并路由历史源有机制开销(继承自 Kimi 2025),论文以等步数汇报,直接等计算研究留作未来工作。(2) 8B 中训路由收益「modest 且集中在 GSM8K/GPQA」,MMLU/MATH/代码统计上无变化,与更激进再训练调度的交互未知。(3) 1B 单次验证通过有约 ±0.07 噪声,靠尾部均值与配对比较缓解;(4) 收益有语料依赖:在 anneal 语料上单头已捕获大部分路由收益(MHAR-单头仅 -0.001/-0.028/-0.006)。我补充观察:(a) 从头最大尺度仅 1B,更大尺度是否保持 U 形最优与收益增长未验证;(b) 下游基准偏窄,作者也承认 broader benchmarks 是未来工作;(c) H=4 与 H=8 在 1B 实际持平,选 H=8 主要靠「与 KV 头对齐」但作者自承无可测量增益,选择偏保守;(d) 即便有融合内核,1B 吞吐仍仅 0.55× 基线,系统成本非平凡。

独立分析的弱点

独立弱点分析:(1) 头数最优的理论解释不足——论文给「被迫妥协随宽度增长」的直觉与探针证据,但为何平底盆地恰落在 H=4-8(而非与 $d$ 或 KV 数成某比例)缺闭合理论;改进方向:建立子空间分歧与最优 $H$ 的定量模型,给出按宽度自适应的 $H$ 调度。(2) 仅 delta 形式做了中训,block/full 形式在大模型上的效果与稳定性未报告;可系统比较三种形式。(3) 等计算对照缺失——虽融合内核把吞吐拉到 0.55-0.88× 基线,但「等墙钟时间下基线多跑 1.1-1.8× 步能否追平」未直接实验;改进方向:做严格等 FLOP 对比。(4) 8B 收益集中在数学/推理任务,通用能力(MMLU/代码)无显著提升;可在更大基座、更长中训或后训练/RL 阶段验证。(5) 系统成本仍非平凡,1B 吞吐仅 0.55× 基线;可进一步优化内核或探索稀疏/块聚合路由降低 $O(L^2)$ 流量。(6) 单查询在 web 语料上从有益退化到有害,说明机制对数据分布敏感,工程上需在目标分布单独验证。

未来方向

作者明确提出的:(1) 直接的等计算(iso-compute)对比纯基线研究;(2) 更广下游基准评测(broader benchmarks);(3) 探究为何最优落在 H=4-8、各路由头具体学到关注哪些层(论文结尾称之为 natural next steps);(4) 8B 中训路由收益如何与更激进再训练调度交互。基于成果可延伸的:(a) 把 MHAR 与其他改残差/连接方法(超连接、REPA 等)正交叠加——因 MHAR 只动 route 内核,理论上可与多数加速方法叠加;(b) 把「深度维度多头化」思想迁移到其他跨层路由机制(如 MoE 路由、长上下文记忆路由);(c) 研究训练过程中头的专门化动力学——附录 D 显示头间偏差近不相关,可设计课程或正则引导更有用的分工;(d) 在多模态/编码器-解码器架构中验证,因为残差流的「被迫妥协」是普遍现象;(e) 结合 speculative/异步路由降低推理时的深度读取开销。

复现评估

复现评估:开源情况较好——8B 中训的 MHAR 模型与日程匹配对照在 HuggingFace 发布(wdlctc/marin-8b-cpt-mhar-delta 与 marin-8b-cpt-plain-lr4e4),Figure 3 给出完整 route 伪代码(约 12 行 einsum+reshape),核心改动极小、易于原型验证。数据公开:anneal_pt_v3 基于 Nemotron-CC-v2、Nemotron pretraining SFT/STEM/code、Nemotron-CC-Math、FinePDFs、Stack-Edu、arXiv、Wikipedia 等公开语料组合。超参已给全(统一峰值 LR 5e-4、20K 步、cosine 调度、bf16、伪查询零初始化、H=8 默认)。算力门槛高:从头训 100M/350M/1B 各 20K 步、8B 中训约 10B token 在 8×H100 上,且融合 Triton 内核需自实现(Appendix E 给完整算法与数值验证)。难度中等偏高:方法本身极易实现,但效率上不退化必须用融合内核,这是主要工程门槛。