大模型面试基础 | Transformer-MHA、MQA、GQA以及MLA技术区别
Transformer的MHA、MQA、GQA以及MLA区别
在具体说MHA、MQA、GQA以及MLA这些技术之前,我们需要有一个印象,这些技术是在Transformer哪一个组件进行优化的,如下图所示,这里展示了Transformer Decoder所有模块的框架图:

我们主要优化的组件就是Q、K、V三个内容,并且为什么在做推理时候,我们只使用KVCache做缓存,而不是QKV Cache,可以看我上一篇文章
多头注意力 (MHA - Multi-Head Attention)
多头注意力 (MHA) 是原始Transformer论文中提出的标准注意力机制。其核心思想:将输入序列的键(Key)、值(Value)和查询(Query)投影到h个(头数)不同的低维空间(子空间)。K、V、Q分别维护h个独立的线性投影权重矩阵,如图所示:
计算步骤
输入
- 输入序列:X∈Rn×dmodel\mathbf{X}\in\mathbb{R}^{n\times d_{\text{model}}}X∈Rn×dmodel
- 头数:HHH
- 每头维度:dk=dmodel/Hd_k = d_{\text{model}}/Hdk=dmodel/H
- K、V带独立的下标h,独立的头空间
-
线性映射(每头独立)
Qh=X WhQ,WhQ∈Rdmodel×dkKh=X WhK,WhK∈Rdmodel×dkVh=X WhV,WhV∈Rdmodel×dk \begin{aligned} \mathbf{Q}_h &= \mathbf{X}\,\mathbf{W}_h^Q,\quad \mathbf{W}_h^Q\in\mathbb{R}^{d_{\text{model}}\times d_k} \\[4pt] \mathbf{K}_h &= \mathbf{X}\,\mathbf{W}_h^K,\quad \mathbf{W}_h^K\in\mathbb{R}^{d_{\text{model}}\times d_k} \\[4pt] \mathbf{V}_h &= \mathbf{X}\,\mathbf{W}_h^V,\quad \mathbf{W}_h^V\in\mathbb{R}^{d_{\text{model}}\times d_k} \end{aligned} QhKhVh=XWhQ,WhQ∈Rdmodel×dk=XWhK,WhK∈Rdmodel×dk=XWhV,WhV∈Rdmodel×dk -
缩放点积注意力(每头)
headh=softmax (QhKh⊤dk)Vh \text{head}_h = \text{softmax}\!\left(\frac{\mathbf{Q}_h\mathbf{K}_h^\top}{\sqrt{d_k}}\right)\mathbf{V}_h headh=softmax(dkQhKh⊤)Vh -
拼接 + 输出映射
MHA(X)=Concat(head1,…,headH) WO,WO∈RHdk×dmodel \text{MHA}(\mathbf{X}) = \text{Concat}(\text{head}_1,\dots,\text{head}_H)\,\mathbf{W}^O,\quad \mathbf{W}^O\in\mathbb{R}^{H d_k\times d_{\text{model}}} MHA(X)=Concat(head1,…,headH)WO,WO∈RHdk×dmodel
KV Cache占用
如下面公式所示,KV Cache大小为num_heads(h头的个数) * seq_len(token长度) * dim(隐藏层):
Kh=X WhK,WhK∈Rdmodel×dkVh=X WhV,WhV∈Rdmodel×dk \begin{aligned} \mathbf{K}_h &= \mathbf{X}\,\mathbf{W}_h^K,\quad \mathbf{W}_h^K\in\mathbb{R}^{d_{\text{model}}\times d_k} \\[4pt] \mathbf{V}_h &= \mathbf{X}\,\mathbf{W}_h^V,\quad \mathbf{W}_h^V\in\mathbb{R}^{d_{\text{model}}\times d_k} \end{aligned} KhVh=XWhK,WhK∈Rdmodel×dk=XWhV,WhV∈Rdmodel×dk
因为每个Attention的头,都需要有独立的K、V进行计算,其消耗的内存会比较多。
Multi-Query Attention(MQA)
MQA是MHA的一种极端简化形式,旨在显著减少KV Cache的大小。
核心区别:所有头共享同一组K\mathbf{K}K与V\mathbf{V}V,仅查询Q\mathbf{Q}Q仍为多头。由下面公式可看到K,V\mathbf{K},\mathbf{V}K,V不再带下标hhh,所有头共用。
计算步骤
输入
- 输入序列:X∈Rn×dmodel\mathbf{X}\in\mathbb{R}^{n\times d_{\text{model}}}X∈Rn×dmodel
- 头数:HHH
- 每头维度:dk=dmodel/Hd_k = d_{\text{model}}/Hdk=dmodel/H
- K,V\mathbf{K},\mathbf{V}K,V不再带下标hhh,所有头共用
-
线性映射(共享K/V)
Qh=X WhQ,WhQ∈Rdmodel×dkK=X WK,WK∈Rdmodel×dkV=X WV,WV∈Rdmodel×dk \begin{aligned} \mathbf{Q}_h &= \mathbf{X}\,\mathbf{W}_h^Q,\quad \mathbf{W}_h^Q\in\mathbb{R}^{d_{\text{model}}\times d_k} \\[4pt] \mathbf{K} &= \mathbf{X}\,\mathbf{W}^K,\quad \mathbf{W}^K\in\mathbb{R}^{d_{\text{model}}\times d_k} \\[4pt] \mathbf{V} &= \mathbf{X}\,\mathbf{W}^V,\quad \mathbf{W}^V\in\mathbb{R}^{d_{\text{model}}\times d_k} \end{aligned} QhKV=XWhQ,WhQ∈Rdmodel×dk=XWK,WK∈Rdmodel×dk=XWV,WV∈Rdmodel×dk -
缩放点积注意力(每头)
headh=softmax (QhK⊤dk)V \text{head}_h = \text{softmax}\!\left(\frac{\mathbf{Q}_h\mathbf{K}^\top}{\sqrt{d_k}}\right)\mathbf{V} headh=softmax(dkQhK⊤)V -
拼接 + 输出映射
MQA(X)=Concat(head1,…,headH) WO,WO∈RHdk×dmodel \text{MQA}(\mathbf{X}) = \text{Concat}(\text{head}_1,\dots,\text{head}_H)\,\mathbf{W}^O,\quad \mathbf{W}^O\in\mathbb{R}^{H d_k\times d_{\text{model}}} MQA(X)=Concat(head1,…,headH)WO,WO∈RHdk×dmodel
KV Cache占用
如下面公式所示,KV Cache大小为1 (h头的个数为1) * seq_len(token长度) * dim(隐藏层):
K=X WK,WK∈Rdmodel×dkV=X WV,WV∈Rdmodel×dk \begin{aligned} \mathbf{K} &= \mathbf{X}\,\mathbf{W}^K,\quad \mathbf{W}^K\in\mathbb{R}^{d_{\text{model}}\times d_k} \\[4pt] \mathbf{V} &= \mathbf{X}\,\mathbf{W}^V,\quad \mathbf{W}^V\in\mathbb{R}^{d_{\text{model}}\times d_k} \end{aligned} KV=XWK,WK∈Rdmodel×dk=XWV,WV∈Rdmodel×dk
Grouped-Query Attention(GQA)
GQA介于MHA与MQA之间,通过将查询头分组共享KV,在性能与KV Cache开销间取得折中。

核心区别:
- 查询头仍保持多头(共H个)
- KV按分组数G共享,每组包含H/G个查询头
- 因此Kg,Vg\mathbf{K}_g,\mathbf{V}_gKg,Vg带下标g∈{1,…,G}g \in \{1,\dots,G\}g∈{1,…,G}
计算步骤
输入
- 输入序列:X∈Rn×dmodel\mathbf{X}\in\mathbb{R}^{n\times d_{\text{model}}}X∈Rn×dmodel
- 头数:HHH
- 分组数:GGG
- 每头维度:dk=dmodel/Hd_k = d_{\text{model}}/Hdk=dmodel/H
-
线性映射(分组共享KV)
对每组g=1,…,Gg=1,\dots,Gg=1,…,G:
Kg=X WgK,WgK∈Rdmodel×dkVg=X WgV,WgV∈Rdmodel×dk \begin{aligned} \mathbf{K}_g &= \mathbf{X}\,\mathbf{W}_g^K,\quad \mathbf{W}_g^K\in\mathbb{R}^{d_{\text{model}}\times d_k} \\[4pt] \mathbf{V}_g &= \mathbf{X}\,\mathbf{W}_g^V,\quad \mathbf{W}_g^V\in\mathbb{R}^{d_{\text{model}}\times d_k} \end{aligned} KgVg=XWgK,WgK∈Rdmodel×dk=XWgV,WgV∈Rdmodel×dk
对每个查询头仍保持多头(共H个):
Qh=X WhQ,WhQ∈Rdmodel×dk \mathbf{Q}_h = \mathbf{X}\,\mathbf{W}_h^Q,\quad \mathbf{W}_h^Q\in\mathbb{R}^{d_{\text{model}}\times d_k} Qh=XWhQ,WhQ∈Rdmodel×dk -
缩放点积注意力(组内共享KV)
对属于组ggg的每个头hhh:
headh=softmax (QhKg⊤dk)Vg \text{head}_h = \text{softmax}\!\left(\frac{\mathbf{Q}_h\mathbf{K}_g^\top}{\sqrt{d_k}}\right)\mathbf{V}_g headh=softmax(dkQhKg⊤)Vg -
拼接 + 输出映射
GQA(X)=Concat(head1,…,headH) WO,WO∈RHdk×dmodel \text{GQA}(\mathbf{X}) = \text{Concat}(\text{head}_1,\dots,\text{head}_H)\,\mathbf{W}^O,\quad \mathbf{W}^O\in\mathbb{R}^{H d_k\times d_{\text{model}}} GQA(X)=Concat(head1,…,headH)WO,WO∈RHdk×dmodel
KV Cache占用
每组共享一对KV,故缓存大小为G×seqlen×dimG \times seq_{len} \times {dim}G×seqlen×dim
当G=H退化为MHA;当G=1退化为MQA
Multi-Head Latent Attention (MLA)

MLA通过低秩压缩键值减少KV Cache,同时保持多头查询表达能力。
核心区别
- 键值先经共享低秩压缩到潜在维度dlat≪dmodeld_{\text{lat}}\ll d_{\text{model}}dlat≪dmodel
- 再按头展开,保证多头查询的独立性
- 因此K,V\mathbf{K},\mathbf{V}K,V不再直接参与注意力计算,而是先通过低秩的方式进行压缩再映射
计算步骤
输入
- 输入序列:X∈Rn×dmodel\mathbf{X}\in\mathbb{R}^{n\times d_{\text{model}}}X∈Rn×dmodel
- 头数:HHH
- 每头维度:dk=dmodel/Hd_k = d_{\text{model}}/Hdk=dmodel/H
- 潜在维度:dlatd_{\text{lat}}dlat(通常dlat≪dmodeld_{\text{lat}}\ll d_{\text{model}}dlat≪dmodel)
- 低秩压缩(共享)
Zk=X Ck,Ck∈Rdmodel×dlatZv=X Cv,Cv∈Rdmodel×dlat \begin{aligned} \mathbf{Z}_k &= \mathbf{X}\,\mathbf{C}_k,\quad \mathbf{C}_k\in\mathbb{R}^{d_{\text{model}}\times d_{\text{lat}}} \\[4pt] \mathbf{Z}_v &= \mathbf{X}\,\mathbf{C}_v,\quad \mathbf{C}_v\in\mathbb{R}^{d_{\text{model}}\times d_{\text{lat}}} \end{aligned} ZkZv=XCk,Ck∈Rdmodel×dlat=XCv,Cv∈Rdmodel×dlat
C_k和C_v分别是键(Key)的低秩压缩矩阵和值(Value)的低秩压缩矩阵,其目的在于将原始高维输入X投影到一个低维的潜在空间(latent space),维度从d_model(例如4096)压缩到d_lat(例如128)。这使得后续生成的键缓存(KV Cache)大小从O(n * d_model)显著减少到O(n * d_lat),从而极大地节省了推理时的内存占用。
-
按头映射
对每个头h=1,…,Hh=1,\dots,Hh=1,…,H:
Qh=X WhQ,WhQ∈Rdmodel×dkkh=Zk Uhk,Uhk∈Rdlat×dkvh=Zv Uhv,Uhv∈Rdlat×dk \begin{aligned} \mathbf{Q}_h &= \mathbf{X}\,\mathbf{W}_h^Q,\quad \mathbf{W}_h^Q\in\mathbb{R}^{d_{\text{model}}\times d_k} \\[4pt] \mathbf{k}_h &= \mathbf{Z}_k\,\mathbf{U}_h^k,\quad \mathbf{U}_h^k\in\mathbb{R}^{d_{\text{lat}}\times d_k} \\[4pt] \mathbf{v}_h &= \mathbf{Z}_v\,\mathbf{U}_h^v,\quad \mathbf{U}_h^v\in\mathbb{R}^{d_{\text{lat}}\times d_k} \end{aligned} Qhkhvh=XWhQ,WhQ∈Rdmodel×dk=ZkUhk,Uhk∈Rdlat×dk=ZvUhv,Uhv∈Rdlat×dk其中:
- UhK,Uhv∈Rdlat×dkU^{K}_h, U^{v}_h \in \mathbb{R}^{d_{lat} \times d_k}UhK,Uhv∈Rdlat×dk:每头K与V的上投影(Up)
-
缩放点积注意力(每头)
headh=softmax (Qhkh⊤dk)vh \text{head}_h = \text{softmax}\!\left(\frac{\mathbf{Q}_h\mathbf{k}_h^\top}{\sqrt{d_k}}\right)\mathbf{v}_h headh=softmax(dkQhkh⊤)vh -
拼接 + 输出映射
MLA(X)=Concat(head1,…,headH) WO,WO∈RHdk×dmodel \text{MLA}(\mathbf{X}) = \text{Concat}(\text{head}_1,\dots,\text{head}_H)\,\mathbf{W}^O,\quad \mathbf{W}^O\in\mathbb{R}^{H d_k\times d_{\text{model}}} MLA(X)=Concat(head1,…,headH)WO,WO∈RHdk×dmodel
KV Cache占用
仅需缓存压缩后的Zk,Zv\mathbf{Z}_k,\mathbf{Z}_vZk,Zv,大小为:
2×seqlen×dlat 2 \times seq_{len} \times d_{lat} 2×seqlen×dlat
与头数HHH无关,显著低于MHA(Multi-Head Attention)与GQA(Grouped-Query Attention);当dlat≪dmodeld_{\text{lat}}\ll d_{\text{model}}dlat≪dmodel时,内存占用接近MQA甚至更低,同时保持MHA级别的表达力。
总结
| 场景 | 推荐技术 | KV Cache大小 | 理由 |
|---|---|---|---|
| 追求最佳效果且资源充足 | MHA | H×seqlen×dimH \times seq_{len} \times dimH×seqlen×dim | 提供最强的模型表达能力,但KV占用最大 |
| 构建现代大语言模型 | GQA | G×seqlen×dimG \times seq_{len} \times dimG×seqlen×dim | 效率与效果的最佳平衡,KV占用适中(通常G=4-8) |
| 部署于极端资源受限环境 | MQA | 1×seqlen×dim1 \times seq_{len} \times dim1×seqlen×dim | 最大化推理性能,KV占用最小,仅为MHA的1/H |
| 处理超长序列(长文档、基因序列等) | MLA | 2×seqlen×dlat2 \times seq_{len} \times d_{lat}2×seqlen×dlat | 解决O(n²)计算复杂度瓶颈,KV占用极低(dlat≪dimd_{lat} \ll dimdlat≪dim) |
KV Cache占用对比示例
假设典型参数:H=32H=32H=32, dim=4096dim=4096dim=4096, seqlen=2048seq_{len}=2048seqlen=2048, G=8G=8G=8, dlat=128d_{lat}=128dlat=128
| 技术 | KV Cache大小 | 相对MHA的比例 |
|---|---|---|
| MHA | 32×2048×4096=268.432 \times 2048 \times 4096 = 268.432×2048×4096=268.4 MB | 100% |
| GQA | 8×2048×4096=67.18 \times 2048 \times 4096 = 67.18×2048×4096=67.1 MB | 25% |
| MQA | 1×2048×4096=8.41 \times 2048 \times 4096 = 8.41×2048×4096=8.4 MB | 3.1% |
| MLA | 2×2048×128=0.52 \times 2048 \times 128 = 0.52×2048×128=0.5 MB | 0.2% |
MQA/GQA:针对自回归推理瓶颈的优化,主要减少KV Cache内存占用
MLA:针对序列长度瓶颈的优化,同时减少内存占用和计算复杂度,选择时需综合考虑模型性能、内存约束和序列长度需求。
随着Transformer技术的不断发展,这些注意力变体在各自适用的场景中发挥着重要作用,为不同规模的模型和应用提供了灵活的技术选择。在实际应用中,GQA因其良好的平衡性成为当前大语言模型的主流选择,而MLA在处理超长序列场景中展现出独特优势。
更多推荐


所有评论(0)