为什么 GQA 推理吞吐先升后降:Flash-Attention 批量大小调优 4 步速查
2026/9/5 18:29:54 网站建设 项目流程

为什么 GQA 推理吞吐先升后降:Flash-Attention 批量大小调优 4 步速查

【免费下载链接】flash-attentionFast and memory-efficient exact attention项目地址: https://gitcode.com/GitHub_Trending/fl/flash-attention

线上 LLM 推理集群把 batch 从 64 提到 256,吞吐不升反降 15%。问题出在 Flash-Attention 对 Grouped-Query Attention(GQA)的处理上。下面按"诊断 → 根因 → 修复 → 验证"四步,讲清 GQA 批量大小优化和 pack_gqa / num_splits 的 Flash-Attention 调参。

⚡️ 症状:H100 上 GQA 吞吐瓶颈排查

复现很简单:固定 H_q=32、H_k=8、序列长度 2K,只改 batch size。

Batch size吞吐(Tokens/s)延迟(ms)
6428,40045.1
12831,20082.7
25626,800192.3

吞吐在 128 附近见顶,256 时掉 15%。背后是两个矛盾的拉扯:

  • 内存带宽 vs 计算并行度:batch 小时 SM 填不满,利用率低;batch 大时 KV 缓存占满 HBM 带宽,计算干等数据。
  • 线程块调度 vs SM 承载:H100 有 132 个 SM,batch 256 × KV 头数对应的线程块远超 SM 承载,块频繁切换就像 CPU 上下文切换,开销直接吃掉并行收益。

🔍 根因:GQA 分组机制下的"隐性税"

GQA 让 $H_q$ 个查询头共享 $H_k$ 个 KV 头,一个 KV 头被 $H_q / H_k$ 个查询头摊薄。最小例子:$H_q=6, H_k=2$ 时,前 3 个 Q 头共用 KV 头 0,后 3 个共用 KV 头 1。接口约束写在 hopper/flash_attn_interface.py:"Q 的头数必须能被 KV 头数整除。"

这个省内存的结构带了两笔隐性税:

  1. batch 小时,每个 KV 分组内的活跃查询头太少,线程块填不满 SM;
  2. 序列短、不是线程块大小的整数倍时,块内尾部浪费被放大。

PackGQA 把同一组的多个查询头打包进一个线程块,摊薄 KV 读取开销,机制见 hopper/pack_gqa.h。但仓库里的启发式说得非常诚实:hopper/heuristics.h 注释——"PackGQA is a bit slower but can help if seqlen_q is small or not near a multiple of kBlockM"。也就是说它用少量计算效率换内存效率,不是白拿的优化;batch 一增大,这笔交换的账就亏回来了,必须配合num_splits把 K/V 维度切开,降低单次访存量。

🛠️ 修复:pack_gqa 与 num_splits 速查表

Batch sizepack_gqanum_splits一句话理由
≤ 32True1小 batch 靠打包填满 SM
33 – 128True1吞吐峰值区,长序列可试 2
129 – 256False4带宽受限,拆分降单次访存
> 256False4 – 8拆分 + combine,注意显存
from flash_attn import flash_attn_func def pick_params(batch_size: int): if batch_size <= 128: return dict(pack_gqa=True, num_splits=1) return dict(pack_gqa=False, num_splits=4) out = flash_attn_func( q, k, v, softmax_scale=1.0 / (q.shape[-1] ** 0.5), causal=True, **pick_params(batch_size), )

按速查表调整后,同一 H100 机器上的实测:

Batch调参前(auto)调参后(按表配置)
12829,500 Tokens/s31,200(+5.8%)
25622,100 Tokens/s26,800(+21.3%)

✅ 验证与进阶

nvidia-smi盯两个数:GPU-UtilMem-Util同时落在 70%–90% 就是健康区间——GPU-Util 高而 Mem-Util 低,试试开pack_gqa;反过来,加num_splits。进阶方向一句话带过:H100 上开 FP8(e4m3)能直接砍一半带宽压力;长序列(8K)配小 batch(32)、短序列(512)配大 batch(128)的动态调度,能进一步抹平波动。

batch 64–128 是 H100 上的吞吐峰值区,256 时改num_splits=4可把 -15% 拉回 +21%。更多参数说明见 README.md 与 Hopper 接口文档。

【免费下载链接】flash-attentionFast and memory-efficient exact attention项目地址: https://gitcode.com/GitHub_Trending/fl/flash-attention

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询