目录

Attention 的多种变体:从 MHA 到 MLA 的演进与 KV Cache 优化

随着大语言模型(LLM)的上下文长度不断增加,推理时的 KV Cache 显存占用 成为了限制模型吞吐量和推理速度的最大瓶颈。基于“卡内通信带宽 > 卡间通信带宽 > 机间通信带宽”的原理,过大的 KV Cache 会迫使模型进行跨设备通信,严重拖慢推理速度。

为了在保证模型效果的前提下尽可能减少 KV Cache 的大小,Attention 机制经历了一系列的演进。本文将带你梳理从 MHA (Multi-Head Attention)MQA (Multi-Query Attention)GQA (Group-Query Attention),再到 DeepSeek 提出的 MLA (Multi-head Latent Attention) 的发展历程及核心数学推导。

Attention 机制的演进:MHA, MQA, GQA
Attention 机制的演进:MHA, MQA, GQA

1. MHA (Multi-Head Attention)

MHA 源自 2017 年的经典论文《Attention is All You Need》。它的核心思想是将输入的 Query (Q)、Key (K)、Value (V) 在特征维度上分割为 $h$ 个独立的 Head,每个 Head 单独计算 Attention,最后将结果拼接(Concat)起来。

假设输入序列为 $\boldsymbol{x}_1, \boldsymbol{x}_2, \cdots, \boldsymbol{x}_t$。为了便于理解后续的数学表达,我们先统一一下符号定义:

  • $t$:当前正在生成的 token 时刻(当前位置)。
  • $i$:历史序列中的 token 索引($i = 1, 2, \dots, t$)。
  • $d$:模型输入的隐藏层维度(Hidden Dimension)。
  • $h$:注意力头(Head)的总数量。
  • $s$:当前计算的第 $s$ 个 Head($s \in [1, h]$)。
  • $d_k, d_v$:每个 Head 内部的 Key 和 Value 的维度,通常 $d_k = d_v = d / h$。
  • $\boldsymbol{W}$:各类可学习的线性映射权重矩阵。
  • 行向量约定:为了契合深度学习代码(如 PyTorch 中 [batch, seq_len, hidden_dim])的张量排布习惯,本文公式中的 $\boldsymbol{x}, \boldsymbol{q}, \boldsymbol{k}, \boldsymbol{v}, \boldsymbol{c}, \boldsymbol{o}$ 均为行向量(Row Vector)。例如 $\boldsymbol{x}_i \in \mathbb{R}^{1 \times d}$。
  • 输出标识 $\boldsymbol{o}_t$:Attention 层的输出记为 $\boldsymbol{o}_t$(Output),代表当前时刻 $t$ 融合了历史上下文后的特征表示。不使用 $\boldsymbol{x}_{t+1}$ 是因为 $\boldsymbol{x}_{t+1}$ 通常专指自回归生成中预测出的“下一个时刻的新输入 token”。

第 $s$ 个 head 上的线性映射与 Attention 计算如下:

$$ \begin{aligned} \boldsymbol{q}_i^{(s)} &= \boldsymbol{x}_i \boldsymbol{W}_q^{(s)} \in \mathbb{R}^{d_k}, & \boldsymbol{W}_q^{(s)} &\in \mathbb{R}^{d \times d_k} \\ \boldsymbol{k}_i^{(s)} &= \boldsymbol{x}_i \boldsymbol{W}_k^{(s)} \in \mathbb{R}^{d_k}, & \boldsymbol{W}_k^{(s)} &\in \mathbb{R}^{d \times d_k} \\ \boldsymbol{v}_i^{(s)} &= \boldsymbol{x}_i \boldsymbol{W}_v^{(s)} \in \mathbb{R}^{d_v}, & \boldsymbol{W}_v^{(s)} &\in \mathbb{R}^{d \times d_v} \end{aligned} $$

第 $t$ 个 Query 与历史 $1 \sim t$ 的所有 Key 和 Value 计算 Attention(公式中的 $\sum_{i \leq t}$ 即代表对当前及所有历史时刻 $i$ 进行求和):

$$ \boldsymbol{o}_t^{(s)} = \text{Attention} \left( \boldsymbol{q}_t^{(s)}, \boldsymbol{k}_{\leq t}^{(s)}, \boldsymbol{v}_{\leq t}^{(s)} \right) \triangleq \frac{\sum_{i \leq t} \exp \left( \boldsymbol{q}_t^{(s)} \boldsymbol{k}_i^{(s)\top} \right) \boldsymbol{v}_i^{(s)}}{\sum_{i \leq t} \exp \left( \boldsymbol{q}_t^{(s)} \boldsymbol{k}_i^{(s)\top} \right)} $$

最后将多个 head 的结果拼接,并通过一个输出权重矩阵 $\boldsymbol{W}^O$ 进行线性映射,得到最终的输出:

$$ \boldsymbol{o}_t = \text{Concat} \left( \boldsymbol{o}_t^{(1)}, \boldsymbol{o}_t^{(2)}, \cdots, \boldsymbol{o}_t^{(h)} \right) \boldsymbol{W}^O $$

在推理时,为了避免不必要的重复计算,我们引入了 KV Cache,将历史 token 中每个 Head 的 $\boldsymbol{k}_i^{(s)}$ 和 $\boldsymbol{v}_i^{(s)}$ 缓存起来。本质上是通过“空间换时间”的方式提升推理速度。

  • 优点:每个 Head 都有独立的 K 和 V,特征表达能力最强,模型效果上限最高。
  • 缺点:KV Cache 的大小与序列长度、Batch Size 以及 Head 的数量呈线性正相关。在处理长文本时,MHA 的 KV Cache 会占用极大的显存。当单卡显存无法容纳时,必须跨卡甚至跨机通信,严重拖慢推理速度。

2. MQA (Multi-Query Attention)

为了解决 MHA 显存占用过大的问题,2019 年的论文《Fast Transformer Decoding: One Write-Head is All You Need》提出了 MQA

  • 核心思想:所有的 Head 共享同一份 Key 和 Value,只有 Query 在不同 Head 之间是独立的。
  • 数学表达:可以看到 $\boldsymbol{k}_i$ 和 $\boldsymbol{v}_i$ 不再带有上标 $(s)$,即没有了 Head 的概念。
$$ \begin{aligned} \boldsymbol{q}_i^{(s)} &= \boldsymbol{x}_i \boldsymbol{W}_q^{(s)} \in \mathbb{R}^{d_k}, & \boldsymbol{W}_q^{(s)} &\in \mathbb{R}^{d \times d_k} \\ \boldsymbol{k}_i &= \boldsymbol{x}_i \boldsymbol{W}_k \in \mathbb{R}^{d_k}, & \boldsymbol{W}_k &\in \mathbb{R}^{d \times d_k} \\ \boldsymbol{v}_i &= \boldsymbol{x}_i \boldsymbol{W}_v \in \mathbb{R}^{d_v}, & \boldsymbol{W}_v &\in \mathbb{R}^{d \times d_v} \end{aligned} $$
  • 缓存变化:KV Cache 只需要缓存 1 个 Head 的 K 和 V,大小直接降为 MHA 的 $1/h$。
  • 优点:极大地减少了显存占用,提升了推理速度。
  • 缺点:过度压缩了 K 和 V 的表征空间,导致模型的效果(准确率)出现一定程度的下降。

3. GQA (Group-Query Attention)

为了在 MHA(效果好但显存大)和 MQA(显存小但效果差)之间取得平衡,2023 年的论文《GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints》提出了 GQA

  • 核心思想:将所有的 Head 分为 $g$ 组。同一组内的 Head 共享一份 K 和 V。
  • 数学表达:将每一份 K 和 V 均分为 $g$ 组。公式中的上标 $([sg/h])$ 表示第 $s$ 个 Head 被分配到了第 $\lfloor sg/h \rfloor$ 个 K-V 组中。每组的 KV 被重复(repeat)$h/g$ 次,正好满足 $h$ 个 Head 所需。
$$ \begin{aligned} \boldsymbol{q}_i^{(s)} &= \boldsymbol{x}_i \boldsymbol{W}_q^{(s)} \in \mathbb{R}^{d_k} \\ \boldsymbol{k}_i^{([sg/h])} &= \boldsymbol{x}_i \boldsymbol{W}_k^{([sg/h])} \in \mathbb{R}^{d_k} \\ \boldsymbol{v}_i^{([sg/h])} &= \boldsymbol{x}_i \boldsymbol{W}_v^{([sg/h])} \in \mathbb{R}^{d_v} \end{aligned} $$
  • 实现细节:当 $g=h$ 时,GQA 退化为 MHA;当 $g=1$ 时,GQA 退化为 MQA。在 LLaMA 2/3-70B 等模型中,通常设置 $g=8$。这样正好每张显卡可以负责计算一组 K 和 V 的 Attention,既保证了 KV 的多样性,又大幅减少了卡间通讯和显存占用。
  • 优点:在推理速度和模型效果之间取得了极佳的平衡。
  • 缺点:需要人为凭经验设定合理的分组数 $g$。

4. MLA (Multi-head Latent Attention)

MLA 是 DeepSeek-V2 中提出的一种极具创新的 Attention 变体。它通过 低秩投影 (Low-Rank Projection) 的方式,在大幅降低 KV Cache 显存占用的同时,取得了比 MHA 更好的效果。

Multi-head Latent Attention (MLA) 架构图
Multi-head Latent Attention (MLA) 架构图

4.1 核心思想:低秩映射与矩阵吸收

MLA 不再直接缓存高维的 K 和 V,而是将输入 $\boldsymbol{x}_i$ 通过一个低秩矩阵降维成一个隐向量(Latent Vector)$\boldsymbol{c}_i \in \mathbb{R}^{d_c}$(多个 head 之间共享):

$$ \boldsymbol{c}_i = \boldsymbol{x}_i \boldsymbol{W}_c \in \mathbb{R}^{d_c}, \quad \boldsymbol{W}_c \in \mathbb{R}^{d \times d_c} $$

然后通过升维矩阵 $\boldsymbol{W}_k^{(s)}$ 和 $\boldsymbol{W}_v^{(s)}$ 将 $\boldsymbol{c}_i$ 重新映射回 $\boldsymbol{k}_i^{(s)}$ 和 $\boldsymbol{v}_i^{(s)}$:

$$ \begin{aligned} \boldsymbol{q}_i^{(s)} &= \boldsymbol{x}_i \boldsymbol{W}_q^{(s)} \in \mathbb{R}^{d_k} \\ \boldsymbol{k}_i^{(s)} &= \boldsymbol{c}_i \boldsymbol{W}_k^{(s)} \in \mathbb{R}^{d_k}, & \boldsymbol{W}_k^{(s)} &\in \mathbb{R}^{d_c \times d_k} \\ \boldsymbol{v}_i^{(s)} &= \boldsymbol{c}_i \boldsymbol{W}_v^{(s)} \in \mathbb{R}^{d_v}, & \boldsymbol{W}_v^{(s)} &\in \mathbb{R}^{d_c \times d_v} \end{aligned} $$
  • 缓存变化:推理时,KV Cache 只需要缓存这个低维的 $\boldsymbol{c}_i$。传统的 KV Cache 大小为 $2 \times h \times d_k \times l$(2 代表 K 和 V,$h$ 为头数,$d_k$ 为头维度,$l$ 为层数),现在骤降为 $d_c \times l$(其中 $d_c \ll h \times d_k$)。
  • 矩阵吸收(Absorbed):你可能会问,推理时每次都要把 $\boldsymbol{c}_i$ 升维回 K 和 V,这不增加计算量吗?实际上,由于矩阵乘法的结合律,我们可以将升维矩阵与 Query 的投影矩阵提前合并(吸收)
$$ \boldsymbol{q}_t^{(s)} \boldsymbol{k}_i^{(s)\top} = \left( \boldsymbol{x}_t \boldsymbol{W}_q^{(s)} \right) \left( \boldsymbol{c}_i \boldsymbol{W}_k^{(s)} \right)^\top = \boldsymbol{x}_t \left( \boldsymbol{W}_q^{(s)} \boldsymbol{W}_k^{(s)\top} \right) \boldsymbol{c}_i^\top $$

这样,$\boldsymbol{W}_q^{(s)}$ 和 $\boldsymbol{W}_k^{(s)\top}$ 在推理前就可以合并为一个权重矩阵。同理,Value 的升维矩阵 $\boldsymbol{W}_v^{(s)}$ 也可以被吸收进最终的输出拼接矩阵 $\boldsymbol{W}^O$ 中。因此,$\boldsymbol{c}_i$ 就能直接参与 Attention 计算,完全不需要在推理时动态计算出 K 和 V。

(注:在 DeepSeek-V2 中,为了节约训练参数量,对 Q 也进行了低秩投影 $\boldsymbol{c}'_i = \boldsymbol{x}_i \boldsymbol{W}'_c$,但这与 KV Cache 优化无关。)

4.2 兼容 RoPE(旋转位置编码)的解耦设计

上述的矩阵吸收有一个前提:Q 和 K 的矩阵乘法中间不能有其他干扰。但是,现代大模型普遍使用 RoPE(旋转位置编码),RoPE 矩阵 $\boldsymbol{\mathcal{R}}$ 会插在 Q 和 K 之间:

$$ \boldsymbol{q}_t^{(s)} \boldsymbol{k}_i^{(s)\top} = \left( \boldsymbol{x}_t \boldsymbol{W}_q^{(s)} \boldsymbol{\mathcal{R}}_t \right) \left( \boldsymbol{c}_i \boldsymbol{W}_k^{(s)} \boldsymbol{\mathcal{R}}_i \right)^\top = \boldsymbol{x}_t \left( \boldsymbol{W}_q^{(s)} \boldsymbol{\mathcal{R}}_{t \rightarrow i} \boldsymbol{W}_k^{(s)\top} \right) \boldsymbol{c}_i^\top $$

由于矩阵相乘不满足交换律,中间夹着的相对位置旋转矩阵 $\boldsymbol{\mathcal{R}}_{t \rightarrow i}$(表示从时刻 $t$ 到时刻 $i$ 的相对位置变换)导致 $\boldsymbol{W}_q^{(s)}$ 和 $\boldsymbol{W}_k^{(s)\top}$ 无法提前合并。

为了解决这个问题,MLA 巧妙地将 Q 和 K 进行了解耦,拆分为“不带 RoPE 的部分”和“带 RoPE 的部分”:

$$ \begin{aligned} \boldsymbol{q}_i^{(s)} &= \left[ \boldsymbol{c}'_i \boldsymbol{W}_{qc}^{(s)}, \boldsymbol{c}'_i \boldsymbol{W}_{qr}^{(s)} \boldsymbol{\mathcal{R}}_i \right] \in \mathbb{R}^{d_k + d_r} \\ \boldsymbol{k}_i^{(s)} &= \left[ \boldsymbol{c}_i \boldsymbol{W}_{kc}^{(s)}, \boldsymbol{x}_i \boldsymbol{W}_{kr} \boldsymbol{\mathcal{R}}_i \right] \in \mathbb{R}^{d_k + d_r} \end{aligned} $$
  • Query:包含没有 RoPE 的多头 $\boldsymbol{c}'_i \boldsymbol{W}_{qc}^{(s)}$,和带有 RoPE 的多头 $\boldsymbol{c}'_i \boldsymbol{W}_{qr}^{(s)} \boldsymbol{\mathcal{R}}_i$。
  • Key:包含没有 RoPE 的多头 $\boldsymbol{c}_i \boldsymbol{W}_{kc}^{(s)}$(由 $\boldsymbol{c}_i$ 升维),和带有 RoPE 的多头共享部分 $\boldsymbol{x}_i \boldsymbol{W}_{kr} \boldsymbol{\mathcal{R}}_i$(直接由输入生成)。

此时,QK 的矩阵计算变为:

$$ \begin{aligned} \boldsymbol{q}_t^{(s)} \boldsymbol{k}_i^{(s)\top} &= \left[ \boldsymbol{c}'_t \boldsymbol{W}_{qc}^{(s)}, \boldsymbol{c}'_t \boldsymbol{W}_{qr}^{(s)} \boldsymbol{\mathcal{R}}_t \right] \left[ \boldsymbol{c}_i \boldsymbol{W}_{kc}^{(s)}, \boldsymbol{x}_i \boldsymbol{W}_{kr} \boldsymbol{\mathcal{R}}_i \right]^\top \\ &= \boldsymbol{c}'_t \left( \boldsymbol{W}_{qc}^{(s)} \boldsymbol{W}_{kc}^{(s)\top} \right) \boldsymbol{c}_i^\top + \left( \boldsymbol{c}'_t \boldsymbol{W}_{qr}^{(s)} \boldsymbol{\mathcal{R}}_t \right) \left( \boldsymbol{x}_i \boldsymbol{W}_{kr} \boldsymbol{\mathcal{R}}_i \right)^\top \end{aligned} $$

通过这种解耦,第一项中不带 RoPE 的权重矩阵 $\boldsymbol{W}_{qc}^{(s)} \boldsymbol{W}_{kc}^{(s)\top}$ 依然可以被完美吸收!

4.3 最终的 KV Cache 大小

在 DeepSeek-V2 的设定中($d_c = 4d_k$, $d_r = d_k / 2$),MLA 最终需要缓存的内容只有两部分:

  1. 降维后的隐向量 $\boldsymbol{c}_i \in \mathbb{R}^{d_c}$
  2. 共享的 RoPE 键向量 $\boldsymbol{k}_t^R = \boldsymbol{x}_i \boldsymbol{W}_{kr} \boldsymbol{\mathcal{R}}_i \in \mathbb{R}^{d_r}$

其 KV Cache 的总大小为 $(d_c + d_r) \times l$,约等于 2.25 组的 GQA

  • 优点:通过低秩投影和矩阵吸收,不仅极大地压缩了显存占用,还因为保留了更丰富的联合特征,其模型效果甚至超越了传统的 MHA。
  • 缺点:实现复杂度较高,且由于矩阵吸收的特性,在训练和推理阶段的代码逻辑差异较大。

总结

从 MHA 到 MLA,Attention 机制的演进清晰地展示了算法与工程(显存、带宽)之间的博弈与妥协:

  • MHA:性能上限高,但 KV Cache 爆炸。
  • MQA / GQA:通过共享 KV 强行压缩显存,但在特征多样性上做出了妥协。
  • MLA:通过低秩隐向量压缩信息,利用矩阵吸收消除计算开销,并巧妙解耦 RoPE,最终实现了“极小显存占用 + 极高模型性能”的完美既要又要。