Q 总投影宽度不变时,更大的 GQA 比例(Q/KV heads 比)意味着更小的 KV Cache。但更小的 KV Cache 不总是意味着更快的推理,还与投影和 softmax 的成本有关。下文通过 roofline 分析,结合硬件的算力和带宽,定量分析调整该比例对 prefill 和 decode 性能的影响。
以一层 GQA 为例,基准配置为 、、、,GQA 比例为 8:1,Q 总投影宽度 。保持 hidden size 和 KV heads 数 不变,通过调整 与 改变比例:比例提高时, 增大、 减小;比例降低时则相反。
数据量变化
权重 、 的参数量合计固定为 ,、 合计为 ,随 同比增减。KV Cache 容量也与 成正比。设 batch 含 条长度为 的独立序列,KV Cache 每元素占 bytes,单层缓存为:
取 、,权重与 KV Cache 均为 BF16。将 GQA 比例从 8:1 提高到 16:1 或降低到 4:1,单层四项投影权重和 KV Cache 的数据量如下:
| Q heads 数 | GQA 比例 | head 维度 | 投影权重 | KV Cache | 两项数据量合计 |
|---|---|---|---|---|---|
| 32 | 4:1 | 256 | 160 MiB | 256 MiB | 416 MiB |
| 64 | 8:1 | 128 | 144 MiB | 128 MiB | 272 MiB |
| 128 | 16:1 | 64 | 136 MiB | 64 MiB | 200 MiB |
比例从 8:1 提高到 16:1 时,两项数据量合计减少约 26%;降低到 4:1 时,增加约 53%。batch 内的请求共用一份投影权重,增大 batch 能摊薄每个请求的权重加载成本,而 KV Cache 随独立序列数增长,在总加载数据中的占比随 batch 数提高。因此,batch 为 8 时,同样将比例从 8:1 提高到 16:1,合计数据量可减少约 45%。
但 KV Cache 缩小的同时,Attention scores 和 softmax 输出的数据量却会增大。每个 Q head 都通过 生成一份 scores,形状在 prefill 中为 ,在单步 decode 中为 。减小 只会缩短点积长度,不会减少 scores 的元素数; 翻倍则让 scores 及同形状的 softmax 输出都增加一倍。各项数据量的变化如下:
| 数据对象 | Prefill 元素数 | 单步 decode 元素数 | Q/KV 比 8:1→16:1 | Q/KV 比 8:1→4:1 |
|---|---|---|---|---|
| KV Cache | 减半 | 两倍 | ||
| scores、softmax 输出(各自) | 两倍 | 减半 |
计算成本分析
Attention 交互
GQA 比例越大,共享同一份 KV 的 Q heads 就越多,交互计算的算术强度也越高。将 、 的 FLOPs 合计,先只考虑 KV 读取,算术强度为:
BF16 下 ,decode 的算术强度就是 ;prefill 则是它的 倍。
以 H200 BF16 为例,硬件强度为 。GQA 比例为 8:1 和 16:1 时, 分别为 8 和 16,均远低于这个临界值,按此估算都处于访存受限区。但提高 GQA 比例仍能使交互计算的算术强度更接近硬件强度。
KV 读取只是交互计算的一部分 IO,实际还包括 、 的读写,以及写回显存的 scores 和 softmax 输出。计入这些中间数据后,交互 FLOPs 不变,算术强度会比前面的估算更低。提高 GQA 比例同时减少 KV 数据量、增加 scores 和 softmax 输出的数据量,两者对 IO 的影响相反,因此实际算术强度不会随比例同比增长。
线性投影
记 token 总数为 (prefill:;单步 decode:)。权重、输入和输出均用 BF16,按权重与输入各读一次、输出写一次估算。以 Q 投影为例,每个 token 的输入有 个元素,输出有 个元素,因此输入读取和输出写入合计为 bytes,权重读取量为固定的 bytes。权重由所有 token 共用,但输入输出的 IO 随 增长。每项投影的成本如下:
| 投影(每项) | FLOPs | IO 字节数 | 算术强度(FLOPs/byte) |
|---|---|---|---|
| 、 | |||
| 、 |
以 Q 投影为例,当 较少,满足 时,IO 主要来自权重读取,算术强度可以近似为:
因此,decode 中增大 batch、prefill 中一次处理更多 token,都能摊薄每个 token 的权重加载成本。对于相同输入,调整 GQA 比例并不改变 ,所以在 的近似下,比例不影响投影的算术强度。但随着 增大,输入输出的 IO 逐渐占主导,这个近似不再成立,继续增加 token 数对算术强度的增益也会减弱。
提高比例对总投影计算量的影响有限。比例从 8:1 提高到 16:1,投影 FLOPs 只减少约 5.6%;降低到 4:1 时,投影 FLOPs 增加约 11.1%。这是因为 Q/O 投影不变,随比例变化的 K/V 投影只占四项投影 FLOPs 的 。上下文越长,交互计算占比越高,投影计算的这点变化对总 FLOPs 的影响就越小。
另外,提高比例使 减小,K/V 投影的权重和输出数据量随之减少,但输入 hidden states 的读取量不变。因此 IO 的降幅小于 FLOPs 的降幅,算术强度反而下降。
非线性操作
提高比例会增加 softmax 的成本。因为 Q heads 数增多,每行 softmax 要处理的 scores 却没有减少,因此总工作量随 heads 数增长。减小 缩短的是生成 scores 时的点积长度,并不能减少后续的归一化计算。
而 softmax 的单位运算耗时也与矩阵乘法不同。矩阵乘法通常由 Tensor Cores 加速,softmax 的加法、乘法等运算使用 CUDA Cores,指数运算通常使用 SFU。指数运算的吞吐远低于 Tensor Cores 的矩阵乘法吞吐,即使运算次数占比小,也可能占用较多时间。1
GQA 比例选择准则
提高比例有利于缓解 KV 读取瓶颈,却会增加非线性操作的负担。因此,比例的选择应以主要瓶颈为依据:
| 环节 | 主要耗时 | 调整方向 | 原因 |
|---|---|---|---|
| 投影 | K/V 投影 | 提高比例 | 减少权重数据量、输出数据量和计算量,算术强度本身不会提高 |
| 投影 | Q/O 投影 | 调整比例无法直接降低这部分成本 | 矩阵形状与计算量不变 |
| Attention 交互:、 | KV 读取 | 提高比例 | 交互 FLOPs 不变,KV 数据量减少,算术强度提高 |
| Attention 交互:、 | 矩阵计算 | 无固定的增减方向 | 总 FLOPs 不变,矩阵形状变化会影响硬件利用率 |
| Softmax | scores 的归一化与读写 | 降低比例 | 减少需要处理的 scores 及对应 IO |
性能取舍由服务目标决定。
Head 维度的选择还需兼顾模型任务表现和 kernel 支持。当进一步减小 head 维度不再合适时,还可以改变 KV 的缓存方式。例如 MLA 缓存 K/V 共享的低维表示和额外的 RoPE key,使缓存大小不再直接取决于展开后的 heads 数和维度。
参考资料
-
Tri Dao, FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision,关于 GEMM 与 softmax 执行吞吐及重叠的分析。 ↩