Skip to content
kefan.life
Go back

GQA 比例与推理成本

Q 总投影宽度不变时,更大的 GQA 比例(Q/KV heads 比)意味着更小的 KV Cache。但更小的 KV Cache 不总是意味着更快的推理,还与投影和 softmax 的成本有关。下文通过 roofline 分析,结合硬件的算力和带宽,定量分析调整该比例对 prefill 和 decode 性能的影响。

以一层 GQA 为例,基准配置为 D=4096D=4096、Hq=64H_q=64、Hkv=8H_{kv}=8、d=128d=128,GQA 比例为 8:1,Q 总投影宽度 Hqd=8192H_qd=8192。保持 hidden size DD 和 KV heads 数 HkvH_{kv} 不变,通过调整 HqH_q 与 dd 改变比例:比例提高时,HqH_q 增大、dd 减小;比例降低时则相反。

数据量变化

权重 WQW_Q、WOW_O 的参数量合计固定为 2DHqd2DH_qd,WKW_K、WVW_V 合计为 2DHkvd2DH_{kv}d,随 dd 同比增减。KV Cache 容量也与 dd 成正比。设 batch 含 BB 条长度为 SS 的独立序列,KV Cache 每元素占 ss bytes,单层缓存为:

MKV=2BSHkvd sM_{\text{KV}}=2BSH_{kv}d\,s

取 B=1B=1、S=32768S=32768,权重与 KV Cache 均为 BF16。将 GQA 比例从 8:1 提高到 16:1 或降低到 4:1,单层四项投影权重和 KV Cache 的数据量如下:

Q heads 数GQA 比例head 维度投影权重KV Cache两项数据量合计
324:1256160 MiB256 MiB416 MiB
648:1128144 MiB128 MiB272 MiB
12816:164136 MiB64 MiB200 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 都通过 QK⊤QK^\top 生成一份 scores,形状在 prefill 中为 S×SS\times S,在单步 decode 中为 1×S1\times S。减小 dd 只会缩短点积长度,不会减少 scores 的元素数;HqH_q 翻倍则让 scores 及同形状的 softmax 输出都增加一倍。各项数据量的变化如下:

数据对象Prefill 元素数单步 decode 元素数Q/KV 比 8:1→16:1Q/KV 比 8:1→4:1
KV Cache2BSHkvd2BSH_{kv}d2BSHkvd2BSH_{kv}d减半两倍
scores、softmax 输出(各自)BHqS2BH_qS^2BHqSBH_qS两倍减半

计算成本分析

Attention 交互

GQA 比例越大,共享同一份 KV 的 Q heads 就越多,交互计算的算术强度也越高。将 QK⊤QK^\top、PVPV 的 FLOPs 合计,先只考虑 KV 读取,算术强度为:

Idecode=4BHqdS2BSHkvd s=2sHqHkv,Iprefill=2BHqdS(S+1)2BSHkvd s=S+1sHqHkv.\begin{aligned} I_{\text{decode}} &=\frac{4BH_qdS}{2BSH_{kv}d\,s} =\frac{2}{s}\frac{H_q}{H_{kv}},\\ I_{\text{prefill}} &=\frac{2BH_qdS(S+1)}{2BSH_{kv}d\,s} =\frac{S+1}{s}\frac{H_q}{H_{kv}}. \end{aligned}

BF16 下 s=2s=2,decode 的算术强度就是 Hq/HkvH_q/H_{kv};prefill 则是它的 (S+1)/2(S+1)/2 倍。

以 H200 BF16 为例,硬件强度为 989.5 TFLOPS4.8 TB/s≈206 FLOPs/byte\frac{989.5\ \text{TFLOPS}}{4.8\ \text{TB/s}}\approx206\ \text{FLOPs/byte}。GQA 比例为 8:1 和 16:1 时,IdecodeI_{\text{decode}} 分别为 8 和 16,均远低于这个临界值,按此估算都处于访存受限区。但提高 GQA 比例仍能使交互计算的算术强度更接近硬件强度。

KV 读取只是交互计算的一部分 IO,实际还包括 QQ、PVPV 的读写,以及写回显存的 scores 和 softmax 输出。计入这些中间数据后,交互 FLOPs 不变,算术强度会比前面的估算更低。提高 GQA 比例同时减少 KV 数据量、增加 scores 和 softmax 输出的数据量,两者对 IO 的影响相反,因此实际算术强度不会随比例同比增长。

线性投影

记 token 总数为 NN(prefill:BSBS;单步 decode:BB)。权重、输入和输出均用 BF16,按权重与输入各读一次、输出写一次估算。以 Q 投影为例,每个 token 的输入有 DD 个元素,输出有 HqdH_qd 个元素,因此输入读取和输出写入合计为 2N(D+Hqd)2N(D+H_qd) bytes,权重读取量为固定的 2DHqd2DH_qd bytes。权重由所有 token 共用,但输入输出的 IO 随 NN 增长。每项投影的成本如下:

投影(每项)FLOPsIO 字节数算术强度(FLOPs/byte)
WQW_Q、WOW_O2NDHqd2NDH_qd2[DHqd+N(D+Hqd)]2[DH_qd+N(D+H_qd)]N1+N/D+N/(Hqd)\dfrac{N}{1+N/D+N/(H_qd)}
WKW_K、WVW_V2NDHkvd2NDH_{kv}d2[DHkvd+N(D+Hkvd)]2[DH_{kv}d+N(D+H_{kv}d)]N1+N/D+N/(Hkvd)\dfrac{N}{1+N/D+N/(H_{kv}d)}

以 Q 投影为例,当 NN 较少,满足 DHqd≫N(D+Hqd)DH_qd\gg N(D+H_qd) 时,IO 主要来自权重读取,算术强度可以近似为:

I=2NDHqd2[DHqd+N(D+Hqd)]≈2NDHqd2DHqd=N.I=\frac{2NDH_qd}{2[DH_qd+N(D+H_qd)]} \approx\frac{2NDH_qd}{2DH_qd}=N.

因此,decode 中增大 batch、prefill 中一次处理更多 token,都能摊薄每个 token 的权重加载成本。对于相同输入,调整 GQA 比例并不改变 NN,所以在 I≈NI\approx N 的近似下,比例不影响投影的算术强度。但随着 NN 增大,输入输出的 IO 逐渐占主导,这个近似不再成立,继续增加 token 数对算术强度的增益也会减弱。

提高比例对总投影计算量的影响有限。比例从 8:1 提高到 16:1,投影 FLOPs 只减少约 5.6%;降低到 4:1 时,投影 FLOPs 增加约 11.1%。这是因为 Q/O 投影不变,随比例变化的 K/V 投影只占四项投影 FLOPs 的 1/91/9。上下文越长,交互计算占比越高,投影计算的这点变化对总 FLOPs 的影响就越小。

另外,提高比例使 dd 减小,K/V 投影的权重和输出数据量随之减少,但输入 hidden states 的读取量不变。因此 IO 的降幅小于 FLOPs 的降幅,算术强度反而下降。

非线性操作

提高比例会增加 softmax 的成本。因为 Q heads 数增多,每行 softmax 要处理的 scores 却没有减少,因此总工作量随 heads 数增长。减小 dd 缩短的是生成 scores 时的点积长度,并不能减少后续的归一化计算。

而 softmax 的单位运算耗时也与矩阵乘法不同。矩阵乘法通常由 Tensor Cores 加速,softmax 的加法、乘法等运算使用 CUDA Cores,指数运算通常使用 SFU。指数运算的吞吐远低于 Tensor Cores 的矩阵乘法吞吐,即使运算次数占比小,也可能占用较多时间。1

GQA 比例选择准则

提高比例有利于缓解 KV 读取瓶颈,却会增加非线性操作的负担。因此,比例的选择应以主要瓶颈为依据:

环节主要耗时调整方向原因
投影K/V 投影提高比例减少权重数据量、输出数据量和计算量,算术强度本身不会提高
投影Q/O 投影调整比例无法直接降低这部分成本矩阵形状与计算量不变
Attention 交互:QK⊤QK^\top、PVPVKV 读取提高比例交互 FLOPs 不变,KV 数据量减少,算术强度提高
Attention 交互:QK⊤QK^\top、PVPV矩阵计算无固定的增减方向总 FLOPs 不变,矩阵形状变化会影响硬件利用率
Softmaxscores 的归一化与读写降低比例减少需要处理的 scores 及对应 IO

性能取舍由服务目标决定。

Head 维度的选择还需兼顾模型任务表现和 kernel 支持。当进一步减小 head 维度不再合适时,还可以改变 KV 的缓存方式。例如 MLA 缓存 K/V 共享的低维表示和额外的 RoPE key,使缓存大小不再直接取决于展开后的 heads 数和维度。

参考资料

  1. Tri Dao, FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision,关于 GEMM 与 softmax 执行吞吐及重叠的分析。 ↩


Share this post on:

Next Post
Prefix Cache 路由:概率调度与平滑切流