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}}}XRn×dmodel
  • 头数:HHH
  • 每头维度:dk=dmodel/Hd_k = d_{\text{model}}/Hdk=dmodel/H
  • K、V带独立的下标h,独立的头空间
  1. 线性映射(每头独立)
    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,WhQRdmodel×dk=XWhK,WhKRdmodel×dk=XWhV,WhVRdmodel×dk

  2. 缩放点积注意力(每头)
    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(dk QhKh)Vh

  3. 拼接 + 输出映射
    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,WORHdk×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,WhKRdmodel×dk=XWhV,WhVRdmodel×dk

因为每个Attention的头,都需要有独立的K、V进行计算,其消耗的内存会比较多。


Multi-Query Attention(MQA)

MQA是MHA的一种极端简化形式,旨在显著减少KV Cache的大小。
在这里插入图片描述

核心区别:所有头共享同一组K\mathbf{K}KV\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}}}XRn×dmodel
  • 头数:HHH
  • 每头维度:dk=dmodel/Hd_k = d_{\text{model}}/Hdk=dmodel/H
  • K,V\mathbf{K},\mathbf{V}K,V不再带下标hhh所有头共用
  1. 线性映射(共享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,WhQRdmodel×dk=XWK,WKRdmodel×dk=XWV,WVRdmodel×dk

  2. 缩放点积注意力(每头)
    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(dk QhK)V

  3. 拼接 + 输出映射
    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,WORHdk×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,WKRdmodel×dk=XWV,WVRdmodel×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}}}XRn×dmodel
  • 头数:HHH
  • 分组数:GGG
  • 每头维度:dk=dmodel/Hd_k = d_{\text{model}}/Hdk=dmodel/H
  1. 线性映射(分组共享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,WgKRdmodel×dk=XWgV,WgVRdmodel×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,WhQRdmodel×dk

  2. 缩放点积注意力(组内共享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(dk QhKg)Vg

  3. 拼接 + 输出映射
    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,WORHdk×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}}dlatdmodel
  • 再按头展开,保证多头查询的独立性
  • 因此K,V\mathbf{K},\mathbf{V}K,V不再直接参与注意力计算,而是先通过低秩的方式进行压缩再映射

计算步骤
输入

  • 输入序列:X∈Rn×dmodel\mathbf{X}\in\mathbb{R}^{n\times d_{\text{model}}}XRn×dmodel
  • 头数:HHH
  • 每头维度:dk=dmodel/Hd_k = d_{\text{model}}/Hdk=dmodel/H
  • 潜在维度:dlatd_{\text{lat}}dlat(通常dlat≪dmodeld_{\text{lat}}\ll d_{\text{model}}dlatdmodel
  1. 低秩压缩(共享)
    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,CkRdmodel×dlat=XCv,CvRdmodel×dlat

C_kC_v分别是键(Key)的低秩压缩矩阵值(Value)的低秩压缩矩阵,其目的在于将原始高维输入X投影到一个低维的潜在空间(latent space),维度从d_model(例如4096)压缩到d_lat(例如128)。这使得后续生成的键缓存(KV Cache)大小从O(n * d_model)显著减少到O(n * d_lat),从而极大地节省了推理时的内存占用。

  1. 按头映射
    对每个头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,WhQRdmodel×dk=ZkUhk,UhkRdlat×dk=ZvUhv,UhvRdlat×dk

    其中:

    • UhK,Uhv∈Rdlat×dkU^{K}_h, U^{v}_h \in \mathbb{R}^{d_{lat} \times d_k}UhK,UhvRdlat×dk:每头K与V的上投影(Up)
  2. 缩放点积注意力(每头)
    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(dk Qhkh)vh

  3. 拼接 + 输出映射
    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,WORHdk×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}}dlatdmodel时,内存占用接近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 dimdlatdim

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在处理超长序列场景中展现出独特优势。

Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐