跳转到主要内容

Attention 架构演化:从 MHA 到 GQA、MLA 与混合注意力

从缩放点积注意力出发,解释 MHA 的多头投影与输出融合,并用 KV 共享、表示压缩、访问稀疏和递推状态四条路线理解现代 LLM Attention。

· 约 6 分钟阅读

现代 LLM 并没有简单地抛弃 MHA。演化的核心是:保留多视角信息选择能力,同时减少 KV Cache、历史扫描和显存读写。

Attention 从 MHA 到混合架构的演化地图

1. 起点:缩放点积注意力

单头缩放点积注意力接收一组 Query、Key 和 Value:

Attention(Q,K,V)=softmax ⁣(QKdk)V\operatorname{Attention}(Q,K,V) = \operatorname{softmax}\!\left(\frac{QK^\top}{\sqrt{d_k}}\right)V
  • QQ 表示当前 token 在找什么。
  • KK 表示每个 token 可以通过什么特征被匹配。
  • VV 表示匹配后真正取回的内容。
  • QKQK^\top 产生 token 两两之间的匹配分数。

除以 dk\sqrt{d_k} 是为了控制高维点积的尺度,避免 softmax 过早饱和。

2. MHA:并行学习多个表示子空间

MHA 使用多组独立、可学习的线性投影:

Qi=XWiQ,Ki=XWiK,Vi=XWiVQ_i=XW_i^Q,\qquad K_i=XW_i^K,\qquad V_i=XW_i^V

ii 个 head 独立计算:

headi=Attention(Qi,Ki,Vi)\operatorname{head}_i =\operatorname{Attention}(Q_i,K_i,V_i)

最后拼接并经过输出投影:

MHA(X)=Concat(head1,,headh)WO\operatorname{MHA}(X) =\operatorname{Concat}(\operatorname{head}_1,\ldots,\operatorname{head}_h)W^O

这里有三个容易混淆的点:

  1. **“投影多次”不是重复迭代。**它表示同一个输入同时乘多组不同参数。
  2. **每个 head 的参数在训练中学得。**训练完成后,它们保存在模型权重中,推理时固定使用。
  3. **WOW^O 不是注意力权重。**它把各 head 的输出重新混合并映射回 dmodeld_{\text{model}},方便后续层和残差连接继续处理。

工程实现通常把各 head 的小矩阵合并成大矩阵,一次 GEMM 生成全部 Q/K/V,再 reshape 出 head 维度。因此“逻辑上多组参数”和“物理上多次小矩阵乘法”不是一回事。

3. 为什么要继续改造 MHA

MHA 的主要成本分布在两个阶段:

阶段主要成本随上下文增长
Prefill大量 Query 与 Key 的两两交互完整 attention 约为 O(S2)O(S^2)
Decode每步读取历史 K/V每步读取量约为 O(S)O(S)
常驻显存保存每层、每 token、每 KV head 的 K/VKV Cache 约为 O(S)O(S)

因此后续架构主要沿四条路线发展:

  • 少存几份 KV:MQA、GQA。
  • 把 KV 表示压小:MLA。
  • 少访问一些历史 token:局部、滑窗、块稀疏注意力。
  • 不用逐 token KV 表保存历史:线性/递推式注意力。

4. MQA 与 GQA:减少独立 KV heads

设 Query heads 数为 hh,KV heads 数为 gg

架构Query headsKV heads共享方式
MHAhhhh每个 Q head 有独立 K/V
GQAhhgg,且 1<g<h1<g<h一组 Q heads 共享一组 K/V
MQAhh11所有 Q heads 共享一组 K/V

它们的核心区别不是 Query 数量,而是需要缓存多少组 K/V:

KV elements/token/layer=2×nkv-heads×dhead\text{KV elements/token/layer} =2\times n_{\text{kv-heads}}\times d_{\text{head}}

MQA 最省 KV Cache,但共享约束最强;GQA 位于 MHA 与 MQA 之间,因而成为常见的效果/效率折中。详细显存公式见 KV Cache

5. MLA:压缩表示,而不只是减少 head

MLA 不把自己描述成“只有一个 KV head”。它先把 hidden state 压到共享 latent:

ctKV=WDKVhtc_t^{KV}=W_{DKV}h_t

再从 latent 形成各 head 使用的内容表示:

ktC=WUKctKV,vt=WUVctKVk_t^C=W_{UK}c_t^{KV},\qquad v_t=W_{UV}c_t^{KV}

推理时主要缓存低维 ctKVc_t^{KV} 和小型 RoPE 分支,而不是完整多头 K/V。矩阵吸收让 decode 可以直接在 latent 空间完成关键计算,避免每步显式恢复全部历史 K/V。

所以两条路线的作用点不同:

GQA / MQA:减少独立 KV 的份数
MLA:      压缩 KV 共同依赖的表示

完整机制见 DeepSeek MLA

6. 局部与稀疏注意力:减少“看谁”

MHA、GQA 和 MLA 都可以保留当前 Query 对完整历史做 softmax 的访问语义,但完整历史扫描在超长上下文中仍然昂贵。

局部或稀疏注意力改变访问范围:

  • Sliding Window 只看最近 WW 个 token。
  • Local/Global Hybrid 让多数层看局部、少数层看全局。
  • Block Sparse / Routed Attention 根据 Query 选择少量历史块。

它们减少计算和访存,但选择机制可能漏掉远距离关键信息。这里优化的是访问范围,不是 KV 表示本身。

7. 线性/递推式注意力:用状态概括历史

线性或递推式结构不再要求每一步读取完整历史 KV,而是维护状态:

St=Update(St1,kt,vt),yt=Read(qt,St)S_t=\operatorname{Update}(S_{t-1},k_t,v_t), \qquad y_t=\operatorname{Read}(q_t,S_t)

这样 decode 主要读写固定大小的 recurrent state,长上下文成本不再随完整 KV 序列同样增长。代价是有限状态可能压缩或遗忘细节,prefill 也需要 scan、chunk 或三角矩阵等并行化技巧。

GDN 与 Chunked Prefill 展示了这种路线在真实 kernel 和 trace 中的工程代价。

8. “精确”要分成两个问题

技术是否与完整 MHA 数学等价变化在哪里
FlashAttention是,忽略浮点舍入顺序差异改变 tiling、online softmax 和 IO,不改变目标公式
MQA / GQA改变 K/V 的参数化和共享约束
MLA否;但可保留完整历史 softmax 访问语义用低秩 latent 约束 KV 表示
局部/稀疏注意力一部分 token 对根本不参与计算
线性/递推式注意力通常否用不同算子或有限状态替代完整 softmax attention

“不与 MHA 数学等价”也不等于“最终模型一定更差”。这些结构通常从训练开始就适应自身约束;节省下来的显存和算力还可能用于更长上下文、更大 batch 或更大模型。

9. 当前收敛方向:混合注意力

单一路线各有短板,因此现代设计逐渐走向混合:

多数层:局部、稀疏或递推式注意力,承担低成本历史建模
少数层:完整 GQA/MLA,承担精确的远距离信息交换
系统层:KV 量化、分页、offload 和 cache-aware routing

这意味着 MHA 的核心思想——多个 Query 视角并行选择信息——仍在延续;真正被持续改造的是:K/V 如何表示、保存多少、每次访问多少,以及哪些 token 值得使用更昂贵的精确路径。

相关页面

修改历史1 次提交