Flash-Attention GQA 性能调优pack_gqa 与 num_splits 两个开关怎么选【免费下载链接】flash-attentionFast and memory-efficient exact attention项目地址: https://gitcode.com/GitHub_Trending/fl/flash-attentionFlash-Attention 的 GQAGrouped-Query Attention多个查询头共享少量 KV 头推理吞吐量并不总随 batch 变大而变好。一个典型现象示例配置$H_q32$、$H_k8$、seqlen_q1 的 decodebatch 从 8 加到 64 时吞吐还能涨但 SM 占用距离打满仍差得远batch 继续翻倍每 token 吞吐的增速却明显放缓。很多人第一反应是把 batch 堆大但真正的调优入口其实是 Python 接口上两个参数pack_gqa和num_splits。下面以仓库 hopper/ 目录的代码为准讲清楚这两个参数怎么选、怎么验证。为什么 GQA 推理会喂不饱 SMFlash-Attention 前向 kernel 按块划分工作每个线程块CTA负责一个 (batch, KV 头) 组合处理其中一段 kBlockM × kBlockN 的序列块。GPU 能并行的总量取决于 (batch, KV 头) 组合的数量。MHA$H_q H_k$下 batch32 就有 32×321024 个工作单元SM 很容易填满而 GQA 里 KV 头少$H_k8$ 时只有 256 个单元decode 阶段 seqlen_q1 时每个单元只有 1 行 querytile 的 M 方向却按 kBlockM如 128开整块里 127 行都是 padding。batch 越小、序列越短浪费越严重。但这里有个坑两种浪费要往不同方向补。单元数不够要靠 num_splits把 KV 切成多段、增加并行单元解决行被 padding要靠 pack_gqa把同组的多个查询头塞进同一个 tile 的 M 方向解决。只堆 batch 两个问题都没解决。两个开关在 kernel 层面各做了什么pack_gqa 由 hopper/pack_gqa.h 实现把 Q 从 (seqlen_q, $H_q$) 重排成 (seqlen_q × qhead_per_khead, $H_k$)让同一个 KV 头下的 $qhead_per_khead$ 个查询头在 tile 的 M 方向上连续。头映射关系用cutlass::FastDivmod预计算、指针在 warp 内用__shfl_sync广播代价是每组行多一次 divmod——所以 hopper/heuristics.h 里的官方注释明确写着 PackGQA is a bit slower只有 tile 填充收益盖过这点开销时才值得开。num_splits 则把 KV 的序列方向切成 num_splits 段每段独立算出部分注意力和 logsumexp最后再 reduce 合并。代价是多一轮部分结果的 HBM 读写收益是并行单元数乘以 num_splits。pack_gqa 怎么选让 0.9 阈值替你决定Python 接口默认传pack_gqaNone此时 C 端调用get_pack_gqa自动决策核心是比较 padding 前后的行效率// hopper/heuristics.h inline bool should_pack_gqa(bool varlen_q, int seqlen_q, int qhead_per_khead, int blockM) { if (varlen_q) return true; 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; }效率 有效行数 / 向上取整到 tile 后的行数。只有不打包的效率低于打包的 90% 时才开 pack。以上面示例为例seqlen_q1、每组 4 个查询头、kBlockM128不打包效率 1/128打包 4/128差 4 倍必开反过来若 seqlen_q 本身是 128 的整数倍两者都是 1.0就不开——此时白付 divmod 开销。三条使用规则可以直接带走varlencu_seqlens 输入无条件打包因为 ragged 场景算不出真实 seqlen打包更稳seqlen_q 小或不是 tile 整数倍时自动决策几乎都会选打包要手动锁定就传 True/False在 hopper/flash_attn_interface.py 里用户指定值优先于启发式。num_splits 怎么选先问SM 填不填得满num_splits0表示交给自动启发式逻辑分三层总块数≈ batch × $H_k$达到 SM 数的 80% 以上说明 SM 基本填满直接返回 1 段。例外是超长 KV非因果、单个 KV 头超过约 50MB L2 且查询量足够时按 L2 比例切段以提升复用。H100 有 132 个 SM示例口径$H_k8$ 时 batch 到 20 左右就够1 段阈值了KV 序列方向本身只有 ≤4 个块时如 hdim128、seqlen_k512没有可切空间返回 1其余情况枚举 1 到 max_splits 的每段数算波形效率 n_waves/ceil(n_waves)即最后一波 SM 的利用率取最大值在达到最优 85% 的候选里挑最小的段数——避免理论峰值多一段、实际多一轮 HBM 流量。换句话说调优原则是大批量长序列基本不用动小 batch decode 才是自动决策发挥空间最大的地方手动干预通常没必要。确要介入时典型调用长这样from flash_attn_interface import flash_attn_with_kvcache out flash_attn_with_kvcache( q, k_cache, v_cache, kk_new, vv_new, cache_seqlenscache_seqlens, causalTrue, num_splits0, # 0 启发式自动决定 pack_gqaNone, # None 自动决定 )两参数组合速查示例口径$H_k8$、H100输入场景pack_gqanum_splitsvarlen 变长序列True强制0自动decode 小 batchseqlen_q1自动通常 True0自动prefill 且 batch×$H_k$ ≥ 0.8×SM 数保持默认多半不打包1非因果 单个 KV 头体积 约 50MB保持默认按 L2 比例切怎么验证10 行代码做对照实验别拿单一配置下结论。挑你业务里真实的 batch 和序列长度对 (pack_gqa, num_splits) 做对照import torch from flash_attn_interface import flash_attn_func for pack_gqa in (None, True, False): for num_splits in (1, 2, 4): ev0, ev1 torch.cuda.Event(True), torch.cuda.Event(True) ev0.record() for _ in range(50): flash_attn_func(q, k, v, causalTrue, pack_gqapack_gqa, num_splitsnum_splits) ev1.record(); torch.cuda.synchronize() print(pack_gqa, num_splits, ev0.elapsed_time(ev1) / 50, ms)两个注意点第一同一份张量布局多轮计时并充分 warmup首次调用会触发模板实例化编译混进延迟里会严重失真第二胜出组合还应满足nvidia-smi或 Nsight 里 SM 占用 70% 以上否则可能只是别的组合更差。做正式回归可以直接用 benchmarks/bench_sm90.py它支持按 batch、序列长度、前向/反向方向扫参。上图是不同序列长度下 H100 的前向/反向吞吐趋势可看出注意力计算对序列长度很敏感用同样的方法扫 batch就能定位你自己的拐点。常见坑点与落地清单三个容易踩的坑认为 pack_gqaTrue 一定更快官方注释写明它略慢序列长且 tile 对齐时强制开启只会白付 divmodnum_splits 拍脑袋调大每多一段就多一轮部分结果的 HBM 读写自动启发式用 85% 阈值已经权衡过手动给大值大概率更差把两个参数当独立开关Split 路径的代码注释表明切分开启时 PackGQA 始终生效用于减少编译组合手动设 num_splits1 且 pack_gqaFalse实际跑的可能不是你以为的组合。落地清单先用默认值pack_gqaNone、num_splits0跑一遍基线记录每个配置的延迟小 batch decode 场景确认自动决策实际开启了 packseqlen_q1 时几乎必为 True大批量 prefillbatch×$H_k$ ≥ 0.8×SM 数保持 num_splits1别动实测与基线偏差超过 10% 时才手动覆盖 True/False 或具体段数并跑上面的对照实验调优后复查 SM 占用仍低于 70% 说明瓶颈在 batch 与 seqlen 的组合而不是这两个参数。【免费下载链接】flash-attentionFast and memory-efficient exact attention项目地址: https://gitcode.com/GitHub_Trending/fl/flash-attention创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考