HyperAIHyperAI

Command Palette

Search for a command to run...

滑动窗口注意力优于线性注意力

Alexia Jolicoeur-Martineau Pashmina Cameron Rhea Sanjay Sukthanker Emy Gervais

摘要

由于二次注意力的特性,大型语言模型(LLM)消耗大量内存和能量。每个新词元(token)的成本都比前一个更高。每增加一个词元,其键(key)和值(value)就必须无限期地存储在内存中,这是不可持续的。为了解决二次缩放问题,人们提出了多种替代方案,其中之一是将LLM改造为使用线性注意力。这一想法因其有望以低成本实现最先进的性能来解决二次缩放问题而备受关注。然而,这一研究方向尚未与更简单的基线方法进行适当比较。在本工作中,我们展示了带有汇点(sink)的滑动窗口注意力(SWA)在性能上达到或优于后训练的线性注意力模型。我们在多种LLM和多个下游任务中观察到了这一结果。对于长上下文推理任务(如“大海捞针”和BABILong),SWA的性能大幅提升(比线性注意力高2到10倍)。SWA无需后训练,速度极快,且内存需求低;因此,它是一种极其廉价且可靠的解决方案。为了降低推理内存成本,我们强烈建议改用SWA,而非后训练线性模型。线性注意力模型可能显示出一些潜力,但即使要达到与SWA相当的性能,它们也可能需要从头训练或进行大量的后训练。

一句话总结

微软应用科学组及独立隶属机构的研究人员表明,带 sink 的滑动窗口注意力(SWA)在多种大语言模型和下游任务上达到或超越后训练线性注意力,在长上下文推理上取得 2102\text{--}10210 倍的性能提升,且无需后训练、内存占用更低,使 SWA 成为一种更廉价且更可靠的替代方案。

核心贡献

  • 证明带 attention sink 的滑动窗口注意力(SWA)在多种大语言模型和下游任务上达到或超越后训练线性注意力模型,且无需任何后训练。
  • 在长上下文推理基准(Needle-in-a-Haystack 和 BABILong)上,SWA 的性能比线性注意力高出 2 到 10 倍,在上下文长度 256 时恢复基线性能的 20% 到 25%,而线性注意力方法 LoLCATs 仅恢复 2.2% 和 5%。
  • 表明带 sink 的 SWA 是比线性注意力后训练更廉价且更可靠的替代方案,提供更高的解码速度、更低的内存成本,并在短上下文推理任务上恢复 99% 的基线性能。

引言

大语言模型面临一个显著瓶颈:其注意力机制随上下文长度呈二次方扩展,导致 KV 缓存带来的内存和计算成本不断增长。线性注意力方法提供了线性替代方案,但存在表达力较低、难以决定保留或遗忘哪些信息,以及训练成本高昂等问题。先前的工作如 LoLCATs 尝试以最小微调将预训练模型转换为线性注意力,但未与更简单的基线进行完整比较。

作者通过直接将后训练线性注意力模型与带 attention sink 的滑动窗口注意力(SWA)进行比较来填补这一空白,后者是一种无需训练的方法。他们证明,带 sink 的 SWA(关注前 kkk 个 token 以及前 4 个 token)在短上下文推理任务上达到或超越大多数线性化模型,恢复 99% 的基线性能。在长上下文任务上,SWA 的准确率显著更高,在 S-NIAH-3 和 BABILong 上分别恢复基线性能的 20% 和 25%,而 LoLCATs 仅为 2.2% 和 5%。作者表明,预训练模型可以在推理时直接使用 SWA,无需任何后训练或专用内核,为固定内存成本推理提供了更简单且更有效的解决方案。

方法

方法

作者基于两种互补的注意力机制设计了一种高效的混合架构:滑动窗口注意力(SWA)和线性注意力。每种机制分别解决标准 softmax 自注意力的不同局限,其组合可实现亚二次方推理成本,同时保持强劲性能。

带 Sink 的滑动窗口注意力

滑动窗口注意力不是关注所有先前 token,而是将每个查询限制为仅关注前 www 个 token。这一约束类似于卷积网络的局部感受野:经过 lll 层后,有效感受野增长至 lwl \cdot wlw,使模型能够通过深度而非直接成对注意力来聚合远距离位置的信息。形式上,SWA 计算:

xt=i=max(1,tw+1)texp(qtki/d)vii=max(1,tw+1)texp(qtki/d),t[1,,L].\mathbf{x}_t = \frac{\sum_{i = \max(1, t - w + 1)}^{t} \exp\left(\mathbf{q}_t \mathbf{k}_i^{\top} / \sqrt{d}\right) \mathbf{v}_i}{\sum_{i = \max(1, t - w + 1)}^{t} \exp\left(\mathbf{q}_t \mathbf{k}_i^{\top} / \sqrt{d}\right)}, \qquad t \in [1, \dots, L].xt=i=max(1,tw+1)texp(qtki/d)i=max(1,tw+1)texp(qtki/d)vi,t[1,,L].

经验上,SWA 通过迫使模型学习超出其局部感受野的依赖关系(而非依赖可能鼓励捷径学习的全局注意力模式),改善了长期记忆和长度外推能力。

然而,一个关键失效模式随之出现:大语言模型将不成比例的高注意力分配给前几个 token,即使这些 token 在语义上无关紧要。这些所谓的 attention sink 充当多余注意力质量的存储库。如果滑动窗口越过这些 sink token,性能会灾难性下降。作者采用了一个简单有效的修复方案:除了滑动窗口中的 w4w - 4w4 个 token 外,模型始终关注前 s=4s = 4s=4 个 token。这保证了 sink token 在每个位置都可见,防止了当它们落在窗口之外时出现的性能崩溃。重要的是,本工作仅关注带固定 sink 的免训练 SWA,避免任何额外的后训练或可学习的 sink 参数。

线性注意力

线性注意力用特征映射 ϕ\phiϕ 替换 softmax 核,使得 exp(qtki)ϕ(qt)ϕ(ki)\exp(\mathbf{q}_t \mathbf{k}_i^{\top}) \approx \phi(\mathbf{q}_t) \phi(\mathbf{k}_i)^{\top}exp(qtki)ϕ(qt)ϕ(ki)。这种分解使注意力计算可以重写为循环更新:

xt=ϕ(qt)i=1tϕ(ki)viϕ(qt)i=1tϕ(ki)=ϕ(qt)stϕ(qt)zt,\mathbf{x}_t = \frac{\phi(\mathbf{q}_t) \sum_{i=1}^{t} \phi(\mathbf{k}_i)^{\top} \mathbf{v}_i}{\phi(\mathbf{q}_t) \sum_{i=1}^{t} \phi(\mathbf{k}_i)^{\top}} = \frac{\phi(\mathbf{q}_t) \mathbf{s}_t}{\phi(\mathbf{q}_t) \mathbf{z}_t},xt=ϕ(qt)i=1tϕ(ki)ϕ(qt)i=1tϕ(ki)vi=ϕ(qt)ztϕ(qt)st,

其中状态变量在每一步增量更新:

st=st1+ϕ(kt)vt,zt=zt1+ϕ(kt).\mathbf{s}_t = \mathbf{s}_{t-1} + \phi(\mathbf{k}_t)^{\top} \mathbf{v}_t, \quad \mathbf{z}_t = \mathbf{z}_{t-1} + \phi(\mathbf{k}_t)^{\top}.st=st1+ϕ(kt)vt,zt=zt1+ϕ(kt).

这种形式相对于序列长度实现了 O(1)\mathcal{O}(1)O(1) 的推理成本,因为只需存储和更新固定大小的状态对 (st,zt)(\mathbf{s}_t, \mathbf{z}_t)(st,zt)。这消除了长上下文中二次方注意力带来的不断增长的内存占用和延迟。

实际挑战在于设计一个核 ϕ\phiϕ,使其在表达力、尖峰性和单调性这三个必要属性之间取得平衡。作者采用 Hedgehog 核,该核应用可学习的线性投影,随后进行双侧指数变换:

ϕ(x)(exp(f(x)),exp(f(x))),\phi(x) \leftarrow \left(\exp(f(x)), \exp(-f(x))\right),ϕ(x)(exp(f(x)),exp(f(x))),

其中 fff 是从维度 DDDD/2D/2D/2 的线性投影。这种构造产生的核具有足够的表达力来捕获复杂的注意力模式,同时保持稳定循环更新所需的单调性和尖峰性特征。

混合注意力的后训练

从头训练线性注意力 Transformer 成本过高,且大多数现有软件和硬件栈针对 softmax 注意力进行了优化。相反,作者通过后训练将预训练的二次方注意力大语言模型转换为线性注意力模型。使用低秩适配(LoRA),他们仅用 4000 万 token 的额外训练即可线性化模型,恢复基线性能的很大一部分。实现这一点的关键在于将表达力强的核(如 Hedgehog)与小型滑动窗口注意力组件相结合。SWA 分支保留局部高保真信息,而循环线性注意力可能会模糊这些信息;线性分支则提供高效的全局上下文聚合。这种混合设计构成了作者方法的基础,在不牺牲原始预训练模型质量的情况下实现高效推理。

实验

实验在通用知识、长上下文推理和效率基准上比较了线性注意力变体、滑动窗口注意力(SWA)和全注意力。SWA 在下游任务上持续优于线性方法,在零训练 token 的情况下恢复最多的基线性能,而 LoLCATs 等线性方法需要微调且仍落后,尤其在较长上下文下。在长上下文任务中,SWA 保持优于线性方法的准确率;在速度和内存测试中,SWA 在较小窗口尺寸下速度最快,内存成本更低或相当。总体而言,SWA 成为全注意力的最有效且最高效的替代方案。

滑动窗口注意力(SWA)在通用知识和推理基准上持续优于线性注意力方法,在 MMLU 上恢复最多的基线性能,并几乎恢复全部平均基准性能。虽然某些线性方法如 QRWKV6 在特定情况下达到或略微超过 SWA,但 SWA 在性能和训练效率之间实现了最佳权衡,且无需额外 token。SWA 在 11 种情况中的 9 种取得最高平均下游性能,仅有来自 LoLCATs 和 QRWKV6 的微小例外。SWA 恢复 MMLU 基线性能的 93.2%,为所有方法中最高,并几乎恢复全部(99.0%)平均基线性能。SWA 需要零后训练 token,而次优的高效方法 LoLCATs 使用 4000 万 token 恢复 MMLU 的 83.2% 和平均性能的 97.5%。QRWKV6 在 Qwen2.5-32B-Instruct 的 MMLU 上匹配基线,而 SWA 略有下降;DiJiang 在 Llama2.0-7B 的 MMLU 上优于 SWA。

滑动窗口注意力(SWA)在通用知识和推理基准上持续优于线性注意力变体,在大多数情况下取得最佳平均性能,且无需微调 token。SWA 还恢复了基线模型性能的高比例,尤其在平均指标上,并且是比较方法中训练效率最高的选项。SWA 在 11 种情况中的 9 种取得最高平均下游性能,仅有 LoLCATs 在 Phi-1.5 和 QRWKV6 在 Qwen2.5 上的微小例外。SWA 在无任何微调 token 的情况下恢复平均基线性能的 99.0% 和 MMLU 基线性能的 93.2%。在训练效率上最接近的竞争者 LoLCATs 需要 4000 万 token 来恢复 MMLU 的 83.2% 和平均基线性能的 97.5%。在 MMLU 上,SWA 是大多数基础模型上的最佳表现者,除 Llama2-7B 外,DiJiang 略微优于它。

在所有测试的窗口尺寸和上下文长度下,SWA 在 Single Needle-in-a-Haystack 任务上持续达到或超越 LoLCATs 和 Liger-GLA。在最长的上下文长度下,SWA 保留了全注意力准确率的有意义部分,而其他方法降至接近零。SWA 在每个窗口尺寸和上下文长度下均取得等于或高于 LoLCATs 和 Liger-GLA 的准确率。在 4K 上下文下,SWA 恢复全注意力准确率的 17.2% 到 23%,而 LoLCATs 和 Liger-GLA 最多达到 5.8% 和 0.8%。更大的窗口尺寸通常会提高所有模型的准确率,但 SWA 保持最大优势。

在 BABILong 基准上,LoLCATs(+SWA) 在短上下文长度(0K 和 1K)下略微优于 SWA,但 SWA 在较长上下文(2K 和 4K)下显示出明显优势。相对于全注意力,两种方法在 0K 时都恢复了很大一部分准确率,但在 4K 时 SWA 比 LoLCATs 保留了更多性能。在 0K 和 1K 上下文下,LoLCATs(+SWA) 得分略高于 SWA(例如,0K 时为 56% 对 55%)。在 2K 和 4K 上下文下,SWA 以较大差距超越 LoLCATs(例如,4K 时为 15% 对 3%)。在 0K 时,两种方法恢复约 74% 到 76% 的全注意力准确率,但在 4K 时 SWA 恢复 25%,而 LoLCATs 仅恢复 5%。

SWA 在通用知识和推理基准上持续优于线性注意力方法,在大多数情况下取得最佳平均性能,且无需微调 token,并恢复基线性能的高比例,尤其在平均指标上。在长上下文任务上,SWA 对 LoLCATs 和 Liger-GLA 等替代方案保持明显优势,尤其在较长上下文下其他方法降至接近零准确率时,尽管 LoLCATs 在 BABILong 的极短上下文下略微超过 SWA。总体而言,SWA 在性能和训练效率之间提供了最佳权衡,仅有特定方法在特定基准上的微小例外。


用 AI 构建 AI

从创意到上线——通过免费 AI 协同编码、开箱即用的环境和最优惠的 GPU 价格,加速您的 AI 开发。

AI 协同编码
开箱即用的 GPU
最优定价

HyperAI Newsletters

订阅我们的最新资讯
我们会在北京时间 每周一的上午九点 向您的邮箱投递本周内的最新更新
邮件发送服务由 MailChimp 提供