AB
AiBoss
Wiki

什么是分组查询注意力(Grouped-Query Attention)?

分组查询注意力(Grouped-Query Attention,GQA)是 arXiv 论文《GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints》提出的一种注意力机制,它让多个查询头共享同一组键值头,键值头数量介于多头注意力与多查询注意力之间,从而在解码速度与模型质量之间取得折中。

分组查询注意力(Grouped-Query Attention,GQA)是一种用于 Transformer 解码器的注意力变体:它把查询头(query head)分成若干组,每组共享同一份键头(key head)与值头(value head),因此键值头的数量多于一个、但少于查询头的数量。它要解决的问题是:多头注意力(Multi-Head Attention,MHA)在自回归解码时,键值缓存(KV cache)随头数线性增长,成为推理速度与显存的瓶颈;而多查询注意力(Multi-Query Attention,MQA)虽然只保留单个键值头、能大幅加速解码,却可能带来质量下降。GQA 试图在两者之间取一个中间点。

为什么重要

在 GQA 出现之前,Transformer 解码器的注意力配置基本落在两个极端上。

一端是标准的多头注意力。每个查询头都配有自己独立的键头和值头。这种配置表达能力强,是绝大多数语言模型预训练时的默认选择。代价在于解码阶段:自回归生成需要缓存历史 token 的键和值,缓存规模与键值头数量成正比。头数越多,缓存越大,每一步解码需要读取的数据量也越大,而解码本身是逐 token 串行的、受内存带宽约束的过程,因此缓存膨胀会直接拖慢推理。

另一端是多查询注意力。它把所有查询头压缩到共享唯一的一组键值头,键值缓存因此被压到最小,解码速度显著提升。但论文指出,MQA 可能导致质量下降(quality degradation)。此外,论文还提到一个工程上的顾虑:如果只是为了更快的推理而单独训练一个 MQA 模型,代价并不划算——已有的多头模型检查点无法直接复用。

于是问题变成两个:能不能把已经训练好的多头模型低成本地转成键值头更少的模型?能不能让键值头的数量成为一个可调的中间量,而不是只能在「全部独立」和「只留一个」之间二选一?GQA 正是针对这两点提出的。

工作机制

按论文的描述,GQA 的核心做法可以拆成以下几个要点。

  • 键值头数量取中间值。GQA 被定义为多查询注意力的推广:它使用的键值头数量多于一个、少于查询头的数量。也就是说,查询头被划分为若干组,组内成员共享同一组键和值。当键值头数量等于查询头数量时,它退化为多头注意力;当键值头数量为一时,它退化为多查询注意力。MHA 与 MQA 因此成为 GQA 的两个端点。
  • 从已有检查点「上训练」(uptraining)。论文提出了一套配方,把现有的多头语言模型检查点转换为带 MQA 的模型,所需算力为原始预训练算力的 5%。这里的「上训练」指的是在已有权重基础上继续训练、让模型适应新的注意力结构,而不是从头预训练。
  • 把上训练与 GQA 结合。论文展示了上训练得到的 GQA 模型可以达到接近多头注意力的质量,同时速度与 MQA 相当。这是论文给出的结论,属于该论文的实验主张,而非对所有模型与任务的普适保证。
  • 转换的关键在于键值头的初始化与聚合。把多头检查点转成键值头更少的结构时,需要把原本多个键值头的参数合并成更少的头。论文的配方正是围绕这一转换与随后的继续训练展开的。

从推理角度看,GQA 的收益来源与 MQA 相同:解码时只需缓存更少的键值头,每一步要读取的缓存数据量随之下降。由于键值头数量是可调的,实践者可以在质量与速度之间选择一个折中点,而不必被迫接受某一端的全部代价。

典型例子

论文本身给出的典型用法,是把已经预训练好的多头注意力语言模型检查点,通过上训练转换为 GQA 模型。论文报告的两个具体设定是:

  • 上训练所需的算力为原始预训练算力的 5%;
  • 上训练得到的 GQA 模型,质量接近多头注意力,速度与多查询注意力相当。

论文同时给出了 MHA、MQA、GQA 三者的对照关系:MQA 只使用单个键值头,能大幅加速解码器推理,但可能带来质量下降;GQA 使用中间数量的键值头,是 MQA 的推广。这三者的对比构成了论文实验设计的基本框架。

该论文投稿于 2023 年 5 月 22 日,修订版 v3 于 2023 年 12 月 23 日提交,被 EMNLP 2023 接收。作者为 Joshua Ainslie、James Lee-Thorp、Michiel de Jong、Yury Zemlyanskiy、Federico Lebrón、Sumit Sanghai。

边界与常见误解

第一,GQA 不是「免费的质量提升」。论文的表述是:上训练得到的 GQA 达到接近多头注意力的质量,同时速度与 MQA 相当。这里的「接近」是相对于论文所测试的模型与任务而言的,论文并未声称 GQA 在所有场景下都能无损替代多头注意力。质量与速度之间的折中依然存在,只是折中点变得可调。

第二,GQA 不等于 MQA。两者常被混为一谈,因为都涉及键值头共享。区别在于共享的粒度:MQA 是所有查询头共享唯一一组键值头,GQA 是分组共享,键值头数量严格多于一个。把 GQA 说成「MQA 的别名」是不准确的;准确的说法是 GQA 把 MQA 作为自己的一个特例包含在内。

第三,「上训练」不等于「微调一下就行」。论文给出的配方需要消耗原始预训练算力的 5%,这是一个明确的开销,并非零成本改造。同时,这套配方针对的是把多头检查点转成 MQA 或 GQA 的场景,论文并未声称它可以无条件套用到任意架构或任意训练阶段。

第四,GQA 的收益主要体现在解码推理阶段。它减少的是键值缓存的规模与读取量,因此对自回归逐 token 生成的场景更有意义。把它当作一种通用的、在所有训练与推理环节都能带来加速的手段,超出了论文所讨论的范围。

第五,论文中出现的具体数字(如 5% 的算力比例)是该项研究的实验设定与结论,不应被当作行业统一标准或对所有模型的性能承诺。

参考资料