Gated Attention for Large Language Models: Non-linearity, Sparsity, and Attention-Sink-Free

发表时间: 2025-05 · arXiv:2505.06708

原文: https://arxiv.org/abs/2505.06708

文章标题:门控注意力在大型语言模型中的应用:非线性、稀疏性与无注意力沉溺
作者/机构:Zihan Qiu∗1, Zekun Wang∗1, Bo Zheng∗1, Zeyu Huang∗2, Kaiyue Wen3, Songlin Yang4, Rui Men1, Le Yu1, Fei Huang1, Suozhi Huang5, Dayiheng LiuB1, Jingren Zhou1, Junyang LinB1
1 Qwen Team, Alibaba Group 2 University of Edinburgh 3 Stanford University 4 MIT 5 Tsinghua University

速读

一句话结论 本文通过在标准 Softmax 注意力的缩放点积输出后引入逐头特定的 Sigmoid 门控机制,在几乎不增加参数和计算延迟的情况下,显著提升了大型语言模型的性能与训练稳定性,并彻底消除了“注意力沉溺”现象。

要解决什么问题 现有的标准 Softmax 注意力机制在底层机制上存在两个关键卡点。首先是线性映射带来的表达能力瓶颈:在多头注意力计算中,值(Value)投影和最终的输出(Output)投影是两个连续的线性层。由于单个注意力头的维度通常远小于模型的整体隐藏层维度,这两个连续的线性变换实质上退化成了一个低秩线性映射。由于中间缺乏非线性激活,这种低秩特性严重限制了注意力机制对复杂特征的表达能力。其次是“注意力沉溺”引发的训练与扩展危机:标准模型在自回归过程中,倾向于将极高的注意力权重(平均高达 46.7%)分配给序列的第一个 token。这种冗余的注意力分配不仅会在隐藏状态中催生极大的激活值(巨幅激活),导致模型在 BF16 混合精度训练时极易触发数值错误和损失尖峰,从而无法容忍更高的学习率;还会破坏模型的长文本外推能力。当使用诸如 YaRN 等技术修改旋转位置编码(RoPE)以扩展上下文长度时,模型原有的注意力沉溺模式难以在无训练的情况下自适应,导致长序列性能断崖式下跌。

怎么做的 核心思路是在注意力层的特定位置插入一个轻量级的门控模块,通过动态生成依赖于当前输入的稀疏分数,对注意力信息流进行选择性调制。具体而言,标准的缩放点积注意力(SDPA)计算过程为: $$ \text{SDPA}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V $$ 作者在全面消融了门控的插入位置(如投影后、SDPA 后、最终输出后)、粒度、共享机制和激活函数后,发现最佳方案是在 SDPA 的输出之后直接应用逐头特定的乘法门控。其形式化定义为: $$ Y' = Y \cdot \sigma(XW_{\theta}) $$ 其中 $Y$ 是 SDPA 的输出,$X$ 是用于计算门控分数的输入查询(Query),$W_{\theta}$ 是每个注意力头独立拥有的可学习权重矩阵,$\sigma$ 是 Sigmoid 激活函数。这一设计通过两个关键机制精准绕开了上述卡点。第一是引入非线性。将门控操作强行置于值投影和输出投影之间,打破了原有的连续线性变换,直接为低秩映射注入了非线性表达能力。第二是引入依赖于查询的稀疏性。Sigmoid 激活函数配合逐头独立的参数,能够生成高度集中在 0 附近的稀疏门控分数。这种稀疏调制不仅能有效过滤掉与当前查询无关的上下文信息,还直接切断了模型向初始 token 倾注冗余注意力的路径。因为门控分数动态控制了输出幅度,模型不再需要依赖初始 token 作为“寄存器”来吸收多余的注意力权重,从而从根本上消除了注意力沉溺。这不仅大幅抑制了随之产生的巨幅激活,使得训练过程更加稳定,也让模型在面对上下文扩展时更具鲁棒性。

效果如何 实验在两种规模的模型上展开:在 400B tokens 上训练的 15B 参数混合专家模型(MoE,含 128 个专家),以及在高达 3.5T tokens 上训练的 1.7B 参数密集模型,训练上下文长度均设为 4096。对比基线主要包括两类路线:一是参数扩展路线,即通过增加键值头、查询头或专家数量来对齐参数量的标准模型;二是旨在提升稳定性的架构修改路线,如引入选择性计算的 Switch Head 以及用于防止残差大激活的 sandwich norm。量化结果显示,在 15B MoE 模型上,引入门控使语言建模困惑度(PPL)持续降低了超过 0.2,MMLU 提升了 2 个点。在 1.7B 密集模型上,门控机制展现出极强的训练稳定性:在 3.5T tokens 训练和 4.5e-3 的高学习率下,基线模型出现发散,而门控模型不仅稳定收敛,还几乎消除了损失尖峰。在长文本外推测试中,使用 YaRN 将上下文从 32K 扩展至 128K 时,门控模型在 RULER 基准测试上取得了比基线高出 10 个点以上的增益。此外,逐层分析证实,门控模型各层分配给第一个 token 的注意力比例从基线的 46.7% 骤降至 4.8%。在代价与局限性方面,该门控机制引入的参数量极小(不到 2M),造成的壁钟时间延迟低于 2%。作者也承认,目前主要通过经验消融来验证效果,对于非线性对整体训练动态的更广泛影响,以及消除注意力沉溺为何能直接改善长序列泛化能力的严格理论解释,仍有待进一步探索。

A1 主要贡献

本文系统性地研究了在标准 Softmax 注意力机制中引入门控(gating)机制的影响。

图 1:左:研究中应用门控操作的位置。中:在不同位置应用门控的 15B MoE 模型的性能比较(测试 PPL 和 MMLU)。在 SDPA 之后(G1)应用门控取得了最佳的综合效果。在 Value 层之后(G2)应用门控也显示出显著的改进,尤其是在 PPL 方面。右:在相同超参数下,基线模型和 SDPA 门控的 1.7B 密集模型在 3.5T token 上的训练损失比较(平滑系数 0.9)。门控带来了更低的最终损失和显著增强的训练稳定性,减少了损失尖峰。这种稳定性使得模型可能使用更高的学习率,并有助于更好的扩展。
图 1:左:研究中应用门控操作的位置。中:在不同位置应用门控的 15B MoE 模型的性能比较(测试 PPL 和 MMLU)。在 SDPA 之后(G1)应用门控取得了最佳的综合效果。在 Value 层之后(G2)应用门控也显示出显著的改进,尤其是在 PPL 方面。右:在相同超参数下,基线模型和 SDPA 门控的 1.7B 密集模型在 3.5T token 上的训练损失比较(平滑系数 0.9)。门控带来了更低的最终损失和显著增强的训练稳定性,减少了损失尖峰。这种稳定性使得模型可能使用更高的学习率,并有助于更好的扩展。

A3 背景知识

多头 Softmax 注意力机制回顾

给定输入 $X \in \mathbb{R}^{n \times d_{model}}$,其中 $n$ 是序列长度,$d_{model}$ 是模型维度,Transformer 注意力层的计算【索引 47,Attention is all you need,A Vaswani,2017,Advances in Neural Information Processing Systems】可分为四个阶段。

A2 方法细节

使用门控机制增强注意力层

表 1:门控变体性能与结果。我们在 400B token 上训练了 15A2B MoE 模型。dk 是头维度,dmodel 是模型的隐藏维度,n 是 token 数量。q 指的是查询头的数量,k 指的是键值头的数量。‘Act Func’ 是公式 5 中的激活函数。‘Score Shape’ 是输入 X ∈ Rn,dmodel 的门控分数形状。‘added param’ 表示增加的参数(百万)。
表 1:门控变体性能与结果。我们在 400B token 上训练了 15A2B MoE 模型。dk 是头维度,dmodel 是模型的隐藏维度,n 是 token 数量。q 指的是查询头的数量,k 指的是键值头的数量。‘Act Func’ 是公式 5 中的激活函数。‘Score Shape’ 是输入 X ∈ Rn,dmodel 的门控分数形状。‘added param’ 表示增加的参数(百万)。

A4 实验环境

A4 实验结果

主要结果

门控注意力在 MoE 模型中的表现

门控注意力在密集模型中的表现

表 2:不同方法在不同学习率、批大小和模型配置下的性能表现。'SDPA' 指的是在公式 3 的 SDPA 之后应用 sigmoid 门控,'sandwitch norm'【索引 16】表示在将 attention/ffn 输出添加到残差连接之前对其进行归一化。使用门控时,我们减小了 FFN 的宽度,以使所有方法的参数数量相同。'-' 表示模型在训练过程中发散。
表 2:不同方法在不同学习率、批大小和模型配置下的性能表现。'SDPA' 指的是在公式 3 的 SDPA 之后应用 sigmoid 门控,'sandwitch norm'【索引 16】表示在将 attention/ffn 输出添加到残差连接之前对其进行归一化。使用门控时,我们减小了 FFN 的宽度,以使所有方法的参数数量相同。'-' 表示模型在训练过程中发散。

分析:非线性、稀疏性与无注意力沉溺

非线性提升了注意力中低秩映射的表达能力

表 3:不同(非)线性增强方法的性能。
表 3:不同(非)线性增强方法的性能。

门控引入了输入依赖的稀疏性

图 3:SDPA 逐元素(左)、Value 逐元素(中)以及 SDPA 逐元素头共享门控(右)的门控分数均值和分布。大多数门控分数小于 0.5,表明门控分数是稀疏的。其中,SDPA 输出门控分数表现出最强的稀疏性。
图 3:SDPA 逐元素(左)、Value 逐元素(中)以及 SDPA 逐元素头共享门控(右)的门控分数均值和分布。大多数门控分数小于 0.5,表明门控分数是稀疏的。其中,SDPA 输出门控分数表现出最强的稀疏性。

SDPA 输出门控减少了注意力沉溺

SDPA 输出门控促进了上下文长度扩展

表 5:不同方法在不同序列长度下的性能表现。‘YaRN Extended’ 表示扩展上下文长度的变体。‘(values)’ 表示扩展上下文长度后性能下降的值。
表 5:不同方法在不同序列长度下的性能表现。‘YaRN Extended’ 表示扩展上下文长度的变体。‘(values)’ 表示扩展上下文长度后性能下降的值。

A7 补充细节

相关工作

A5 结论

局限性

A6 附录

A.1 Switch Head 基线实验

表 6:不同 switch head 方法在不同参数增加和配置下的性能。‘switch kv’ 和 ‘switch v’ 分别指在键值和值组件中引入选择性计算。‘Switch kv, 8top8’ 意味着有 8 个键和值映射专家,每个 token 选择 top8 专家。注意‘Switch v, 1top1’ 等同于表 1 行 (11) 中的 v Headwise Gate。
表 6:不同 switch head 方法在不同参数增加和配置下的性能。‘switch kv’ 和 ‘switch v’ 分别指在键值和值组件中引入选择性计算。‘Switch kv, 8top8’ 意味着有 8 个键和值映射专家,每个 token 选择 top8 专家。注意‘Switch v, 1top1’ 等同于表 1 行 (11) 中的 v Headwise Gate。

A.2 关于稀疏门控分数的更多讨论

图 4:门控前后的平均绝对值。基线和门控后的值相似。
图 4:门控前后的平均绝对值。基线和门控后的值相似。

图 5:门控后低于阈值的 SDPA 输出值比例(左:1e-2,右:1e-3)。我们还包括了通过将平均门控分数与门控前隐藏状态相乘得到的稀疏性度量。
图 5:门控后低于阈值的 SDPA 输出值比例(左:1e-2,右:1e-3)。我们还包括了通过将平均门控分数与门控前隐藏状态相乘得到的稀疏性度量。

A.3 逐层的巨幅激活和注意力沉溺

图 6:不同门控配置下巨幅激活和注意力沉溺现象的比较。第 1 行(基线):第 6 层后出现显著的巨幅激活和注意力沉溺。第 2 行(SDPA 门控):激活减少,未观察到注意力沉溺。第 3 行(Value 层门控):激活与第 2 行相似,但存在残余的注意力沉溺。第 4-5 行(通过跨头共享和 NS-sigmoid 减少稀疏性):巨幅激活和注意力沉溺与基线相似。
图 6:不同门控配置下巨幅激活和注意力沉溺现象的比较。第 1 行(基线):第 6 层后出现显著的巨幅激活和注意力沉溺。第 2 行(SDPA 门控):激活减少,未观察到注意力沉溺。第 3 行(Value 层门控):激活与第 2 行相似,但存在残余的注意力沉溺。第 4-5 行(通过跨头共享和 NS-sigmoid 减少稀疏性):巨幅激活和注意力沉溺与基线相似。

A.4 更多逐层门控分数分析

图 7:SDPA 输出门控变体在不同约束下门控分数的分布。
图 7:SDPA 输出门控变体在不同约束下门控分数的分布。

A.5 稳定训练的其他尝试