实现 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 倍

GQA 快速回顾

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。

PyTorch 标准实现的问题

标准的 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))
    # ...

性能问题:

Flash Attention 中的高效实现