/ LLM Algorithms
[LLM Algorithms-3] Standard MHA: From Multi-Head Computation to KV Cache
A shape-by-shape guide to standard Multi-Head Attention, covering Q/K/V projections, per-head computation, output concatenation, prefill, decode, KV cache, and the differences between MHA, GQA, and MQA.
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 满足:
也常简写为:
这里的“一一对应”描述的是 Attention 计算时的配对关系,并不表示 Q、K、V 使用相同的参数。它们来自不同的线性投影,参数彼此独立。
1. 从输入和 Head Dimension 开始#
以 self-attention 为例,输入 hidden states 记为:
其中:
- :batch size;
- :序列长度;
- :模型隐藏维度。
假设模型配置为:
batch_size = 8
sequence_length = 2048
hidden_size = 4096
num_heads = 32每个 head 的维度为:
这要求 能被 整除。拆分 head 并不会改变总特征维度,因为:
2. 从 Hidden States 生成 Q、K、V#
输入 分别经过三组线性投影:
在标准 MHA 中,若忽略 bias:
所以投影后的张量形状仍然是:
接下来把最后一个维度拆成 个 head,并把 head 维移动到序列维之前:
[B, S, d_model]
-> [B, S, H, d_h]
-> [B, H, S, d_h]在前面的例子中:
[B, S, 4096] -> [B, 32, S, 128]于是:
工程实现常把三次投影融合成一次大的 QKV projection,以减少 kernel launch 和内存访问;这只改变执行方式,不改变数学结构。
3. 每个 Head 如何计算 Attention#
第 个 Attention Head 只使用与它对应的 、 和 :
其中 是 attention mask。这个过程可以拆成三步。
3.1 计算注意力分数#
若:
那么:
表示在第 个表示子空间里,第 个 query token 与第 个 key token 的匹配分数。它还不是概率,因为其中可能包含负值,且每行之和不等于 1。
3.2 缩放、Mask 与 Softmax#
为什么要除以 ?
如果 Q、K 各维近似独立且方差为 1,那么点积 的方差会随 增长。较大的 logits 容易让 Softmax 过早进入饱和区,产生接近 one-hot 的分布和很小的梯度。除以 可以把 logits 的尺度稳定在更合适的范围。
Mask 通常在 Softmax 之前加入。对不允许访问的位置,工程实现会加入负无穷或一个足够大的负数,使其 Softmax 权重接近 0。
在 Causal Language Model 中,因果 Mask 保证第 个 token 不能看到未来位置 :
3.3 对 Value 加权求和#
形状变化为:
[B, S, S] @ [B, S, d_h] -> [B, S, d_h]因此:
注意力权重由 Q 和 K 决定,真正被聚合到输出中的内容来自 V。可以用“Q 发起查询、K 提供地址、V 携带内容”来形成直觉,但三者本质上都是模型学习到的向量表示。
4. 拼接所有 Head 并做输出投影#
所有 head 独立计算后,先沿 head 对应的特征维拼接:
形状从:
H 个 [B, S, d_h]恢复为:
[B, S, H * d_h] = [B, S, d_model]即:
最后经过输出投影:
其中:
不只是做形状变换。它会重新混合各个 head 产生的特征,让后续层能够联合使用不同 head 提取的信息。
5. 标准 MHA 的完整公式#
一般形式可以写成:
其中:
而 Scaled Dot-Product Attention 是:
在 self-attention 中,公式里的 Q、K、V 来自同一组输入 hidden states;在 cross-attention 中,Query 可以来自 decoder,而 Key 和 Value 来自 encoder。MHA 描述的是多头计算与投影方式,并不只限于 self-attention。
6. 为什么需要多个 Head#
如果只使用一个维度为 的大 head,模型只会产生一组注意力分布。多头结构则允许模型在不同的低维表示子空间里,同时建立多组 token 关系。
从直觉上看,不同 head 可能更偏向:
局部相邻关系
句法依赖
实体指代
长距离关联
段落或分隔符结构不过,这只是帮助理解的观察方式,并不意味着每个 head 都一定具有稳定、唯一、可由人命名的语义。
从参数角度看,每个 head 都对应独立的投影切片:
因此,各个 head 可以学习不同的 query-key 匹配规则和 value 表示。多头的主要价值不是增加总隐藏维度,而是让模型并行学习多组不同的交互模式。
7. 自回归推理为什么需要 KV Cache#
训练或 Prefill 时,模型通常可以并行处理一段 token;但在 Decode 阶段,自回归模型每一步只生成一个新 token。
假设当前已经有 个可见 token。新 token 的 Query 形状为:
历史 Key 和 Value 的形状为:
新 Query 会查询所有历史位置:
新 token 的 K 和 V 在生成后被追加到 cache,供下一步使用。历史 token 的 K/V 只依赖已经确定的 hidden states,不会在后续步骤变化,因此没有必要反复计算。这就是 KV Cache 的基本原理。
对单步 Decode 来说,当前 token 已经位于序列末尾,没有未来位置可见,所以通常不需要显式构造一个完整的三角 causal mask;但实现仍需正确处理 padding、滑动窗口或分块等边界。
8. MHA 的 KV Cache 有多大#
忽略对齐、分页和元数据开销,对单层模型而言,标准 MHA 的 KV Cache 大小近似为:
其中:
- 前面的 2 表示一份 K Cache 和一份 V Cache;
- 是 batch size 或同时缓存的序列数;
- 是缓存长度;
- 是 KV Heads 数量;
- 是每个 head 的维度;
- 是每个元素的字节数。
对标准 MHA, 且 ,所以单层也可以写成:
如果模型共有 层,总量近似为:
8.1 一个 16 GiB 的具体例子#
假设模型配置为:
layers = 32
batch = 1
sequence_length = 32768
num_kv_heads = 32
head_dim = 128
dtype = FP16 / BF16, 2 bytes那么:
这还只是 batch 为 1 时的理论数据大小。若 batch 增加到 8,在所有请求都缓存 32K token 的前提下:
KV Cache ~= 128 GiB实际推理引擎还可能有 block 对齐、未填满页面、内存碎片和管理元数据等额外开销。反过来,KV Cache 量化、滑动窗口 Attention 或前缀共享也可能降低实际物理占用。
9. MHA 的计算复杂度#
对于长度为 的 self-attention,每个 head 的 QK 计算是:
所有 head 的复杂度约为:
Attention probabilities 与 V 相乘的复杂度处于相同量级。因此,若只看 Attention 主体,它对序列长度呈二次增长:
不过,完整 Transformer 层还包含 QKV projection、输出投影和 MLP。线性投影的成本通常是 ,所以在上下文较短、隐藏维度很大时,Attention 不一定占据全部计算;随着 增大,二次项才会越来越突出。
FlashAttention 可以通过 tiling 和在线 Softmax 避免把完整 注意力矩阵写回显存,显著降低中间显存占用与 IO,但它并没有把标准 Softmax Attention 的算术复杂度从 变成 。
10. Prefill 与 Decode 是两种不同的瓶颈#
同一套 MHA 数学结构,在 Prefill 和 Decode 阶段会呈现完全不同的硬件特征。
10.1 Prefill:大矩阵,通常更偏计算密集#
假设一次输入 16K prompt:
q_len = 16384
kv_len = 16384每个 head 的逻辑 Attention 矩阵大小为:
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。
假设 :
| 结构 | Q Heads | KV Heads | 共享方式 | 相对 MHA 的理论 KV Cache |
|---|---|---|---|---|
| MHA | 32 | 32 | 每个 Q Head 独享一组 KV | |
| GQA | 32 | 8 | 每 4 个 Q Heads 共享一组 KV | |
| MQA | 32 | 1 | 所有 Q Heads 共享一组 KV |
它们的连接关系可以画成:
MHA
Q0 -> KV0
Q1 -> KV1
Q2 -> KV2
Q3 -> KV3
GQA
Q0 --+
Q1 --+--> KV0
Q2 --+
Q3 --+
MQA
Q0 --+
Q1 --+
Q2 --+--> KV0
Q3 --+
... |更一般地,若 能被 整除,每组 KV 服务的 Query Heads 数量是:
对同样的层数、序列长度、head dimension 和数据类型,KV Cache 相对标准 MHA 的缩放比例近似为:
减少 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 结构对称#
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 HeadsQ/K/V 完全对称。不过,“容易切分”仍依赖头数能否整除 TP size、具体并行布局和通信实现,并不是任何配置下都自动成立。
13. 标准 MHA 的局限#
13.1 KV Cache 最大#
由于:
在 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 同时面对两种增长:
前者带来容量和 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它最关键的结构特征是:
MHA 为每个 Query Head 保留独立的 Key/Value 表示,结构对称、表达约束少,也是理解 Transformer Attention 的标准起点。它的代价同样明确:KV Cache 容量最大,Decode 阶段需要读取的历史 KV 最多,在长上下文和高并发服务中容易受到显存容量与带宽限制。
GQA 和 MQA 并没有改变“Query 对历史 Key 做匹配,再聚合 Value”这一基本逻辑。它们真正改变的是 Query Heads 共享 KV Heads 的程度,用一定的表示共享换取更小的 KV Cache 和更高效的 Decode。