实现 GQA/MQA 支持,让多个 Query Head 共享 KV,优化 KV Cache 内存占用。
配套代码
在前面的章节中,我们实现了完整的 Flash Attention Kernel,支持任意序列长度、多维并行和 Causal Masking。本章将添加最后一个关键特性:Grouped Query Attention (GQA)。
这是 Llama 2/3、Mistral、Falcon 等主流开源模型的标配技术,能够在几乎不损失模型质量的前提下,将 KV Cache 内存占用减少 4-8 倍。
Grouped Query Attention (GQA) 通过让多个 Query heads 共享同一组 Key/Value heads 来减少 KV Cache 内存:
# MHA: 每个 Q head 有独立的 K/V head
Q: (B, 8, N, D)
K: (B, 8, N, D) # 8 个 KV heads
V: (B, 8, N, D)
# GQA: 多个 Q heads 共享 KV heads (groups=4)
Q: (B, 8, N, D)
K: (B, 2, N, D) # 只有 2 个 KV heads → 内存减少 4x
V: (B, 2, N, D)
# Q heads 0-3 → KV head 0
# Q heads 4-7 → KV head 1
Head 索引映射公式:
如果您对 GQA/MQA 的概念、演进历史、内存分析和设计权衡还不熟悉,建议先阅读 Attention 机制详解 中的详细介绍。
本章重点关注如何在 Flash Attention 中零拷贝、零额外内存地实现 GQA。
标准的 GQA 实现需要先扩展 KV heads:
systems/flash_attention/07_gqa.py
def standard_attention(Q, K, V, causal=False):
B, H_Q, N, D = Q.shape
H_KV = K.shape[1]
num_groups = H_Q // H_KV
# 将 KV heads 扩展到与 Q heads 匹配
K = K.repeat_interleave(num_groups, dim=1) # (B, 2, N, D) → (B, 8, N, D)
V = V.repeat_interleave(num_groups, dim=1)
# 然后进行标准的 MHA 计算
score = torch.matmul(Q, K.transpose(-2, -1))
# ...
性能问题: