優(yōu) 4 步速查)
為什么 GQA 推理吞吐先升后降Flash-Attention 批量大小調(diào)優(yōu) 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的處理上。下面按診斷 → 根因 → 修復(fù) → 驗證四步講清 GQA 批量大小優(yōu)化和 pack_gqa / num_splits 的 Flash-Attention 調(diào)參。?? 癥狀H100 上 GQA 吞吐瓶頸排查復(fù)現(xiàn)很簡單固定 H_q32、H_k8、序列長度 2K只改 batch size。Batch size吞吐Tokens/s延遲ms6428,40045.112831,20082.725626,800192.3吞吐在 128 附近見頂256 時掉 15%。背后是兩個矛盾的拉扯內(nèi)存帶寬 vs 計算并行度batch 小時 SM 填不滿利用率低batch 大時 KV 緩存占滿 HBM 帶寬計算干等數(shù)據(jù)。線程塊調(diào)度 vs SM 承載H100 有 132 個 SMbatch 256 × KV 頭數(shù)對應(yīng)的線程塊遠超 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 的頭數(shù)必須能被 KV 頭數(shù)整除。這個省內(nèi)存的結(jié)構(gòu)帶了兩筆隱性稅batch 小時每個 KV 分組內(nèi)的活躍查詢頭太少線程塊填不滿 SM序列短、不是線程塊大小的整數(shù)倍時塊內(nèi)尾部浪費被放大。PackGQA 把同一組的多個查詢頭打包進一個線程塊攤薄 KV 讀取開銷機制見 hopper/pack_gqa.h。但倉庫里的啟發(fā)式說得非常誠實hopper/heuristics.h 注釋——PackGQA is a bit slower but can help if seqlen_q is small or not near a multiple of kBlockM。也就是說它用少量計算效率換內(nèi)存效率不是白拿的優(yōu)化batch 一增大這筆交換的賬就虧回來了必須配合num_splits把 K/V 維度切開降低單次訪存量。? 修復(fù)pack_gqa 與 num_splits 速查表Batch sizepack_gqanum_splits一句話理由≤ 32True1小 batch 靠打包填滿 SM33 – 128True1吞吐峰值區(qū)長序列可試 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), )按速查表調(diào)整后同一 H100 機器上的實測Batch調(diào)參前auto調(diào)參后按表配置12829,500 Tokens/s31,2005.8%25622,100 Tokens/s26,80021.3%? 驗證與進階用nvidia-smi盯兩個數(shù)GPU-Util和Mem-Util同時落在 70%–90% 就是健康區(qū)間——GPU-Util 高而 Mem-Util 低試試開pack_gqa反過來加num_splits。進階方向一句話帶過H100 上開 FP8e4m3能直接砍一半帶寬壓力長序列8K配小 batch32、短序列512配大 batch128的動態(tài)調(diào)度能進一步抹平波動。batch 64–128 是 H100 上的吞吐峰值區(qū)256 時改num_splits4可把 -15% 拉回 21%。更多參數(shù)說明見 README.md 與 Hopper 接口文檔?!久赓M下載鏈接】flash-attentionFast and memory-efficient exact attention項目地址: https://gitcode.com/GitHub_Trending/fl/flash-attention創(chuàng)作聲明:本文部分內(nèi)容由AI輔助生成(AIGC),僅供參考