更新于 2026年7月23日

在前面的一系列文章中,我们给大家介绍了多种以降低模型计算资源为动机而提出来的算法模型,可以看出这也成为了大模型研究领域的一个重要方向。

例如到目前为止,我们已经遇到过了多个由类似动机所提出的模型技术:「① 希望减少模型参数,但又不太损失精度——模型蒸馏;」 「② 希望加快推理速度,又能保持模型精度—— RMSNorm」「③希望减少模型重复计算过程,降低缓存使用同时提高推理效率——KV Cache」「④ 希望不仅仅只通过堆资源,而是充分利用训练数据来提高模型性能——LLaMA 1」

MHA、MQA 和 GQA 结构对比图
MHA、MQA 和 GQA 结构对比图

在接下来的两篇文章当中,将继续一口气给大家介绍由同一家公司(谷歌)中的不同部门(Google Research 和 Google Brain)分别提出来的两个基于多头注意力机制改造而来的算法,其中一个作者还是 Transformer 论文的二作。

关键词:多头注意力、self-attention、MHA、MQA、GQA、LLaMA、Transformer

1. 动机#

在前面介绍 Transformer 的时候,我们就详细介绍了多头注意力机制(Multi-Head Attention, MHA)的原理及实现过程,简单来说多头注意力机制包含两个关键元素:自注意力机制和多头。关于自注意力机制的动机原理我们这里就不再赘述,我们稍微来回顾一下提出“多头”的动机是什么。

1.1 多头动机#

30 秒时间,你也思考一下为什么要使用多头?

自注意力机制在对当前位置的信息进行编码时会将注意力过度集中于自身所在的位置,因为在计算注意力权重时自己与自己所在位置的相似性是最高的,而这可能导致模型忽略了其它位置上的信息。

如过大家还有印象的话就知道,上面这段话便是我们在介绍“多头”动机时的描述。

图 1. 多头注意力权重可视化结果[1]
图 1. 多头注意力权重可视化结果[1]

如图1所示是一张在真实场景下多头注意力中不同权重矩阵的可视化结果。可以发现,在不同的权重矩阵中注意力值的分配并不相同,并且可以明显地看出在第3张注意力权重可视化结果中,模型将过多的注意力集中到了每个字符自身所在的位置而忽略了其它位置上的信息。

因此,在一定上相同维度中多头 $h$ 的值越大,整个模型的表达能力就越强,越能使得模型对于注意力权重进行合理分配。

此时,我们终于想起原来使用“多头”的目的是为了解决模型将注意力过度集中于自身位置的问题,即让模型能拥有更多的注意力权重分配方式。

敲黑板,上面加粗是重点,等下要用到!

不过,所谓成也萧何败萧何,多头既让模型增强了表达能力,可同时也显著增加了模型的参数量,当然更重要的一点是增加了推理时 KV 缓存的开销。

所以,在这样的背景下,如何才能做到既可以满足多头的动机,同时又能最大程度的降低所使用的计算资源呢?

1.2 MQA 动机#

自多头注意力提出以来,它的优越性能就使得其成为了替代 RNN 的不二选择。尽管得益于自注意力机制在序列长度上的并行化计算,模型在训练时的速度通常要远远快于同量级的 RNN 网络;但是,由于在推理过程中生成内容需要逐时刻进行,这就意味着无法实现并行,此时模型需要反复加载大量的的 “keys” 和 “values” 状态,使得内存带宽开销很大 [2]。

While training these layers is generally fast and simple, due to parallelizability across the length of the sequence, incremental inference (where such paralleization is impossible) is often slow, due to the memory-bandwidth cost of repeatedly loading the large “keys” and “values” tensors.

图 2. MHQ 注意力 Key-Value 对示意图
图 2. MHQ 注意力 Key-Value 对示意图

从图2可以看出,在多头注意力当中每个头都有对应的 Key、Value 和 Query,然而在推理过程中,生成每一个新词模型都必须同时使用之前时刻所有的 Key 和 Value,而这些缓存(KV缓存)通常很大且不能并行,所以这就是自注意力机制的瓶颈所在。

更多关于 KV 缓存的内容,可以参见文章 「大模型为什么需要KV缓存是何动机与原理?为什么Q不用缓存?」

在这样的动机下,2019年11月 Noam Shazeer 提出了变种的多头注意力机制——多查询注意力机制(Multi Query Attention, MQA) ,即同一个 Key 和 Value 在不同的头之间共享,这样便显著地降低了在推理过程中对于计算资源的开销 [2]。

We propose a variant called multi-query attention, where the keys and values are shared across all of the different attention “heads”, greatly reducing the size of these tensors and hence the memory bandwidth requirements of incremental decoding.

图 3. MQA 注意力 Key-Value 对示意图
图 3. MQA 注意力 Key-Value 对示意图

从图3我们可以看出,此时虽然只有一个 Key 和 Value ,但是多个 Query 与同一个 Key 一样可以计算得到多个不同的注意力权重分配矩阵,然后再与同一个 Value 作用得到最后多头的计算结果。

你看,这样是不是既满足了多头的动机,又达到了降低模型参数的目的?

那这样行不行呢?感觉上好像行,但是也说不准。

怎么办?实验试一试,如果效果差不多那就证明这样是可以的,如过差很远那就证明这样的做法是不行的。

最后,经过作者实验后发现,通过 MQA 训练得到的模型在解码时的速度将会明显变快,并且与基线模型相比结果只有些许的下降。

We verify exper imentally that the resulting models can indeed be much faster to decode, and incur only minor quality degradation from the baseline.

不过这里需要注意的是,上面的结论仅仅只是在作者自己论文中实验的结论,但是实际情况是 MQA 在大模型推理中的效果并不太理想,因为单一的 Key 和 Value 大大降低了模型的表达能力,因此用得更多的是我们在下一篇文章中将要介绍的 GQA ,它是 MHA 和 MQA 之间的一种折衷做法。

不过,这依旧不影响我们去学习和了解 MQA ,因为它也是 GQA 的一种特殊情况。

这里顺便提一下 Noam Shazeer 这个人,他是 Transformer 的作者之一、然后自己提出了 MQA、接着又提出了之前我们已经介绍过的 SwiGLU,后来又参与了 Pathways Language Model (PaLM) 等等。

2. MQA 原理与实现#

在介绍完 MQA 和 GQA 各自提出的动机以后,我们再来看一下各自具体的原理部分。为了能够循序渐进地来介绍这两部分内容,我们先来带着大家简单回顾一下 MHA 的原理。

2.1 MHA 原理#

假设多头输入 Query、Key 和 Value 分别为 $X_q\in\mathbb{R}^{T_t\times d}$ 、 $X_k\in\mathbb{R}^{T_s\times d}$ 和 $X_v\in\mathbb{R}^{T_s\times d}$ ,其中 $T_s$ 和 $T_t$ 分别表示源序列和目标序列的长度,$d$ 表示模型维度;第 $i$ 个头对应的线性层映射权重为 $W^i_q\in\mathbb{R}^{d\times d_h}$ 、 $W^i_k\in\mathbb{R}^{d\times d_h}$ 和 $W^i_v\in\mathbb{R}^{d\times d_h}$ ,其中 $i=1,2,...,h$,且通常 $d=h\times d_h$ 。

此时,第 $i$ 个头的注意力计算过程可以表示为

$$ \begin{aligned} Q_i &= X_qW^i_q\in\mathbb{R}^{T_t\times d_h}\\[2ex] \quad K_i &= X_kW^i_k\in\mathbb{R}^{T_s\times d_h}\\[2ex] \quad V_i &= X_vW^i_v\in\mathbb{R}^{T_s\times d_h}\\[2ex] \text{head}_i(Q_i,K_i,V_i)&=\text{softmax}\left(\frac{Q_iK_i^T}{\sqrt{d_h}}\right)V_i\in\mathbb{R}^{T_t\times d_h} \end{aligned}\tag{1} $$

进一步,多头注意力的计算过程为

$$ \begin{aligned} Y_{\text{MHA}}=\text{Concat}(\text{head}_1,...,\text{head}_h)W^O \end{aligned} \tag{2} $$

其中 $W^O\in\mathbb{R}^{d\times d}$ 。

接着,根据式(1)我们可以得到计算 3 个线性变换的时间复杂度分别为 $O(T_tdd_hh)$ 、 $O(T_sdd_hh)$ 和 $O(T_sdd_hh)$ ,$Q_iK_i^T$ 对应的时间复杂度为 $O(T_td_hT_sh)$ ,计算注意力权重的时间复杂度为 $O(T_tT_sd_hh)$ 。因为模型超参数是固定的,MHA 的整体时间复杂度为 $O(T_tT_sd_hh)$。同时,由于 MHA 中每个头拥有独立的 Query 、Key 和 Value ,所以缓存 $K$ 和 $V$ 的空间复杂度均为 $O(T_sd_hh)$。

此时,如果仅考虑推理时的场景那么 $T_t=T_s=n$,则 MHA 整体的时间复杂度为 $O(n^2d)$ ,空间复杂度为 $O(nd)$ 。

以上整个计算过程便是 MHA 的原理,下面我们再来看 MQA 的原理。

2.2 MQA 原理#

在介绍 MQA 思想的时候我们谈到,MQA 中同一个 Key 和 Value 会在不同的头之间共享,所以仅有第 $i$ 个头对应的 Query 有线性层映射权重为 $W^i_q\in\mathbb{R}^{d\times d_h}$ ,其中 $i=1,2,...,h$ ;而 Key 和 Value 则分别只有一个 $W_k\in\mathbb{R}^{d\times d_h}$ 和 $W_v\in\mathbb{R}^{d\times d_h}$ ;通常 $d=h\times d_h$ 。

此时,第 $i$ 个头的注意力计算过程可以表示为

$$ \begin{aligned} Q_i &= X_qW^i_q\in\mathbb{R}^{T_t\times d_h}\\[2ex] \quad K &= X_kW_k\in\mathbb{R}^{T_s\times d_h}\\[2ex] \quad V &= X_vW_v\in\mathbb{R}^{T_s\times d_h}\\[2ex] \text{head}_i(Q_i,K,V)&=\text{softmax}\left(\frac{Q_iK^T}{\sqrt{d_h}}\right)V\in\mathbb{R}^{T_t\times d_h} \end{aligned}\tag{3} $$

进一步,多头注意力的计算过程为

$$ \begin{aligned} Y_{\text{MQA}}=\text{Concat}(\text{head}_1,...,\text{head}_h)W^O \end{aligned} \tag{4} $$

其中 $W^O\in\mathbb{R}^{d\times d}$ 。

同时,根据式(3)我们可以得到计算 3 个线性变换的时间复杂度分别为 $O(T_tdd_hh)$ 、 $O(T_sdd_h)$ 和 $O(T_sdd_h)$ ,$Q_iK^T$ 对应的时间复杂度为 $O(T_td_hT_sh)$ ,计算注意力权重的时间复杂度同样为 $O(T_tT_sd_hh)$ 。因为模型超参数是固定的,MHA 的整体时间复杂度为 $O(T_tT_sd_hh)$。同时,由于 MHA 中所有共享一个 Key 和 Value ,所以缓存 $K$ 和 $V$ 的空间复杂度均为 $O(T_sd_h)$。

此时,如果仅考虑推理时的场景那么 $T_t=T_s=n$,则 MHA 整体的时间复杂度为 $O(n^2d)$ ,空间复杂度为 $O(nd_h)$ 。

以上整个计算过程便是 MQA 的原理。

在看完 MQA 的原理之后脑子里就在想,明明(事后)看起来很简单且自然的想法,为什么自己却没有这样一丝的想法呢? 可能还是动机不一样了。

回想起自己在上研究生那会儿先是做聚类研究的,因为真的是出于兴趣去研究它,所以导师在给了我第一个想法并做完以后,自己脑子里面马上就蹦出来了进一步改进的做法,然后马上实验做论文,可谓是一气呵成。

所以,不仅仅每个算法模型的提出有它自己的动机,我们做每件事情的时候内心当中应该同样也要有自己的动机,这样才能做到不随波逐流。

一下扯远了,拉回来!

2.3 MQA 实现#

在介绍完 MQA 的具体原理以后,我们再来看一下如何实现这部分代码。从式(3)可以知道,只需要在初始化 $W_k$ 和 $W_v$ 时仅保留一个头就行,其它地方的代码不变,如下:

 1 class MultiQueryAttention(nn.Module):
 2     def __init__(self, embed_dim=64, num_heads=8, max_batch_size=32, max_seq_len=128):
 3         super().__init__()
 4         self.num_heads = num_heads
 5         self.embed_dim = embed_dim
 6         self.head_dim = embed_dim // num_heads
 7         self.max_batch_size = max_batch_size
 8         self.max_seq_len = max_seq_len
 9         self.wq = nn.Linear(self.embed_dim, self.embed_dim, bias=False)
10         self.wk = nn.Linear(self.embed_dim, self.head_dim, bias=False)
11         self.wv = nn.Linear(self.embed_dim, self.head_dim, bias=False)
12         self.wo = nn.Linear(self.embed_dim, self.embed_dim, bias=False)
13         self.cache_k = torch.zeros((self.max_batch_size, self.max_seq_len, self.head_dim))
14         self.cache_v = torch.zeros((self.max_batch_size, self.max_seq_len, self.head_dim))
15 
16     def forward(self, x, start_pos, mask=None):
17         bsz, seq_len, _ = x.shape
18         xq = self.wq(x)
19         # [ bsz, seq_len, embed_dim] @ [embed_dim, embed_dim] = [ bsz, seq_len, embed_dim]
20         xk, xv = self.wk(x), self.wv(x)
21         # [ bsz, seq_len, embed_dim] @ [embed_dim, head_dim] = [ bsz, seq_len, head_dim]
22         self.cache_k[:bsz, start_pos: start_pos + seq_len] = xk
23         self.cache_v[:bsz, start_pos: start_pos + seq_len] = xv
24         keys = self.cache_k[:bsz, : start_pos + seq_len]  # [bsz, cache_len + seq_len, head_dim]
25         values = self.cache_v[:bsz, : start_pos + seq_len]  # [bsz, cache_len + seq_len, head_dim]
26         print(f"keys(values) shape [bsz, cache_len + seq_len, num_heads, head_dim]: {values.shape}")
27         queries = xq.view(bsz, seq_len, self.num_heads, self.head_dim)  
28         # [bsz, seq_len, num_heads, head_dim]
29         keys, values = keys.unsqueeze(1), values.unsqueeze(1)  
30         # [bsz, 1 cache_len + seq_len, head_dim]
31         queries = queries.transpose(1, 2)  # [bsz, num_heads, seq_len, head_dim]
32         scores = torch.matmul(queries, keys.transpose(2, 3)) / math.sqrt(self.head_dim)
33         if mask is not None:  # [1, 1, seq_len, cache_len + seq_len]
34             scores = scores + mask
35         scores = F.softmax(scores.float(), dim=-1).type_as(xq)
36         output = torch.matmul(scores, values)  # [bsz, num_heads, seq_len, head_dim]
37         output = output.transpose(1, 2).contiguous().view(bsz, seq_len, -1)
38         return self.wo(output)

以上代码便是实现 MQA 的全过程,整体上同基于 KV 缓存的 MHA 多头注意力机制实现过程差异不大(这部分内容可参考上一篇文章「大模型为什么需要KV缓存是何动机与原理?为什么Q不用缓存?」),唯一的区别在于第10~11行中对于 $W_k$ 和 $W_v$ 形状的定义,以及后续在每一行代码中各个变量形状的变化。

同时,在第32行和第36行的计算过程中,因为此时 keysvalues 中只有一个头,因此在进行计算的时候底层会有一个广播机制的作用,也就是在实际计算的时候会将 keysvalues 复制到每个头中参与计算,如图4所示。

图 4.  multi-query-attention 中 keys 和 values 复制过程图
图 4. keys 和 values 复制过程图

其它部分的代码我们在这里就不再赘述,大家可以根据上述代码的标注进行理解,完整示例代码可参考Code/C07_MA/C03_MQA_kv_cache.py 文件。

3. 复杂度分析#

在介绍完 MQA 的原理及实现过程以后我们再来对比一下它与 MHA 在各个方面的差异之处。

3.1 复杂度分析对比#

经过上面内容的介绍我们分析得到,如果仅考虑推理时的场景那么 $T_t=T_s=n$,则 MHA 整体的时间复杂度为 $O(n^2d)$ ,空间复杂度为 $O(nd)$ ;MHA 整体的时间复杂度为 $O(n^2d)$ ,空间复杂度为 $O(nd_h)$ 。

可以看到,与 MHA 相比 MQA 在时间复杂度上并没有降低,这因为 Query 的数量没有变,所以仍旧需要计算 $h$ 次 $Q_iK^T$ ,所以在时间开销上并没有减少。

但是,从空间复杂度上来看,MHA 为 $O(nd)$ 而 MQA 为 $O(nd_h)$ ,所以后者相较于前者会减少约 $h$ 倍存储空间,因此MQA 对显存节省非常显著,尤其在长上下文推理时。

3.2 参数缓存对比#

我们继续再来简单看一下 MHA 和 MQA 的参数量和存储大小的差异。在 MHA 中因为每个注意力头有独立的 Query、Key、Value 投影矩阵,所以由式(1)可知其参数为 $(3dd_h)h=3d^2$ 。在 MQA 中,所有的头共享一个 Key 和 Value ,所以分别对应只有一个投影矩阵,所以由式(3)可知其参数为 $hdd_h+2dd_h=d^2+2dd_h$ 。

可以看出,MQA 参数量比 MHA 少,但减少幅度有限,通常小于 10%,因为大部分参数在 Query 和 FFN 中。

同时,由于在 MHA 中 KV Cache 随 $h$ 数线性增长,所以长序列推理时极容易出现显存爆炸的情况;但是在 MQA 中大约能节省 $h$ 倍的显存。

现在假设模型在推理时 $h=32$,$n=2048$,$d_h=128$,参数精度为 float16 (即每个参数占用 2 个字节)。由此可以得到此时如果是MHA 其对应 KV 缓存占用显存大小为 $(nd_hh2)2=2048\times128\times32\times2\times2=33554432$ 字节(32MB) ,而 MQA 对应 KV 缓存大小为 $(nd_h2)2=2048\times128\times1\times2\times2=1048576$ 字节(1MB)。

4. 总结#

综上,关于 MHQ 和 MQA 我们可以得到如下总结

对比维度 MHA MQA 提升点
参数量 略低 小幅减少
训练速度 较慢 略快 节省 KV 投影参数
推理速度 快 2-3x 少 KV 读取
空间复杂度 低很多 缓存减少 约$h$倍
KV Cache 大小 极小 长上下文显著节省

总结就是,MQA 主要在推理性能(速度 & 显存)上优势巨大,特别是长上下文场景;但是在参数量和训练速度提升有限,因为时间复杂度理论上没变,但是带宽瓶颈减少也带来了实际的加速。

引用#

[1] https://github.com/BAI-Yeqi/Statistical-Properties-of-Dot-Product/blob/master/proof.pdf

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

[3] Ainslie J, Lee-Thorp J, De Jong M, et al. Gqa: Training generalized multi-query transformer models from multi-head checkpoints[J]. arXiv preprint arXiv:2305.13245, 2023.

阅读 --

彻底搞懂KV Cache大模型推理加速的核心!

本文从 ChatGPT、Claude 等大模型首 Token 慢、其余瞬间输出的现象切入,系统讲清 KV Cache 原理:自回归推理中如何缓存每一层的 Key 与 Value 矩阵并避免每步重算历史 token,以及 …

第3节 网络结构与自注意力实现

本文是 Transformer 精读系列第3讲,详细讲解如何用 PyTorch 一步步实现 MultiHeadAttention 模块。先回顾多层 Transformer 与单层 Encoder-Decoder 结构,再系统拆解自注意力实现 …

大模型为什么需要KV缓存?为什么Q不用缓存?

本文是 LLaMA 2 模型原理铺垫篇,从大模型自回归解码的「重复计算」动机讲起,对比不采用 KV Cache 与采用 KV Cache 两套推理流程的算力开销,重点解答两个常见疑问:一是为什么只缓存 K 和 V 而不缓存 Q(Q 与当前位 …