Grouped Query Attention
Grouped-Query Attention
GQA 把 Hq 个 Query heads 分成若干组,每组共享一个 K/V head。它保留多种 query 子空间,却只维护中间数量的历史 KV,是 Multi-Head Attention 与 Multi-Query Attention 之间的结构旋钮。
形状和映射
- Query heads:
Hq;KV heads:Hkv。 - 通常要求
Hq % Hkv == 0,每个 KV head 服务g=Hq/Hkv个 Q head。 - Q:
[B,Hq,Tq,Dh];K/V:[B,Hkv,Tk,Dh]。 - 相对 MHA,KV Cache 容量近似按
Hkv/Hq缩减,但实际还受层数、block 粒度和元数据影响。
1 | Q0 Q1 Q2 Q3 Q4 Q5 Q6 Q7 |
效率从哪里来
decode 每一步都要读取历史 K/V。GQA 减少写入和读取的 KV head 数,因此:
1 | Hkv 下降 → KV cache 变小 → 每步内存搬运减少 |
它主要改善内存侧瓶颈,不会把 Q/K 点积的语义计算变成零成本。若 GPU 已经被 GEMM 或通信占满,GQA 的端到端收益可能小于 KV 公式给出的比例。
repeat_kv 的边界
实现常把 K/V 逻辑扩展到 Hq 以适配矩阵乘法;这不代表缓存必须把同一数据复制 g 份。是否物理复制取决于 kernel、layout 和编译路径。评估时应查看实际显存峰值与内存读流量,而不是从变量名推断。
Hkv=Hq 退化为 MHA,Hkv=1 退化为 MQA。GQA 的中间点不是必然最优:质量、TPOT、吞吐和长上下文能力要在目标模型与任务上联合验证。论文中的 uptraining 结果也不能直接替代你的模型回归。
关系与验证
它向上服务 Autoregressive Generation 的 decode,向下依赖 Tensor Shape 和 KV Cache,横向比较 MHA/MQA。验证要固定质量集、硬件、batch 和长度分布,记录质量、KV 峰值、TTFT、TPOT 与吞吐。
参考资料
Hkv 是模型容量和系统带宽之间的旋钮
固定 Hq 时,Hkv 越小,多个 query head 看到的 K/V 子空间越相似;Hkv 越大,读取和保存历史的代价越接近 MHA。这个取舍不是线性的质量曲线:某些层可能对共享更敏感,某些层的 head 高度冗余。若做结构改造,按层记录 KV head 数和任务回归,比全模型统一采用一个比例更有信息量。
物理布局决定理论收益能否落地
一个常见路径是把 [B,Hkv,Tk,Dh] 通过 view/expand 映射成 [B,Hq,Tk,Dh],让后续 kernel 看起来像 MHA;如果 expand 保持 stride 0,读取可以共享底层存储,若随后调用 contiguous,就会真的复制。优化 GQA 时应检查 graph 中是否出现 materialize、KV cache block 的 occupancy 以及 HBM 读流量。否则只凭 Hkv/Hq 估算,容易得到虚假的吞吐预期。
除了质量和 TPOT,还要验证 batch 变化下的行为。GQA 的收益通常在长历史、低 Tq、较高并发时更明显;短 prompt 或 compute-bound 的 prefill 可能几乎看不出来。把 prompt 长度、生成长度、并发和硬件列为固定实验轴,才能解释为什么同一改动在两个服务场景结论不同。