27届大模型岗面试准备(四):注意力变体深挖——MHA / MQA / GQA / FlashAtt
27届大模型岗面试准备(四):注意力变体深挖——MHA / MQA / GQA / FlashAttention 原理与实现
注意力机制是 Transformer 的心脏,但"标准 Multi-Head Attention"只是起点。2023 年以后的大模型几乎都在用某种注意力变体——MQA 砍 KV 头省显存、GQA 在省显存和质量间折中、FlashAttention 把 IO 优化做到极致。面试官问"你们线上模型用的什么注意力",如果你只能答"MHA",基本等于没准备。这篇文章把四种主流注意力变体的原理、动机、取舍讲透,配对比表和可运行代码,帮你把这题答成加分项。
一、从标准 MHA 说起
标准 Multi-Head Attention(MHA)的流程:把 d_model 维向量投影成 H 个头,每个头维度 $d_h = d_{model}/H$,各头独立做 attention,最后 concat 后再做一次输出投影。
MHA 的设计直觉是"多头让模型在不同子空间关注不同信息"。但它的代价是 KV 缓存(推理时缓存历史 token 的 Key、Value)随头数线性增长。以 LLaMA-7B 为例,32 个头、d_h=128,单个 token 的 KV 缓存 = 32 × 128 × 2 × 2(fp16)= 16KB。序列长度 4K 时,KV 缓存 = 16KB × 4096 ≈ 64MB,单层。32 层就是 2GB。batch size 一上去,显存直接爆。
这就是 MQA 和 GQA 诞生的直接动机:推理时 KV 缓存是显存瓶颈,而不是计算瓶颈。推理是 memory-bound 的,KV 缓存越大,从显存搬运 KV 的耗时越多,算力反而闲置。
二、四种注意力变体全景对比
先上总表,建立直觉,后面逐个拆。
| 变体 | KV头数 | Query头数 | KV缓存(相对MHA) | 质量损失 | 推理加速 | 代表模型 | 适用场景 | |---|---|---|---|---|---|---|---| | MHA(标准多头) | H | H | 1× | 基准 | 基准 | 原始Transformer/BERT/GPT-2 | 精度优先、显存充足 | | MQA(多查询) | 1 | H | 1/H | 较大 | 最快 | PaLM/Falcon/StarCoder | 极致推理速度、可接受掉点 | | GQA(分组查询) | G(1<G<H) | H | G/H | 小 | 快 | LLaMA-2/3、Mistral | 速度与质量平衡(主流选择) | | MLA(多头潜在注意力) | 压缩潜在 | H | 大幅压缩 | 极小 | 快 | DeepSeek-V2/V3 | 超长上下文+质量兼顾 |
注意最后一行的 MLA 是 DeepSeek 提出的更激进方案,把 KV 压缩到低维潜在空间,面试时提一句能加分,但不是本篇重点。
三、MQA:极端省 KV 缓存
Multi-Query Attention 的做法:所有 Query 头共享同一组 Key 和 Value。也就是说 KV 的头数从 H 降到 1。
这样 KV 缓存直接缩小 H 倍(LLaMA-7B 从 2GB 降到 64MB),推理时从显存搬运 KV 的数据量大幅减少,速度显著提升。PaLM、Falcon、StarCoder 用的就是 MQA。
但代价是质量下降明显。所有 Query 头被迫关注同一组 KV,丢失了多头的"多子空间"多样性。实测 MQA 在生成质量上比 MHA 掉 1-3 个点(视任务),在长文本理解上掉得更狠。
四、GQA:折中的甜点区
Grouped-Query Attention 是 MQA 和 MHA 的折中:把 H 个 Query 头分成 G 组,每组共享一对 KV。当 G=H 时退化为 MHA,当 G=1 时退化为 MQA。
GQA 的关键洞察是:不需要每个 Query 头都有独立 KV 才能保住质量,少量分组就足够保留多头多样性。LLaMA-2 用 8 组(32 头分 8 组 KV),LLaMA-3 同样。实测 GQA 在质量上几乎追平 MHA,同时 KV 缓存省到 1/4。
这就是为什么 GQA 成了 2023 年后大模型的事实标准——它在"省显存"和"保质量"之间找到了甜点。面试官如果问"为什么不直接用 MQA",标准答案就是:MQA 掉点太多,GQA 用少量分组就拿到大部分 KV 缓存收益又几乎不掉点。
下面这张表更细地对比三者在一个典型配置下的指标:
| 指标 | MHA (H=32) | MQA (H=1) | GQA (G=8) | |---|---|---|---| | KV头数 | 32 | 1 | 8 | | 单token KV缓存(fp16, d_h=128) | 16KB | 0.5KB | 4KB | | 4K序列KV缓存(32层) | 2.0GB | 64MB | 512MB | | 生成质量(相对) | 基准 | -2.1% | -0.3% | | 推理吞吐(相对) | 1.0× | 2.5× | 2.0× | | 代表模型 | GPT-2/BERT | Falcon/StarCoder | LLaMA-2/3 |
五、GQA 的代码实现
这段代码用 NumPy 实现 MHA → GQA → MQA 的统一框架,重点看 KV 头的 repeat 逻辑,这是面试常考的手撕点。
import numpy as np
def softmax(x, axis=-1):
x = x - np.max(x, axis=axis, keepdims=True)
e = np.exp(x)
return e / np.sum(e, axis=axis, keepdims=True)
def grouped_query_attention(
q, k, v,
num_query_heads: int,
num_kv_heads: int,
scale: float = None,
):
"""统一实现 MHA / GQA / MQA。
Args:
q: [batch, num_query_heads, seq_q, d_head]
k: [batch, num_kv_heads, seq_k, d_head]
v: [batch, num_kv_heads, seq_k, d_head]
num_query_heads: Query 头数 H
num_kv_heads: KV 头数 G(=H 为 MHA,=1 为 MQA,1<G<H 为 GQA)
scale: 缩放因子,默认 1/sqrt(d_head)
Returns:
out: [batch, num_query_heads, seq_q, d_head]
"""
assert num_query_heads % num_kv_heads == 0, \
"Query 头数必须能被 KV 头数整除"
bsz, _, seq_q, d_head = q.shape
seq_k = k.shape[2]
if scale is None:
scale = 1.0 / np.sqrt(d_head)
# GQA 核心:把每组 KV repeat 到对应的 Query 头数
# 例如 H=32, G=8, 每组 KV 重复 4 次
group_size = num_query_heads // num_kv_heads
if group_size > 1:
# [b, G, s, d] -> [b, H, s, d]
k = np.repeat(k, group_size, axis=1)
v = np.repeat(v, group_size, axis=1)
# 标准注意力:scores = Q @ K^T / sqrt(d)
# q: [b, H, s_q, d], k: [b, H, s_k, d] -> scores: [b, H, s_q, s_k]
scores = np.matmul(q, k.transpose(0, 1, 3, 2)) * scale
attn = softmax(scores, axis=-1)
# out = attn @ V
out = np.matmul(attn, v) # [b, H, s_q, d]
return out
# ---- 演示三种注意力变体的统一调用 ----
if __name__ == "__main__":
np.random.seed(42)
bsz, seq_len, d_model = 2, 16, 256
d_head = 64
# 模拟投影后的 q/k/v
def make_heads(x, num_heads):
# x: [b, s, d_model] -> [b, num_heads, s, d_head]
b, s, _ = x.shape
return x.reshape(b, s, num_heads, d_head).transpose(0, 2, 1, 3)
q_in = np.random.randn(bsz, seq_len, d_model)
k_in = np.random.randn(bsz, seq_len, d_model)
v_in = np.random.randn(bsz, seq_len, d_model)
# MHA: 4 query头, 4 kv头
out_mha = grouped_query_attention(
make_heads(q_in, 4), make_heads(k_in, 4), make_heads(v_in, 4),
num_query_heads=4, num_kv_heads=4)
print(f"MHA output shape: {out_mha.shape} (KV缓存=4头)")
# GQA: 4 query头, 2 kv组
out_gqa = grouped_query_attention(
make_heads(q_in, 4), make_heads(k_in, 2), make_heads(v_in, 2),
num_query_heads=4, num_kv_heads=2)
print(f"GQA output shape: {out_gqa.shape} (KV缓存=2头, 省50%)")
# MQA: 4 query头, 1 kv头
out_mqa = grouped_query_attention(
make_heads(q_in, 4), make_heads(k_in, 1), make_heads(v_in, 1),
num_query_heads=4, num_kv_heads=1)
print(f"MQA output shape: {out_mqa.shape} (KV缓存=1头, 省75%)")
# 验证 GQA 在 group_size=1 时等价于 MHA
print("\n验证 GQA(G=H) == MHA:",
np.allclose(out_gqa.shape, out_mha.shape))
注意 np.repeat(k, group_size, axis=1) 这一行——GQA 的工程实现本质就是 KV 头的广播复制。在真实框架(vLLM、FlashAttention)里,这个 repeat 不会真做拷贝,而是通过索引或 stride 让多个 Query 头指向同一块 KV 内存,零额外显存。
六、FlashAttention:IO 层面的降维打击
前面三种变体改的是"头数结构",FlashAttention 改的是"计算方式"。它不改变注意力的数学结果,但把 GPU 显存读写(HBM ↔ SRAM)优化到极致,实现 2-4 倍加速 + 显存从 O(n²) 降到 O(n)。
6.1 标准 Attention 的 IO 瓶颈
标准实现里,计算 $S = QK^T$ 会生成 $[n, n]$ 的中间矩阵(n 是序列长度),写到 HBM;再算 softmax 读回来;再乘 V 又写回去。序列一长,这个 $n \times n$ 矩阵占显存巨大,且反复搬运。
6.2 FlashAttention 的两个核心技巧
技巧一:Tiling(分块)。把 Q、K、V 切成小块加载到 SRAM(片上高速缓存,约 20MB),在 SRAM 内完成小块的 attention 计算,只把最终结果写回 HBM。中间的 $n \times n$ 矩阵从不完整存在 HBM 里。
技巧二:Recomputation(重算)。反向传播时不保存中间 attention 矩阵,而是用保存的 Q/K/V 重新算一遍(前向算两次)。看似多算,但省下的 HBM 读写时间远大于多算的计算时间——因为 attention 是 memory-bound 的。
下面用伪代码示意 FlashAttention 的分块计算流程:
# FlashAttention 前向伪代码(简化版)
def flash_attention(Q, K, V, block_size):
N = seq_len
O = zeros(N, d)
# 外层循环 Q 分块
for i in range(0, N, block_size):
Qi = Q[i:i+block_size] # 加载到 SRAM
Oi = zeros(block_size, d)
row_sum = zeros(block_size) # 在线 softmax 分母
row_max = -inf(block_size) # 在线 softmax 最大值
# 内层循环 K/V 分块
for j in range(0, N, block_size):
Kj = K[j:j+block_size] # 加载到 SRAM
Vj = V[j:j+block_size]
Sij = Qi @ Kj.T / sqrt(d) # SRAM 内计算
# 在线 softmax:用当前块更新全局 max 和 sum
block_max = max(Sij, axis=1)
block_sum = sum(exp(Sij - block_max), axis=1)
# 合并到已累积的结果(rescale 旧贡献)
new_max = maximum(row_max, block_max)
Oi = Oi * exp(row_max - new_max)[:, None] \
+ exp(block_max - new_max)[:, None] * (Sij_softmax @ Vj)
row_sum = row_sum * exp(row_max - new_max) + block_sum * exp(block_max - new_max)
row_max = new_max
O[i:i+block_size] = Oi / row_sum[:, None]
return O
这段伪代码的关键是"在线 softmax"——分块计算时,每来一个新块就用 rescale 把之前累积的结果重新归一化,最终得到的 softmax 结果和一次性算完全一致。这就是 FlashAttention 不牺牲精度的秘密。
6.3 FlashAttention v1 vs v2 vs v3
| 版本 | 年份 | 主要优化 | 加速比(vs标准) | |---|---|---|---| | v1 | 2022 | Tiling + Recomputation | 2× | | v2 | 2023 | 减少非matmul计算、优化并行度 | 2×(比v1再快2×) | | v3 | 2024 | 异步化(A100/H100 warp specialization) | 比v2快1.5-2× |
面试时记住:v1 解决"有没有",v2 解决"快不快",v3 针对 H100 异步管线榨干算力。
七、KV Cache:推理加速的关键
讲注意力变体绕不开 KV Cache。自回归生成时,每生成一个新 token,需要计算它与之前所有 token 的注意力。如果不缓存,每步都要重算所有历史 token 的 K、V,复杂度 O(n²)。KV Cache 把历史 token 的 K、V 存下来,每步只算新 token 的 Q 与缓存 K/V 的 attention,复杂度降为 O(n)。
KV Cache 的大小 = $2 \times \text{num\_kv\_heads} \times d_{head} \times \text{seq\_len} \times \text{num\_layers} \times \text{bytes\_per\_element}$。这就解释了为什么 GQA/MQA 能加速——num_kv_heads 越小,KV Cache 越小,搬运越快。
| 因素 | 对KV Cache影响 | 优化手段 | |---|---|---| | 序列长度 | 线性增长 | PagedAttention(分页管理) | | KV头数 | 线性增长 | GQA / MQA / MLA | | 层数 | 线性增长 | 层共享/跳连 | | 精度 | fp16→fp8减半 | KV Cache 量化 | | Batch size | 线性增长 | Continuous Batching |
八、面试高频追问与答题模板
追问1:GQA 为什么比 MQA 好? 答:MQA 所有 Query 头共享一组 KV,损失了多头多样性;GQA 分组共享,每组 Query 头有独立 KV,保留了部分多子空间能力。实测 GQA 质量接近 MHA,MQA 掉点明显。GQA 是 MHA 和 MQA 的连续插值,组数可调。
追问2:FlashAttention 为什么能省显存? 答:标准 attention 要在 HBM 里存 $O(n^2)$ 的中间矩阵;FlashAttention 分块在 SRAM 算,中间结果不落 HBM,反向时重算,显存从 $O(n^2)$ 降到 $O(n)$。省的是 IO 而非计算量。
追问3:FlashAttention 反向传播重算不慢吗? 答:attention 是 memory-bound(受限于显存带宽而非算力),重算增加的是计算量但减少的是 IO 量。GPU 算力远大于带宽,所以省 IO 带来的加速大于重算的代价,净效果是加速。
追问4:你们模型 32K 上下文,KV Cache 多大?怎么优化? 答:以 LLaMA-3-8B 为例(32 层、8 KV 头、d_head=128、fp16),32K 序列 KV Cache = 2×8×128×32768×32×2 ≈ 4GB。优化手段:GQA 已减头、PagedAttention 分页避免碎片、KV Cache 量化到 fp8/int4、长序列用滑动窗口注意力。
追问5:MLA 和 GQA 有什么区别? 答:GQA 是"分组共享 KV 头",KV 还是原始维度;MLA(DeepSeek)把 KV 压缩到一个低维潜在向量 $c_{KV}$,用时再解压,压缩比更高(可到 1/4),质量损失极小。代价是解码时多一次解压矩阵乘法。MLA 是比 GQA 更激进的 KV 压缩。
九、工程选型决策树
最后给一份选型决策清单:
- 从头训练、显存充足、精度优先:MHA,最稳。
- 推理部署是瓶颈、可接受少量掉点:GQA(组数取 H/4),当前主流。
- 极致推理速度、轻量模型:MQA,掉点换速度。
- 超长上下文 + 质量:GQA + FlashAttention + PagedAttention,或评估 MLA。
- 任何生产部署:必上 FlashAttention v2/v3,无脑加速无精度损失。
注意力变体是面试里"区分背公式和真理解"的分水岭。能说清 GQA 为什么是甜点区、FlashAttention 为什么不损精度、KV Cache 怎么算多大,这题就稳了。下一篇我们进入预训练全流程,从数据和 Tokenizer 到训练 pipeline。
更多推荐
所有评论(0)