返回博客

/ LLM算法

[LLM算法-3] 标准 MHA:从多头计算到 KV Cache

从张量形状和完整公式出发,拆解标准 Multi-Head Attention 的 Q、K、V 投影、多头计算与输出拼接,并分析 Prefill、Decode、KV Cache,以及 MHA、GQA、MQA 的核心差异。

14 minLLM · Transformer · Attention · MHA

Multi-Head Attention(MHA,多头注意力)是 Transformer 中最标准的 Attention 结构。它把隐藏空间拆成多个较小的子空间,让各个 Attention Head 独立建立 token 之间的关系,再把结果拼接起来。

一句话概括:

标准 MHA 中,Query、Key、Value 的头数相同,每个 Query Head 都有与自己一一对应的 Key Head 和 Value Head。

如果模型有 32 个 Attention Heads,那么它会生成:

32 个 Query Heads
32 个 Key Heads
32 个 Value Heads

对应关系是:

Q head 0  <-> K head 0,  V head 0
Q head 1  <-> K head 1,  V head 1
Q head 2  <-> K head 2,  V head 2
...
Q head 31 <-> K head 31, V head 31

因此标准 MHA 满足:

Hq=Hk=Hv,H_q=H_k=H_v,

也常简写为:

Hq=Hkv.H_q=H_{kv}.

这里的“一一对应”描述的是 Attention 计算时的配对关系,并不表示 Q、K、V 使用相同的参数。它们来自不同的线性投影,参数彼此独立。

1. 从输入和 Head Dimension 开始#

以 self-attention 为例,输入 hidden states 记为:

XRB×S×dmodel,X\in\mathbb R^{B\times S\times d_{\text{model}}},

其中:

  • BB:batch size;
  • SS:序列长度;
  • dmodeld_{\text{model}}:模型隐藏维度。

假设模型配置为:

batch_size     = 8
sequence_length = 2048
hidden_size    = 4096
num_heads      = 32

每个 head 的维度为:

dh=dmodelH=409632=128.d_h=\frac{d_{\text{model}}}{H} =\frac{4096}{32} =128.

这要求 dmodeld_{\text{model}} 能被 HH 整除。拆分 head 并不会改变总特征维度,因为:

Hdh=dmodel.H d_h=d_{\text{model}}.

2. 从 Hidden States 生成 Q、K、V#

输入 XX 分别经过三组线性投影:

Q=XWQ,K=XWK,V=XWV.Q=XW_Q,\qquad K=XW_K,\qquad V=XW_V.

在标准 MHA 中,若忽略 bias:

WQ,WK,WVRdmodel×dmodel.W_Q,W_K,W_V \in\mathbb R^{d_{\text{model}}\times d_{\text{model}}}.

所以投影后的张量形状仍然是:

Q,K,VRB×S×dmodel.Q,K,V\in\mathbb R^{B\times S\times d_{\text{model}}}.

接下来把最后一个维度拆成 HH 个 head,并把 head 维移动到序列维之前:

[B, S, d_model]
  -> [B, S, H, d_h]
  -> [B, H, S, d_h]

在前面的例子中:

[B, S, 4096] -> [B, 32, S, 128]

于是:

Q,K,VRB×H×S×dh.Q,K,V\in\mathbb R^{B\times H\times S\times d_h}.

工程实现常把三次投影融合成一次大的 QKV projection,以减少 kernel launch 和内存访问;这只改变执行方式,不改变数学结构。

3. 每个 Head 如何计算 Attention#

hh 个 Attention Head 只使用与它对应的 QhQ_hKhK_hVhV_h

Oh=softmax(QhKhdh+M)Vh,O_h= \operatorname{softmax} \left( \frac{Q_hK_h^\top}{\sqrt{d_h}}+M \right)V_h,

其中 MM 是 attention mask。这个过程可以拆成三步。

3.1 计算注意力分数#

Ah=QhKh.A_h=Q_hK_h^\top.

若:

Qh,KhRB×S×dh,Q_h,K_h\in\mathbb R^{B\times S\times d_h},

那么:

AhRB×S×S.A_h\in\mathbb R^{B\times S\times S}.

Ah[i,j]A_h[i,j] 表示在第 hh 个表示子空间里,第 ii 个 query token 与第 jj 个 key token 的匹配分数。它还不是概率,因为其中可能包含负值,且每行之和不等于 1。

3.2 缩放、Mask 与 Softmax#

Ph=softmax(Ahdh+M).P_h= \operatorname{softmax} \left( \frac{A_h}{\sqrt{d_h}}+M \right).

为什么要除以 dh\sqrt{d_h}

如果 Q、K 各维近似独立且方差为 1,那么点积 qkq^\top k 的方差会随 dhd_h 增长。较大的 logits 容易让 Softmax 过早进入饱和区,产生接近 one-hot 的分布和很小的梯度。除以 dh\sqrt{d_h} 可以把 logits 的尺度稳定在更合适的范围。

Mask 通常在 Softmax 之前加入。对不允许访问的位置,工程实现会加入负无穷或一个足够大的负数,使其 Softmax 权重接近 0。

在 Causal Language Model 中,因果 Mask 保证第 ii 个 token 不能看到未来位置 j>ij>i

Mij={0,ji,,j>i.M_{ij}= \begin{cases} 0,&j\le i,\\ -\infty,&j>i. \end{cases}

3.3 对 Value 加权求和#

Oh=PhVh.O_h=P_hV_h.

形状变化为:

[B, S, S] @ [B, S, d_h] -> [B, S, d_h]

因此:

OhRB×S×dh.O_h\in\mathbb R^{B\times S\times d_h}.

注意力权重由 Q 和 K 决定,真正被聚合到输出中的内容来自 V。可以用“Q 发起查询、K 提供地址、V 携带内容”来形成直觉,但三者本质上都是模型学习到的向量表示。

4. 拼接所有 Head 并做输出投影#

所有 head 独立计算后,先沿 head 对应的特征维拼接:

Oconcat=Concat(O1,,OH).O_{\text{concat}} =\operatorname{Concat}(O_1,\ldots,O_H).

形状从:

H 个 [B, S, d_h]

恢复为:

[B, S, H * d_h] = [B, S, d_model]

即:

OconcatRB×S×dmodel.O_{\text{concat}} \in\mathbb R^{B\times S\times d_{\text{model}}}.

最后经过输出投影:

Y=OconcatWO,Y=O_{\text{concat}}W_O,

其中:

WORdmodel×dmodel,YRB×S×dmodel.W_O\in\mathbb R^{d_{\text{model}}\times d_{\text{model}}}, \qquad Y\in\mathbb R^{B\times S\times d_{\text{model}}}.

WOW_O 不只是做形状变换。它会重新混合各个 head 产生的特征,让后续层能够联合使用不同 head 提取的信息。

5. 标准 MHA 的完整公式#

一般形式可以写成:

MultiHead(Q,K,V)=Concat(head1,,headH)WO,\operatorname{MultiHead}(Q,K,V) =\operatorname{Concat}(\operatorname{head}_1,\ldots,\operatorname{head}_H)W_O,

其中:

headh=Attention(QWQ(h),KWK(h),VWV(h)),\operatorname{head}_h =\operatorname{Attention} \left( QW_Q^{(h)}, KW_K^{(h)}, VW_V^{(h)} \right),

而 Scaled Dot-Product Attention 是:

Attention(Q,K,V)=softmax(QKdh+M)V.\operatorname{Attention}(Q,K,V) =\operatorname{softmax} \left( \frac{QK^\top}{\sqrt{d_h}}+M \right)V.

在 self-attention 中,公式里的 Q、K、V 来自同一组输入 hidden states;在 cross-attention 中,Query 可以来自 decoder,而 Key 和 Value 来自 encoder。MHA 描述的是多头计算与投影方式,并不只限于 self-attention。

6. 为什么需要多个 Head#

如果只使用一个维度为 dmodeld_{\text{model}} 的大 head,模型只会产生一组注意力分布。多头结构则允许模型在不同的低维表示子空间里,同时建立多组 token 关系。

从直觉上看,不同 head 可能更偏向:

局部相邻关系
句法依赖
实体指代
长距离关联
段落或分隔符结构

不过,这只是帮助理解的观察方式,并不意味着每个 head 都一定具有稳定、唯一、可由人命名的语义。

从参数角度看,每个 head 都对应独立的投影切片:

WQ(h),WK(h),WV(h).W_Q^{(h)},\quad W_K^{(h)},\quad W_V^{(h)}.

因此,各个 head 可以学习不同的 query-key 匹配规则和 value 表示。多头的主要价值不是增加总隐藏维度,而是让模型并行学习多组不同的交互模式。

7. 自回归推理为什么需要 KV Cache#

训练或 Prefill 时,模型通常可以并行处理一段 token;但在 Decode 阶段,自回归模型每一步只生成一个新 token。

假设当前已经有 TT 个可见 token。新 token 的 Query 形状为:

QnewRB×H×1×dh.Q_{\text{new}} \in\mathbb R^{B\times H\times 1\times d_h}.

历史 Key 和 Value 的形状为:

Kcache,VcacheRB×H×T×dh.K_{\text{cache}},V_{\text{cache}} \in\mathbb R^{B\times H\times T\times d_h}.

新 Query 会查询所有历史位置:

Onew=softmax(QnewKcachedh)Vcache.O_{\text{new}} =\operatorname{softmax} \left( \frac{Q_{\text{new}}K_{\text{cache}}^\top}{\sqrt{d_h}} \right)V_{\text{cache}}.

新 token 的 K 和 V 在生成后被追加到 cache,供下一步使用。历史 token 的 K/V 只依赖已经确定的 hidden states,不会在后续步骤变化,因此没有必要反复计算。这就是 KV Cache 的基本原理。

对单步 Decode 来说,当前 token 已经位于序列末尾,没有未来位置可见,所以通常不需要显式构造一个完整的三角 causal mask;但实现仍需正确处理 padding、滑动窗口或分块等边界。

8. MHA 的 KV Cache 有多大#

忽略对齐、分页和元数据开销,对单层模型而言,标准 MHA 的 KV Cache 大小近似为:

2×B×S×Hkv×dh×b,2\times B\times S\times H_{kv}\times d_h\times b,

其中:

  • 前面的 2 表示一份 K Cache 和一份 V Cache;
  • BB 是 batch size 或同时缓存的序列数;
  • SS 是缓存长度;
  • HkvH_{kv} 是 KV Heads 数量;
  • dhd_h 是每个 head 的维度;
  • bb 是每个元素的字节数。

对标准 MHA,Hkv=HH_{kv}=HHdh=dmodelHd_h=d_{\text{model}},所以单层也可以写成:

2×B×S×dmodel×b.2\times B\times S\times d_{\text{model}}\times b.

如果模型共有 LL 层,总量近似为:

KV Cache Bytes=2LBSHkvdhb.\boxed{ \text{KV Cache Bytes} =2LBSH_{kv}d_hb }.

8.1 一个 16 GiB 的具体例子#

假设模型配置为:

layers          = 32
batch           = 1
sequence_length = 32768
num_kv_heads    = 32
head_dim        = 128
dtype           = FP16 / BF16, 2 bytes

那么:

KV Cache=2×32×1×32768×32×128×2=17,179,869,184 bytes=16 GiB.\begin{aligned} \text{KV Cache} &=2\times32\times1\times32768\times32\times128\times2\\ &=17{,}179{,}869{,}184\ \text{bytes}\\ &=16\ \text{GiB}. \end{aligned}

这还只是 batch 为 1 时的理论数据大小。若 batch 增加到 8,在所有请求都缓存 32K token 的前提下:

KV Cache ~= 128 GiB

实际推理引擎还可能有 block 对齐、未填满页面、内存碎片和管理元数据等额外开销。反过来,KV Cache 量化、滑动窗口 Attention 或前缀共享也可能降低实际物理占用。

9. MHA 的计算复杂度#

对于长度为 SS 的 self-attention,每个 head 的 QK 计算是:

[S,dh]×[dh,S][S,S].[S,d_h]\times[d_h,S]\rightarrow[S,S].

所有 head 的复杂度约为:

O(BHS2dh)=O(BS2dmodel).O(BHS^2d_h) =O(BS^2d_{\text{model}}).

Attention probabilities 与 V 相乘的复杂度处于相同量级。因此,若只看 Attention 主体,它对序列长度呈二次增长:

O(S2).O(S^2).

不过,完整 Transformer 层还包含 QKV projection、输出投影和 MLP。线性投影的成本通常是 O(BSdmodel2)O(BSd_{\text{model}}^2),所以在上下文较短、隐藏维度很大时,Attention 不一定占据全部计算;随着 SS 增大,二次项才会越来越突出。

FlashAttention 可以通过 tiling 和在线 Softmax 避免把完整 S×SS\times S 注意力矩阵写回显存,显著降低中间显存占用与 IO,但它并没有把标准 Softmax Attention 的算术复杂度从 O(S2)O(S^2) 变成 O(S)O(S)

10. Prefill 与 Decode 是两种不同的瓶颈#

同一套 MHA 数学结构,在 Prefill 和 Decode 阶段会呈现完全不同的硬件特征。

10.1 Prefill:大矩阵,通常更偏计算密集#

假设一次输入 16K prompt:

q_len  = 16384
kv_len = 16384

每个 head 的逻辑 Attention 矩阵大小为:

16384×16384.16384\times16384.

FlashAttention 不会将这个矩阵完整保存在 HBM 中,但大部分 query-key 交互仍然需要计算。Prefill 中矩阵较大、数据复用机会较多,通常更容易利用设备的矩阵计算吞吐,因此往往更偏 compute-bound。

10.2 Decode:小 Query,长 KV,通常更偏带宽受限#

单步 Decode 时:

q_len  = 1
kv_len = 16384

每个 head 的核心乘法近似为:

[1, d_h] @ [d_h, 16384]

单个新 token 的算术量比 Prefill 小得多,但每一层、每一步都要读取整段历史 K/V。计算对数据的复用率较低,因此 Decode Attention 经常是 memory-bandwidth bound。

标准 MHA 的 KV Heads 最多,所以它不仅缓存量最大,Decode 时需要读取的历史 KV 数据量通常也最大。这正是 GQA 和 MQA 被广泛采用的重要原因。

11. MHA、GQA 与 MQA 的根本区别#

三种结构都可以拥有相同数量的 Query Heads,核心差异是有多少组独立的 Key/Value Heads,以及多少个 Query Heads 共享一组 KV。

假设 Hq=32H_q=32

结构Q HeadsKV Heads共享方式相对 MHA 的理论 KV Cache
MHA3232每个 Q Head 独享一组 KV11
GQA328每 4 个 Q Heads 共享一组 KV1/41/4
MQA321所有 Q Heads 共享一组 KV1/321/32

它们的连接关系可以画成:

MHA
Q0 -> KV0
Q1 -> KV1
Q2 -> KV2
Q3 -> KV3
 
GQA
Q0 --+
Q1 --+--> KV0
Q2 --+
Q3 --+
 
MQA
Q0 --+
Q1 --+
Q2 --+--> KV0
Q3 --+
...   |

更一般地,若 HqH_q 能被 HkvH_{kv} 整除,每组 KV 服务的 Query Heads 数量是:

G=HqHkv.G=\frac{H_q}{H_{kv}}.

对同样的层数、序列长度、head dimension 和数据类型,KV Cache 相对标准 MHA 的缩放比例近似为:

HkvHq.\frac{H_{kv}}{H_q}.

减少 KV Heads 主要降低的是 K/V projection 参数、KV Cache 容量和 Decode 读取量;Query Heads 仍然可以保持较多,以尽量保留多组查询表示。代价是多个 Query Heads 必须共享同一组 K/V 表示,可能形成表达瓶颈。

12. 标准 MHA 的优势#

12.1 每个 Head 拥有独立的 K/V 表示#

MHA 不要求多个 Query Heads 共享 K/V,每个 head 都能学习自己的匹配空间和被聚合内容,表达约束最少。

12.2 结构对称#

Hq=Hk=Hv.H_q=H_k=H_v.

Q、K、V 的头数和形状完全对称,理解、调试和通用 kernel 实现都比较直接。

12.3 训练与算子生态成熟#

原始 Transformer 就采用标准 MHA。它拥有成熟的训练方法、数值经验和高度优化的 Attention kernel,也是理解其他 Attention 变体的基准。

12.4 Tensor Parallel 切分自然#

例如:

num_heads = 32
TP size   = 8

在常见的 head-wise 切分中,每张设备可以分到:

4 个 Q Heads
4 个 K Heads
4 个 V Heads

Q/K/V 完全对称。不过,“容易切分”仍依赖头数能否整除 TP size、具体并行布局和通信实现,并不是任何配置下都自动成立。

13. 标准 MHA 的局限#

13.1 KV Cache 最大#

由于:

Hkv=Hq,H_{kv}=H_q,

在 Query Heads 和 head dimension 相同的前提下,MHA 的 KV Cache 比 GQA 和 MQA 更大。

13.2 Decode 带宽压力大#

每生成一个新 token,都要读取所有层、所有历史位置、所有 KV Heads 的 K/V。上下文越长、并发越高,这部分带宽成本越突出。

13.3 同等显存下可容纳的并发更低#

KV Cache 往往是服务系统中重要的动态显存消费者。单请求缓存越大,同样的显存容量可以同时驻留的 token 和请求通常越少。

13.4 长上下文成本高#

标准全局 MHA 同时面对两种增长:

KV Cache=O(S),\text{KV Cache}=O(S), Prefill Attention FLOPs=O(S2).\text{Prefill Attention FLOPs}=O(S^2).

前者带来容量和 Decode 带宽压力,后者带来长 Prompt 的计算压力。二者不能混为一谈。

14. 一个简化的 PyTorch 实现#

下面的实现展示标准 causal self-attention 的核心数据流。它适合帮助理解张量形状,不包含 KV Cache、RoPE、dropout 和生产级 kernel 优化。

import math
 
import torch
import torch.nn as nn
 
 
class MultiHeadAttention(nn.Module):
    def __init__(self, hidden_size: int, num_heads: int) -> None:
        super().__init__()
 
        if hidden_size % num_heads != 0:
            raise ValueError("hidden_size must be divisible by num_heads")
 
        self.hidden_size = hidden_size
        self.num_heads = num_heads
        self.head_dim = hidden_size // num_heads
 
        self.q_proj = nn.Linear(hidden_size, hidden_size, bias=False)
        self.k_proj = nn.Linear(hidden_size, hidden_size, bias=False)
        self.v_proj = nn.Linear(hidden_size, hidden_size, bias=False)
        self.o_proj = nn.Linear(hidden_size, hidden_size, bias=False)
 
    def forward(
        self,
        x: torch.Tensor,
        attention_mask: torch.Tensor | None = None,
    ) -> torch.Tensor:
        batch_size, seq_len, _ = x.shape
 
        q = self.q_proj(x)
        k = self.k_proj(x)
        v = self.v_proj(x)
 
        # [B, S, H * D] -> [B, H, S, D]
        q = q.view(
            batch_size, seq_len, self.num_heads, self.head_dim
        ).transpose(1, 2)
        k = k.view(
            batch_size, seq_len, self.num_heads, self.head_dim
        ).transpose(1, 2)
        v = v.view(
            batch_size, seq_len, self.num_heads, self.head_dim
        ).transpose(1, 2)
 
        # [B, H, S, D] @ [B, H, D, S] -> [B, H, S, S]
        scores = torch.matmul(q, k.transpose(-2, -1))
        scores = scores / math.sqrt(self.head_dim)
 
        if attention_mask is not None:
            scores = scores + attention_mask
 
        attention_weights = torch.softmax(scores, dim=-1)
 
        # [B, H, S, S] @ [B, H, S, D] -> [B, H, S, D]
        output = torch.matmul(attention_weights, v)
 
        # [B, H, S, D] -> [B, S, H * D]
        output = output.transpose(1, 2).contiguous()
        output = output.view(batch_size, seq_len, self.hidden_size)
 
        return self.o_proj(output)

其中 attention_mask 需要能够广播到 [B, H, S, S],而且应该是加性 mask:允许访问的位置为 0,禁止访问的位置为负无穷或足够大的负数。

生产级大模型推理通常还会使用:

  • fused QKV projection;
  • fused RoPE;
  • FlashAttention;
  • PagedAttention;
  • KV Cache;
  • continuous batching;
  • 针对不同 Prefill/Decode 形状优化的 fused kernels。

这些技术改变的是数据布局、内存管理和 kernel 执行方式,标准 MHA 的数学语义仍然不变。

15. 如何从模型配置识别 MHA#

Hugging Face 风格的模型配置通常包含:

{
  "num_attention_heads": 32,
  "num_key_value_heads": 32
}

如果:

num_attention_heads == num_key_value_heads

通常就是标准 MHA。

GQA 则可能是:

{
  "num_attention_heads": 32,
  "num_key_value_heads": 8
}

MQA 通常是:

{
  "num_attention_heads": 32,
  "num_key_value_heads": 1
}

不过,不同模型库的字段命名并不完全一致。有些旧配置不显式提供 num_key_value_heads,此时实现可能默认令它等于 num_attention_heads。判断时最好同时查看配置解析逻辑和 Attention 模块,而不是只依赖某一个字段。

总结#

标准 MHA 的完整数据流可以压缩成:

hidden states
  -> Q/K/V projection
  -> 拆成 H 个 heads
  -> 每个 Q head 与对应 K/V head 独立做 attention
  -> 拼接所有 head
  -> output projection

它最关键的结构特征是:

Hq=Hk=Hv=Hkv.\boxed{H_q=H_k=H_v=H_{kv}}.

MHA 为每个 Query Head 保留独立的 Key/Value 表示,结构对称、表达约束少,也是理解 Transformer Attention 的标准起点。它的代价同样明确:KV Cache 容量最大,Decode 阶段需要读取的历史 KV 最多,在长上下文和高并发服务中容易受到显存容量与带宽限制。

GQA 和 MQA 并没有改变“Query 对历史 Key 做匹配,再聚合 Value”这一基本逻辑。它们真正改变的是 Query Heads 共享 KV Heads 的程度,用一定的表示共享换取更小的 KV Cache 和更高效的 Decode。