LLM 推理优化技术:Attention 从 MHA 到 GQA

1. Multi Head Attention

目前绝大部分的 LLM 使用的是 Decoder-only Transformer[1],网络结构如下图所示:

其中 xN 表示的是 N 次重复这个部分。在这个网络结构中主要有两个部分组成,分别为多头注意力(Multi Head Attention,MHA)和前馈神经网络。在多头注意力机制中,注意力的计算公式为:

Attention(Q,K,V)=softmax(QKdk)V\mathrm{Attention}(Q,K,V)=\mathrm{softmax}(\frac{QK^\top}{\sqrt{d_k}})V

其中 Q,K,VQ,K,V 分别是查询、键、值向量。在 Decoder-only Transformer 结构中,需要在 Attention 的计算过程中加入掩码矩阵,以防止前面的 Token 与后面的 Token 做注意力的计算,因此上述的计算公式在 Decoder-only Transformer 结构中应该为:

Maksed_Attention(Q,K,V)=softmax(QKdk+M)V\mathrm{Maksed\_Attention}(Q,K,V)=\mathrm{softmax}(\frac{QK^\top}{\sqrt{d_k}}+M)V

其中,矩阵 M 称为掩码矩阵,格式为:

M=(000000000000000000000)M=\begin{pmatrix} 0 & -\infty & -\infty & -\infty & -\infty & -\infty \\ 0 & 0 & -\infty & -\infty & -\infty & -\infty \\ 0 & 0 & 0 & -\infty & -\infty & -\infty \\ 0 & 0 & 0 & 0 & -\infty & -\infty \\ 0 & 0 & 0 & 0 & 0 & -\infty \\ 0 & 0 & 0 & 0 & 0 & 0 \\ \end{pmatrix}

掩码矩阵是一个下三角矩阵,这样经过 softmax 计算后,对角线以上的部分都为 0,保证了前面的 Token 不会看见后面的 Token。

在推理的过程中,每一个 Token 在计算得到 Q,K,VQ,K,V 向量后,根据 Attention 的计算公式,得到 Attention 的值,如下图所示:

上述便是单个 Attention 的计算,在原始的 Encoder-Decoder[1] 模型中,采用的是多头注意力模型(Multi Head Attention,MHA),其作用是在同一组输入上,同时学习多种不同的表示(或者能力),然后能将其组合起来,这种好处就无需多说了,相比较只用一个头,能够增强模型的综合能力。具体来说,对于每一个头,有自己的 WQ,WK,WVW_Q,W_K,W_V,会将输入投影到不同的表示子空间中,并在子空间中计算 Attention。例如,第 ii 个头的计算如下:

Qi=XWQ(i)Q_i=XW_Q^{(i)}

Ki=XWK(i)K_i=XW_K^{(i)}

Vi=XWV(i)V_i=XW_V^{(i)}

ii 个 Head 的 Attention 的计算如下:

headi=softmax(QiKidk)Vi\mathrm{head}_i=\mathrm{softmax}(\frac{Q_iK_i^\top}{\sqrt{d_k}})V_i

最后将多个头组合起来:

MHA(X)=Concat(head1,,headh)WO\mathrm{MHA}\left(X\right)=Concat\left(head_1​,\cdots ,head_h​\right)W_O

其具体过程如下图所示:

2. KV Cache

注意到有一点,在预测新的 Token 时,最终的 Classifier 只会用到 Embedding 的最后一行,而最后一行是 Attention 的最后一行经过 FFN 运算后得到的。因此是不是每次只需要得到上图中 Attention 矩阵的最后一行就可以了。其具体过程如下图所示:

带回到 Attention 的计算过程中,实际更关注的是上一个 Token 对应的 QQ 和所有的 K,VK,V,如下图所示:

此时就考虑将历史 Token 的 K,VK,V 缓存起来,这便是 KV Cache 的基本思想。KV Cache[2] 是大模型推理性能优化的一个常用技术,该技术可以在不影响任何计算精度的前提下,通过空间换时间的思想,提高推理性能,其具体过程如下图所示:

更多的 KV Cache 相关的内容可以参阅参考文献[2]。

假设有 hh 个头,假设 QQ 的维度为 nn,上下文的长度为 ll,每一个头中在 KV Cache 中需要存储的 KK 的维度为 n×ln\times l,总的需要存储的数据量为 (n×l)×2×h\left(n\times l\right)\times 2\times h。上下文长度 ll 越长,显存的占用就越大,能否有办法进一步减少存储的数据量呢?

3. Multi-Query Attention 与 Grouped-Query Attention

GPU 是 LLM 推理的重要硬件,除了计算性能之外,显存容量是一个重要的指标,KV Cache 是一种用空间换时间的方法,如果能进一步减少 KV Cache 对显存的占用,能够放入更长的上下文,也能进一步提升推理速度。

3.1. Multi-Query Attention

MQA[3] 是在 2019 年由 Google 提出的方法,全称是 Multi-Query Attention。在 MHA 中,每一个 Head 都有自己的 KKVV,MQA 的思路很直接,直接让所有 Attention Head 共享同一个 KKVV,对于每一个 Head 有自己的 QiQ_i

Qi=XWQ(i)Q_i=XW_Q^{(i)}

KKVV 则是所有的 Head 共享的:

K=XWKK=XW_K

V=XWVV=XW_V

ii 个 Head 的 Attention 的计算如下:

headi=softmax(QiKdk)V\mathrm{head}_i=\mathrm{softmax}(\frac{Q_iK^\top}{\sqrt{d_k}})V

最后将多个头组合起来:

MQA(X)=Concat(head1,,headh)WO\mathrm{MQA}\left(X\right)=Concat\left(head_1​,\cdots ,head_h​\right)W_O​

这样,总的需要存储的数据量为 (n×l)×2\left(n\times l\right)\times 2,MQA 直接将 KV Cache 减少到了 MHA 的 1h\frac{1}{h},两者的比较如下图所示:

使用 MQA 的模型包括 PaLM、StarCoder、Gemini 等。

3.2. Grouped-Query Attention

Google 在 2013 年提出的 GQA(Grouped-Query Attention)[4] 是 MQA 和 MHA 的折中版本,GQA 的方法也很简单,它就是将所有 Head 分为 gg 个组,每组共享同一对 KKVV,对于每一个 Head 有自己的 QiQ_i

Qi=XWQ(i)Q_i=XW_Q^{(i)}

KKVV 则是每一组 jgj\in g 的 Head 共享的:

Kj=XWK(j)K_j=XW_K^{(j)}

Vj=XWV(j)V_j=XW_V^{(j)}

ii 个 Head ,其属于第 jj 个分组,其 Attention 的计算如下:

headi=softmax(QiKjdk)Vj\mathrm{head}_i=\mathrm{softmax}(\frac{Q_iK_j^\top}{\sqrt{d_k}})V_j

最后将多个头组合起来:

MQA(X)=Concat(head1,,headh)WO\mathrm{MQA}\left(X\right)=Concat\left(head_1​,\cdots ,head_h​\right)W_O​

与 MHA 以及 MQA 的对比如下图所示:

GQA 的模型包括 LLAMA2-70B,以及 LLAMA3 全系列,此外使用 GQA 的模型还有 TigerBot、DeepSeek-V1、StarCoder2、Yi、ChatGLM2、ChatGLM3、Qwen2 等,相比使用 MQA 的模型更多。

4. 总结

在 MHA 框架下,结合 KV Cache 的空间换时间,能够极大加速 LLM 的推理,但是对于推理来说,GPU 的显存资源尤其珍贵,这就出现了针对 KV Cache 进一步缩减空间的方案,MQA 和 GQA 就是针对 MHA 的精简方案,通过缩减 K 矩阵和 V 矩阵,达到缩减 KV Cache 的效果。

参考文献

[1] Vaswani A , Shazeer N , Parmar N ,et al.Attention Is All You Need[J].arXiv, 2017.DOI:10.48550/arXiv.1706.03762.

[2] LLM 推理优化技术:KV Cache

[3] Shazeer N. Fast transformer decoding: One write-head is all you need[J]. arXiv preprint arXiv:1911.02150, 2019.

[4] Ainslie J, Lee-Thorp J, De Jong M, et al. Gqa: Training generalized multi-query transformer models from multi-head checkpoints[C]//Proceedings of the 2023 conference on empirical methods in natural language processing. 2023: 4895-4901.