GQA 吞吐在 batch 128 后为何不再涨:flash-attention 的 pack_gqa 与 num_splits 调优全解
GQA 吞吐在 batch 128 后为何不再涨flash-attention 的 pack_gqa 与 num_splits 调优全解【免费下载链接】flash-attentionFast and memory-efficient exact attention项目地址: https://gitcode.com/GitHub_Trending/fl/flash-attention帮一个 GQA 模型的线上部署做 Flash-Attention 调优时我们碰到过这种情况batch 一路加吞吐却不再线性上涨过了 128 之后平台期个别配置甚至回退。锅不在模型在两个开关pack_gqa和num_splits。这篇讲清楚它们分别在什么条件下该开、该关、该调多大。先看怪象batch 越大反而可能越慢先摆现象不讲原理。在 flash-attention 的 HopperH100前向路径上跑 GQA——也就是 Q 头数多于 KV 头数的配置比如 32 个 Q 头共享 8 个 KV 头——batch 扫描下来通常撞见三段式曲线小 batch18GPU 利用率上不去。一个 batch 的 KV 头就那么几个凑不满全卡的 SM一半计算单元在空转。中 batch32128吞吐稳步爬升这是舒服的区间。大 batch128 以上曲线走平继续加大 batch 时吞吐可能不涨反跌。具体的吞吐数字取决于卡型、序列长度、head dim 和精度以你的环境实测为准别拿别人的表直接抄。共性的是拐点这个形状本身先升、后平、再可能回落。为什么会这样KV 头共享省了显存但桌子会坐满看不懂怪象就别急着拧参数。先把内核在做什么讲透。GQA 本质是一桌人拼一份菜单一个 KV 头要服务H_q / H_k个 Q 头。类比拼桌四个人Q 头坐一桌只点一份菜KV账单按人头摊。省下的就是显存——KV 缓存的大小只跟 KV 头数挂钩跟 Q 头数无关序列越长省得越多。内核层面这个拼桌体现在 hopper/pack_gqa.h 里Q 被摊平成(每组Q头数, 序列位置)的一维行号写回时靠cutlass::FastDivmod把行号拆回组内第几个头、序列第几个位置。这个 divmod 映射就是拼桌的座位表。PackGQA把多个 Q 头塞进同一块 tile默认调度下一个线程块负责1 个 Q 头 × kBlockM 个序列位置kBlockM 是 tile 的 M 维长度Hopper 多数配置为 128 行。问题来了如果seqlen_q很短比如推理时只有一两个 token一块 128 行的 tile 里真正有效的只有几行其余全在空转——照样计费。PackGQA 的做法是让一块 tile 的 128 行由多个 Q 头 × 序列位置拼满行行有效。代价是 Q 的加载要在不同头之间跳跃所以官方注释很诚实hopper/heuristics.h 里写着Heuristic: PackGQA is a bit slower but can help if seqlen_q is small or not near a multiple of kBlockM翻译PackGQA 本身略慢但当seqlen_q很小、或不是 kBlockM 整数倍时它能帮上忙——省下的空转比多花的跳转多。大 batch 为什么反而变慢前向的并行度约等于batch × KV头数 × ceil(seqlen_q / kBlockM)个块。batch 小块数比 SM 数A100 为 108H100 为 132还少SM 吃不满——这是小 batch 怪象的根源。batch 大块数是 SM 的几倍甚至几十倍调度尾部效应开始显形同时单位时间要从 HBM 拉取的 KV 数据量随 batch 线性膨胀带宽顶到天花板后继续加 batch 就不产生收益了——这是大 batch 怪象的根源。所以曲线先升后平不是玄学是两种瓶颈交接的必然结果。怎么调两个参数三条条件规则规则一看序列长度形状定 pack_gqa如果seqlen_q短明显小于 2 × kBlockM或者不是 kBlockM 的整数倍开pack_gqaTrue。decode、增量解码这类场景最常命中这条。如果seqlen_q长且接近 kBlockM 整数倍保持None自动或False。tile 本来就没空行打包白付跳转成本。拿不准先留None。接口默认值就是走启发式自动决策hopper/flash_attn_interface.py 里pack_gqaNone它内部就是按上面那句注释做的判断。规则二SM 吃不满就用 num_splits 补并行num_splits是把每个块的 KV 序列维切成几段、各自独立并行再合并。如果 batch 小、nvidia-smi 里 GPU-Util 明显填不满设num_splits0官方启发式自动选段数或显式给 24。块不够时切 KV 是人为制造并行度的最直接手段代价是多一次flash_attn_combine合并。如果 batch 已经不小回到num_splits1。SM 已经吃饱再切只是增加合并开销还会引入 fp32 的累积缓冲区显存和耗时双输。参考 hopper/flash_attn_interface.py 的 docstringnum_splits1不切、1按段数切、0走启发式。规则三一个最小起手配置from flash_attn import flash_attn_func # Hopper (FA3) 入口见 hopper/ out flash_attn_func( q, k, v, causalTrue, pack_gqaTrue, # seqlen_q 短 / 非 kBlockM 整数倍时 num_splits0, # 0 启发式自动; 1 不切; 1 切 N 段 )改动原则一次只动一个变量其余保持默认测完再动下一个。自检清单动手调之前过一遍调参前把这张表跑完能省掉大部分弯路正确性前提Q 头数必须能被 KV 头数整除接口 docstring 明确要求不满足直接报错先确认配置合法。序列长度检查seqlen_q对 kBlockM典型 128取余是否为 0不是 →pack_gqa优先开。瓶颈定位边跑边看nvidia-smi或nvidia-smi dmon两个数GPU-Util 长期低于 70% → 并行度不足 → 用num_splits补显存带宽Mem-Util已接近饱和 → 参数再调也快不过去了出路是降 batch、缩序列、换精度而不是继续拧pack_gqa。基线先行先用默认组合pack_gqaNone, num_splits1跑一遍当基线之后每次只改一处。扫出你的拐点batch 按 8 → 16 → 32 → 64 → 128 → 256 扫一遍记吞吐峰值就是你的最优 batch而不是越大越好。对应的决策路径收敛成一句话版序列短或不成 tile 整数倍→pack_gqaTrue否则维持自动GPU-Util 填不满→num_splits0或 24填得满 →num_splits1带宽已打满→ 停止调参改 batch 或精度⚠️ 最后提醒一句以上所有阈值128、0.7 等是方向性参考不是铁律。同一张卡换 head dim、换 causal/非 causal、换 varlen 接口拐点都会挪。以实测为准永远以实测为准。【免费下载链接】flash-attentionFast and memory-efficient exact attention项目地址: https://gitcode.com/GitHub_Trending/fl/flash-attention创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考