目录
1 传统MHA存在什么弊端
2 多查询注意力 MQA
2.1 什么是多查询注意力MQA
2.2 MQA的问题
3 分组查询注意力 GQA
3.1 什么是分组查询注意力
3.2 分组查询注意力的工作原理
4 MHA GQA MQA他们三者之间的区别和联系
4.1 GQA是MHA和MQA的推广
4.2 GQA、MHA与MQA的比较
4.3 将MHA转换为GQA
5 总结
参考文献:
abstract:
MHA:每个头各自一套 KV
分组查询注意力GQA就是分组共享键值KV
多查询注意力MQA就是所有头共享键值KV
联系:组数 = 头数 → 退化为 MHA;组数 = 1 → 退化为 MQA
1 传统MHA存在什么弊端
传统的多头注意力由于每个头都需要缓存各自的 K V,因此,多头注意力中的KV缓存增长非常快,占用大量内存。这使得推理变得缓慢且成本高昂,尤其是对于长序列和大型模型。
因此,我们需要一种更智能的方法,在不损失模型输出质量的前提下,减少KV缓存的大小。首先,让我们理解多查询注意力(MQA),这是解决该问题的首次尝试。然后我们将学习分组查询注意力(GQA),这是一个更好的解决方案。
2多查询注意力 MQA
2.1 什么是多查询注意力MQA
多查询注意力(MQA)是一种策略,所有head共享相同的键和值,但每个head仍拥有自己的查询。
让我们来分解一下这个术语:
多查询注意力 = 多查询 + 单一共享键和值
在多头注意力中,每个头都有自己的Q、K和V。在多查询注意力中,每个头都有自己的Q,但所有头共享一个K和一个V。
假设我们有8个头:
多头注意力(MHA):
- head1:Q₁、K₁、V₁
- head2:Q₂、K₂、V₂
- head3:Q₃、K₃、V₃
- ...(8个独立的K和V套装)
多查询注意力(MQA):
- head1:Q₁、K_shared、V_shared
- head2:Q₂,K_shared,V_shared
- head3:Q₃、K_shared、V_shared
- ...(所有head仅设置1K和1V)
现在,KV缓存只需要存储1组密钥和1组值,而不是8组。KV缓存体积会变小8倍。
2.2MQA的问题
MQA大幅减少内存,但存在权衡。由于所有头共享相同的键和值,模型失去了从不同角度观察输入的能力。输出质量可能会下降,训练也可能变得不稳定。
需要折中:既省内存又尽量保持质量,这就是 GQA
3分组查询注意力 GQA
3.1 什么是分组查询注意力
分组查询注意力(GQA)是一种策略,将头项划分为组,组内所有头共享相同的键和值,但每个头仍拥有自己的查询。
分组查询注意力 = 分组查询 + 每组共享键和值
简单来说,GQA不是像MHA那样给每个头有自己的K和V,也不是像MQA那样给所有头共用一个K和V,而是将头分组,每个组共享一个K和V。
学习这点的最好方法是举个例子。
假设我们有8个头,然后把它们分成两组(每组4个头):
第一组:头1,头2,头3,头4——共享K_group1,V_group1
第二组:第五个头,第六个头,第七个头,八个头——共享K_group2,V_group2
所以:
- head1:Q₁、K_group1、V_group1
- head2:Q₂,K_group1,V_group1
- head3:Q₃、K_group1、V_group1
- head4:Q₄、K_group1、V_group1
- head5:Q₅、K_group2、V_group2
- head6:Q₆、K_group2、V_group2
- head7:Q₇、K_group2、V_group2
- head8:Q₈、K_group2、V_group2
现在,我们不再像MHA那样存储8套K和V,而是只存储2套。KV缓存体积会变小4倍。
而且我们不是像MQA那样只用一个共享的K和V,而是用两个集合。所以,模型仍然有一定能力从不同角度看待输入。
GQA是MHA和MQA之间的最佳平衡点。它节省了接近MQA的内存,同时保持了接近MHA的质量。
注:在GQA中,只有键和值在组内共享。查询仍然是每个头的独立。这很重要,因为查询让每个头都能从不同的角度看待输入。通过保持查询分开,GQA保持了注意力模式的多样性。KV缓存缩小是因为我们只在推理时存储键和值,而不是查询时。
3.2分组查询注意力的工作原理
让我们一步步走过整个流程。
第一步:把头分成几组。组数是我们在训练前选择的设定。假设我们有8个头和2组。
第二步:每个头用自己的权重矩阵计算自己的查询(Q)。所以,这8个头都有各自独立的查询。这和MHA是一样的。
第三步:每个组使用该组的权重矩阵计算一个共享键(K)和一个共享值(V)。第一组计算K_group1和V_group1。第2组计算K_group2和V_group2。
第四步:每个头部使用自己的查询运行注意力机制,但共享其组的键和值。第1到第4头使用K_group1和V_group1。第5到第8号用K_group2和V_group2。
步骤5:所有头部的输出都被串接并通过最终投影,就像MHA一样。
结果输出与MHA相同。但KV缓存要小得多,因为我们只为每个组存储K和V,而不是每个头。
4 MHA GQA MQA他们三者之间的区别和联系
4.1GQA是MHA和MQA的推广
当组数=正面数时:每个群体恰好有一个头。每个头都有自己的K和V。这正是多头注意力(MHA)的体现。
当组数 = 1 时:所有头颅都在同一组。所有头部共用相同的K和V。这正是多查询注意力(MQA)。
当组数介于1到头数之间时:这就是分组查询注意力(GQA)。
4.2GQA、MHA与MQA的比较
MHA (8 query heads, 8 KV sets - one per head):
[Q1] [Q2] [Q3] [Q4] [Q5] [Q6] [Q7] [Q8]
| | | | | | | |
v v v v v v v v
[K1] [K2] [K3] [K4] [K5] [K6] [K7] [K8]
[V1] [V2] [V3] [V4] [V5] [V6] [V7] [V8]
MQA (8 query heads, 1 KV set - shared by all heads):[Q1] [Q2] [Q3] [Q4] [Q5] [Q6] [Q7] [Q8]
\ \ \ | | / / /
+-----+----+---+-----+---+----+-----+
|
v
[K_shared]
[V_shared]
GQA (8 query heads, 2 groups - 1 KV set per group):[Q1] [Q2] [Q3] [Q4] [Q5] [Q6] [Q7] [Q8]
\ | | / \ | | /
+--+----+--+ +-+----+-+
| |
v v
[K_group1] [K_group2]
[V_group1] [V_group2]
- 在MHA中,每个查询都有自己的私钥和私值。8个查询,8个KV集。最高质量,最大内存。
- 在MQA中,所有查询都指向一个共享的键和值。8个查询,1KV组。内存最少,但多样性降低。
- 在GQA中,查询被划分为多个组,每个组共享一个键和值。8个查询,2个KV组。两者兼得。
4.3将MHA转换为GQA
接下来一个大问题是:我们是否总是需要从零开始训练GQA模型?答案是否定的。
最初的GQA论文表明,我们可以将现有的多头注意力模型以非常低的成本转换为GQA模型。这叫做uptraining(上训练/继续微调)。
流程很简单:
第一步:以一个已经训练过的现有MHA模型为例。
第二步:对于每个组,取该组中所有head的键权矩阵并取平均。对Value权重矩阵也同样操作。这样我们每个组就有一个共享的密钥和一个共享的值。
步骤3:对模型进行短时间微调。原始论文显示,仅仅用约5%的预训练计算量进行上训练,就足以恢复接近全MHA的质量。
这也是GQA被迅速采用的主要原因之一。实验室无需丢弃现有的MHA模型,也不必花费大量计算量来训练新模型。他们可以直接把他们升格为GQA。
5 总结
- 多头注意力(MHA):每个head都有自己的查询、键和值。质量最好,但KV缓存在推理过程中会变得非常大。
- KV缓存问题:在文本生成过程中,模型会为每个head存储每个前单词的键和值。由于多head和长序列,这需要大量GPU内存。
- 多查询注意力(MQA):所有head共享一个键和一个值,但每个头仍然有自己的查询。KV缓存会变得非常小,但输出质量可能会下降。
- 分组查询注意力(GQA):head被划分为多个组。每个组共享一个键和一个值,而每个头仍然拥有自己的查询。这是MHA和MQA之间的最佳平衡点。
- GQA是一个概括:当组数等于头数时,GQA变为MHA。当组数为1时,GQA变为MQA。介于两者之间的都是GQA。
- 重要性:GQA节省的内存接近MQA,同时保持接近MHA的质量。它现已被应用于许多流行型号,如LLaMA 2、LLaMA 3和Mistral 7B。
- 上级培训:我们不需要从零开始训练GQA模型。我们可以拿现有的MHA模型,平均每个组内的K和V权重,然后短时间微调,转化为GQA模型。这也是GQA被迅速采用的主要原因之一。
参考文献:
Grouped Query Attention