
为什么 GQA 有时快有时慢Flash-Attention 调参 pack_gqa 与 num_splits 完整指南【免费下载链接】flash-attentionFast and memory-efficient exact attention项目地址: https://gitcode.com/GitHub_Trending/fl/flash-attention这篇文章围绕 GQAGrouped-Query Attention把多个 Query 头共享同一份 K/V 头的注意力变体展开基于 Flash-Attention 这个快速且省内存的精确注意力开源实现FlashAttention / FlashAttention-2 / 3 / 4讲清楚为什么 GQA 的性能会随 batch size 起伏不定以及pack_gqa和num_splits这两个参数到底该怎么配。先说结论GQA 省的是显存花的是调度上的心思。头数配对了性能曲线是另一回事。误区一K/V 头数调小就完事了不用管别的GQA 的用法很直接Q 有 32 个头K、V 给 8 个每个 KV 头服务 4 个 Q 头hopper/flash_attn_interface.py 的接口文档里就写着Q 的头数必须能被 KV 头数整除。这一步确实能砍掉 75% 的 KV 缓存显存但它只解决了装得下的问题没解决跑得快的问题。真正决定速度的是下面两个问题每个线程块block一次算了多少有效的 Q 行整个 GPU 上的 SM流式多处理器有没有被喂饱。batch 大小一变这两个量的变化趋势完全不同性能曲线自然就不是单调的了。误区二pack_gqaTrue是免费午餐常开就好pack_gqa是 Hopper 上 FlashAttention-3 引入的开关核心实现在 hopper/pack_gqa.h。它做的事一句话概括把共享同一个 KV 头的 4 个 Q 头拼进同一个计算块里算KV 只从显存搬一次。听起来很美好但代价是拼进去的 Q 头越多块tile在 M 方向上被撑得越大。如果你的序列本身就很长本来一个块就能塞满那么拼头之后多出来的行大部分是 padding白算了。官方源码里写得很直白hopper/heuristics.h// Heuristic: PackGQA is a bit slower but can help if // seqlen_q is small or not near a multiple of kBlockM float nopack_gqa_efficiency float(seqlen_q) / float(round_up(seqlen_q, blockM)); float pack_gqa_efficiency float(seqlen_q * qhead_per_khead) / float(round_up(seqlen_q * qhead_per_khead, blockM)); return nopack_gqa_efficiency 0.9 * pack_gqa_efficiency;也就是说只有当不拼时的填充浪费明显大于拼时差到 0.9 倍以上才默认启用。传None让库自己判断通常是最稳的选择varlen变长序列场景则一律拼因为每个样本长度不齐不拼浪费更狠。所以pack_gqa不是布尔开关而是一个按序列长度分档的决策。误区三batch 越大吞吐越高线性放大很多人默认 batch 翻倍、tokens/s 也翻倍GQA 下恰恰不是这样。关键在 SM 占用度。GPU 上的并行度 ≈batch × KV头数 × 块数。batch 小的时候比如 batch1、H_k8总共可能只有 8 个块而 H100 有 132 个 SM大部分 SM 在干等——这时再大的算力都是空转。这就是为什么官方提供num_splits把一个 KV 序列切成多段不同段分给不同块并行算算完再做一次归并reduce相当于凭空多造出几个并行任务。但 batch 大到一定程度后另一个方向的问题出现任务已经多到 SM 忙不过来再拆 split 只是徒增 HBM 读写每多一个 split中间结果就多一次落盘再读回。于是吞吐量在某个 batch 上见顶再往上反着走——参考文章里batch 256 比 64 还慢的观感根源就在这而不是KV 缓存占带宽那么简单。误区四num_splits调得越高越好看 hopper/heuristics.h 里的num_splits_heuristic逻辑相当老练若块数已经够 80% 的 SM就不拆num_splits1若拆的段太少num_n_blocks 4也不拆否则遍历候选 split 数算每档的波次填充效率取能达到最优效率 85% 的最小拆分数。注意最后一条它不选效率最高的那档而是够用就行的最小拆分因为拆分本身有通信与显存成本。手动调参时方向应该和它一致——先试 1不够再加别直接拉满。落地一张能直接抄的决策表在 hopper/flash_attn_interface.py 的flash_attn_func里两个参数都标着 Can be tuned for speedout flash_attn_func( q, k, v, causalTrue, pack_gqaNone, # 默认交给启发式多数场景别动 num_splits0, # 0 自动选拆分数 )场景pack_gqanum_splits原因短序列2K、小 batch交给默认通常为开1Q 块拼不满拼头补行SM 吃不饱就少拆长序列≥4K、大 batch交给默认通常为关1拼头引入 padding白算batch1、H_k 很小如 MQA开2~4并行任务太少必须靠拆分喂饱 SMvarlen 变长推理开自动长度不齐时不拼浪费更大KV 特别长、单头 KV 超 L2 容量按默认自动1源码里对size_one_kv_head 50MB有专门的拆分逻辑⚡ 一个实用的验证方法固定输入分别跑num_splits1,2,4和pack_gqaNone/True/False用nsys或torch.cuda.Event计时谁快用谁——这个组合空间很小十分钟能扫完。最佳实践检查清单✅默认先不碰pack_gqaNonenum_splits0自动已经内置了官方启发式80% 的场景这就是最优解。✅短序列/小 batch 优先怀疑并行度不够先看batch × H_k是否远小于 SM 数不够就调num_splits。✅长序列优先怀疑padding 浪费这时强行开pack_gqa可能反而变慢用上面 2×3 的扫描确认。✅varlen 场景不要手动关 PackGQA源码对varlen_q是直接返回true的。✅别追求极端值num_splits的收益曲线是先快后平再掉头选够用的最小值。✅用官方基准校准跑 hopper/test_flash_attn.py 和 benchmarks/bench_sm90.py 确认改动方向而不是只看主观感受。调参的本质就一句话batch 决定任务够不够多序列长度决定每个块实不实pack_gqa和num_splits分别回答这两个问题。想清楚你在哪个格子参数基本就不用猜了。【免费下载链接】flash-attentionFast and memory-efficient exact attention项目地址: https://gitcode.com/GitHub_Trending/fl/flash-attention创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考