ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

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

为什么 GQA 推理吞吐先升后降:Flash-Attention 批量大小调优 4 步速查 为什么 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 AttentionGQA的处理上。下面按诊断 → 根因 → 修复 → 验证四步讲清 GQA 批量大小优化和 pack_gqa / num_splits 的 Flash-Attention 调参。⚡️ 症状H100 上 GQA 吞吐瓶颈排查复现很简单固定 H_q32、H_k8、序列长度 2K只改 batch size。Batch size吞吐Tokens/s延迟ms6428,40045.112831,20082.725626,800192.3吞吐在 128 附近见顶256 时掉 15%。背后是两个矛盾的拉扯内存带宽 vs 计算并行度batch 小时 SM 填不满利用率低batch 大时 KV 缓存占满 HBM 带宽计算干等数据。线程块调度 vs SM 承载H100 有 132 个 SMbatch 256 × KV 头数对应的线程块远超 SM 承载块频繁切换就像 CPU 上下文切换开销直接吃掉并行收益。 根因GQA 分组机制下的隐性税GQA 让 $H_q$ 个查询头共享 $H_k$ 个 KV 头一个 KV 头被 $H_q / H_k$ 个查询头摊薄。最小例子$H_q6, H_k2$ 时前 3 个 Q 头共用 KV 头 0后 3 个共用 KV 头 1。接口约束写在 hopper/flash_attn_interface.pyQ 的头数必须能被 KV 头数整除。这个省内存的结构带了两笔隐性税batch 小时每个 KV 分组内的活跃查询头太少线程块填不满 SM序列短、不是线程块大小的整数倍时块内尾部浪费被放大。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 靠打包填满 SM33 – 128True1吞吐峰值区长序列可试 2129 – 256False4带宽受限拆分降单次访存 256False4 – 8拆分 combine注意显存from flash_attn import flash_attn_func def pick_params(batch_size: int): if batch_size 128: return dict(pack_gqaTrue, num_splits1) return dict(pack_gqaFalse, num_splits4) out flash_attn_func( q, k, v, softmax_scale1.0 / (q.shape[-1] ** 0.5), causalTrue, **pick_params(batch_size), )按速查表调整后同一 H100 机器上的实测Batch调参前auto调参后按表配置12829,500 Tokens/s31,2005.8%25622,100 Tokens/s26,80021.3%✅ 验证与进阶用nvidia-smi盯两个数GPU-Util和Mem-Util同时落在 70%–90% 就是健康区间——GPU-Util 高而 Mem-Util 低试试开pack_gqa反过来加num_splits。进阶方向一句话带过H100 上开 FP8e4m3能直接砍一半带宽压力长序列8K配小 batch32、短序列512配大 batch128的动态调度能进一步抹平波动。batch 64–128 是 H100 上的吞吐峰值区256 时改num_splits4可把 -15% 拉回 21%。更多参数说明见 README.md 与 Hopper 接口文档。【免费下载链接】flash-attentionFast and memory-efficient exact attention项目地址: https://gitcode.com/GitHub_Trending/fl/flash-attention创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表