← 返回题目列表

MHA、MQA、GQA、MLA 有什么区别?为什么大模型要改注意力的 KV?

高频 中等 第 7 / 25 题 更新于 2026/07/28
多头注意力MQAGQAMLAKV Cache

简化版

它们都是多头注意力的变体,核心区别在于多个 Query 头共享多少个 Key/Value 头,目的都是压缩推理时的 KV Cache、加快解码

  • MHA(标准多头):每个 Query 头配一套独立的 K/V 头,质量最好但 KV Cache 最大。
  • MQA(多查询):所有 Query 头共享同一套 K/V,KV Cache 缩小 h 倍(h=头数),解码最快,但质量略降。
  • GQA(分组查询):折中——把 Query 头分成 g 组,每组共享一套 K/V,g 介于 1(=MQA)和 h(=MHA)之间,兼顾质量与显存。Llama 2/3 用的就是它。
  • MLA(多头潜注意力,DeepSeek):把 K/V 压缩成一个低秩潜向量缓存,用时再解压,KV Cache 比 GQA 还小,质量还能接近 MHA。

详细版

为什么要动 KV

自回归解码时,为避免重复计算,历史 token 的 Key/Value 会被缓存下来(KV Cache)。它的大小 = 层数 × KV头数 × head_dim × 序列长度 × 2(K和V) × batch长上下文 + 大 batch 时,KV Cache 会占用巨量显存,成为推理瓶颈(甚至超过模型权重)。减少 KV 头数,就能成倍压缩缓存、提高吞吐。

四者对比

变体Query 头KV 头KV Cache 大小质量代表模型
MHAhh基准(最大)最好原始 Transformer、GPT-3
MQAh1缩小 h 倍略降PaLM、Falcon
GQAhg(1<g<h)缩小 h/g 倍接近 MHALlama 2 70B、Llama 3、Mistral
MLAh低秩潜向量比 GQA 更小接近/超过 MHADeepSeek-V2/V3

GQA 为什么成主流

MQA 省显存但质量掉得明显(尤其大模型、长上下文);MHA 质量好但显存吃紧。GQA 用「分组共享」在两者之间取一个甜点:比如 64 个 Query 头分成 8 组、共 8 套 KV,KV Cache 只有 MHA 的 1/8,质量却几乎无损。这就是当代开源大模型普遍选 GQA 的原因。

完整版教学

一、先复习多头注意力,锁定”KV 头”这个变量

标准多头注意力(MHA)把注意力拆成 h 个头,每个头有独立的投影矩阵,分别产生自己的 Q、K、V:

head_i = Attention(Q·W_Q^i, K·W_K^i, V·W_V^i)
MHA = Concat(head_1, ..., head_h) · W_O

关键点:MHA 里 Q、K、V 的头数都是 h,各头互相独立。推理时要缓存的是每个 K/V 头的历史,所以 KV Cache 正比于 KV 头的数量。MQA/GQA/MLA 动的全是这个「KV 头」维度——Query 头一般保持 h 不变(保留多头的表达多样性),只削减 KV 头。

二、KV Cache 为什么是推理的命门

自回归生成第 t 个 token 时,注意力需要用到前面所有 token 的 K 和 V。如果每步都重算历史的 K/V,复杂度是 O(t²) 且大量重复。KV Cache 的做法是把历史 K/V 存起来,每步只算新 token 的 K/V 并追加,把每步降到 O(t)。

代价是显存。KV Cache 大小:

KVCache = 2 × 层数 × KV头数 × head_dim × 序列长度 × batch × 精度字节

举例:一个 80 层、64 头、head_dim=128 的模型,处理 32K 上下文、batch=16,用 MHA 的 KV Cache 能轻松到几十上百 GB,远超模型权重本身。长上下文和高并发场景下,KV Cache 直接决定了能塞多大 batch、多长上下文——这就是为什么削减 KV 头如此值得。

记住这条因果链:削减 KV 头数 → KV Cache 变小 → 能放更大 batch / 更长上下文 → 吞吐和成本大幅改善。这是 MQA/GQA/MLA 的共同动机。

三、MQA:一套 KV 走天下

MQA 把所有 Query 头共享同一套 K、V(KV 头数=1):

head_i = Attention(Q·W_Q^i, K·W_K, V·W_V)   # W_K, W_V 全体共享
  • 好处:KV Cache 缩小 h 倍,解码时访存大减,速度提升明显。
  • 代价:所有头看同一份 K/V,注意力的多样性下降,质量有可损,尤其在大模型和长上下文上更明显。
  • 适合对延迟极敏感、模型不太大的场景。

四、GQA:分组是关键的折中艺术

GQA 把 h 个 Query 头分成 g 组,每组内的 Query 头共享一套 K/V,共 g 套 KV:

  • g=1:所有头共享 → 退化为 MQA
  • g=h:每头独立 → 退化为 MHA
  • 1<g<h:折中,KV Cache 是 MHA 的 g/h。

典型配置:64 个 Query 头分 8 组,KV 头=8,缓存降到 1/8,而质量几乎追平 MHA。为什么分组就能保住质量? 因为「组内共享、组间独立」保留了相当程度的注意力多样性——不同组仍能关注不同信息,不像 MQA 那样一刀切。

工程上还有个便利:MHA 训练好的模型可以通过对 KV 头做均值池化”升级”成 GQA 再微调(uptraining),不必从头训,Llama 2 论文就是这么做的。

五、MLA:换个思路,压缩而非共享

DeepSeek 的 MLA(Multi-head Latent Attention)不走「减少头数」的路,而是低秩压缩

  • 把每个 token 的 K/V 通过一个下投影压成一个低维潜向量(latent),只缓存这个潜向量,而不是完整的多头 K/V。
  • 用时再用上投影把潜向量解压回各头的 K/V。
  • 为兼容 RoPE(旋转位置编码不能简单低秩分解),MLA 采用解耦 RoPE:一部分维度专门承载位置信息、单独处理。

效果:KV Cache 比 GQA 还小(相当于只存一个压缩向量),质量却能接近甚至超过 MHA。这让 DeepSeek 在长上下文推理上有显著的显存和吞吐优势。理解 MLA 抓一句:「缓存一个低秩潜向量,而不是缓存整套多头 KV」

六、面试拆解算例

这类题最怕只讲术语,最好把它落成一次资源账。假设模型有 32 层、hidden size 为 4096、序列长度从 2K 增到 16K,标准注意力的相关度矩阵规模会从 2K × 2K 变成 16K × 16K,理论元素数量放大 64 倍。即使具体算子不会真的把所有中间矩阵都落到 HBM,复杂度曲线仍然决定了 prefill 延迟和显存压力会快速上升。

attention_scores_per_head = seq_len * seq_len
2K  ->  2,048 * 2,048   ≈ 4.19M
16K -> 16,384 * 16,384  ≈ 268.44M
放大倍数 ≈ 64
观察维度面试要说清的问题工程判断
张量形状Q/K/V、hidden state 或路由权重怎样变化能否解释实现差异
复杂度随层数、序列长度、head 数怎样增长谁先成为瓶颈
质量风险是否改变训练分布或表达能力会不会掉点
部署代价算子、框架、缓存是否支持能不能稳定上线
输入 token -> embedding -> 注意力/FFN/归一化模块 -> hidden state -> logits
                |             |                 |
             位置/掩码       显存与吞吐        质量与稳定性

所以回答「MHA、MQA、GQA、MLA 有什么区别?为什么大模型要改注意力的 KV?」时,推荐先讲结构变化,再讲这条变化怎样影响资源曲线,最后补一句质量和部署的边界。这样比单纯背「某某结构更快、更省」更像工程答案。

七、常见误区与追问

  • 误区:MQA/GQA 减少的是 Query 头。 减的是 KV 头,Query 头通常保持 h 不变。
  • 误区:改这些是为了省训练算力。 主要为省推理显存(KV Cache)和加快解码,训练算力变化不大。
  • 追问:GQA 的 g 怎么选? 权衡质量与显存,常见取 8;g 越小越省显存但质量风险越大。
  • 追问:为什么不干脆都用 MQA? 大模型/长上下文下 MQA 质量掉得明显,GQA 才是稳妥折中。
  • 追问:MLA 和 GQA 谁更好? MLA 在同等 KV Cache 下质量更高(或同等质量下缓存更小),但实现更复杂(需解耦 RoPE);GQA 简单通用、生态成熟。
  • 追问:这些和 FlashAttention 冲突吗? 不冲突,正交。FlashAttention 优化的是注意力的计算/访存方式,MQA/GQA/MLA 优化的是 KV 的存储量,二者叠加使用。

八、加强记忆

一条主线钉死四者:它们都在削 KV、保 Query,目的是压缩推理时的 KV Cache。 MHA=每个 Query 头一套 KV(质量最好、缓存最大);MQA=全体共享一套 KV(最省、质量略降);GQA=分 g 组各享一套 KV(折中甜点,Llama 系首选);MLA=把 KV 压成低秩潜向量缓存(DeepSeek,缓存更小质量还高)。记住因果:KV 头越少 → 缓存越小 → batch 更大、上下文更长 → 吞吐更高,而分组/压缩是为了在省显存的同时守住质量。