優(yōu)全解)
GQA 吞吐在 batch 128 后為何不再漲flash-attention 的 pack_gqa 與 num_splits 調(diào)優(yōu)全解【免費下載鏈接】flash-attentionFast and memory-efficient exact attention項目地址: https://gitcode.com/GitHub_Trending/fl/flash-attention幫一個 GQA 模型的線上部署做 Flash-Attention 調(diào)優(yōu)時我們碰到過這種情況batch 一路加吞吐卻不再線性上漲過了 128 之后平臺期個別配置甚至回退。鍋不在模型在兩個開關(guān)pack_gqa和num_splits。這篇講清楚它們分別在什么條件下該開、該關(guān)、該調(diào)多大。先看怪象batch 越大反而可能越慢先擺現(xiàn)象不講原理。在 flash-attention 的 HopperH100前向路徑上跑 GQA——也就是 Q 頭數(shù)多于 KV 頭數(shù)的配置比如 32 個 Q 頭共享 8 個 KV 頭——batch 掃描下來通常撞見三段式曲線小 batch18GPU 利用率上不去。一個 batch 的 KV 頭就那么幾個湊不滿全卡的 SM一半計算單元在空轉(zhuǎn)。中 batch32128吞吐穩(wěn)步爬升這是舒服的區(qū)間。大 batch128 以上曲線走平繼續(xù)加大 batch 時吞吐可能不漲反跌。具體的吞吐數(shù)字取決于卡型、序列長度、head dim 和精度以你的環(huán)境實測為準(zhǔn)別拿別人的表直接抄。共性的是拐點這個形狀本身先升、后平、再可能回落。為什么會這樣KV 頭共享省了顯存但桌子會坐滿看不懂怪象就別急著擰參數(shù)。先把內(nèi)核在做什么講透。GQA 本質(zhì)是一桌人拼一份菜單一個 KV 頭要服務(wù)H_q / H_k個 Q 頭。類比拼桌四個人Q 頭坐一桌只點一份菜KV賬單按人頭攤。省下的就是顯存——KV 緩存的大小只跟 KV 頭數(shù)掛鉤跟 Q 頭數(shù)無關(guān)序列越長省得越多。內(nèi)核層面這個拼桌體現(xiàn)在 hopper/pack_gqa.h 里Q 被攤平成(每組Q頭數(shù), 序列位置)的一維行號寫回時靠cutlass::FastDivmod把行號拆回組內(nèi)第幾個頭、序列第幾個位置。這個 divmod 映射就是拼桌的座位表。PackGQA把多個 Q 頭塞進同一塊 tile默認(rèn)調(diào)度下一個線程塊負(fù)責(zé)1 個 Q 頭 × kBlockM 個序列位置kBlockM 是 tile 的 M 維長度Hopper 多數(shù)配置為 128 行。問題來了如果seqlen_q很短比如推理時只有一兩個 token一塊 128 行的 tile 里真正有效的只有幾行其余全在空轉(zhuǎn)——照樣計費。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 本身略慢但當(dāng)seqlen_q很小、或不是 kBlockM 整數(shù)倍時它能幫上忙——省下的空轉(zhuǎn)比多花的跳轉(zhuǎn)多。大 batch 為什么反而變慢前向的并行度約等于batch × KV頭數(shù) × ceil(seqlen_q / kBlockM)個塊。batch 小塊數(shù)比 SM 數(shù)A100 為 108H100 為 132還少SM 吃不滿——這是小 batch 怪象的根源。batch 大塊數(shù)是 SM 的幾倍甚至幾十倍調(diào)度尾部效應(yīng)開始顯形同時單位時間要從 HBM 拉取的 KV 數(shù)據(jù)量隨 batch 線性膨脹帶寬頂?shù)教旎ò搴罄^續(xù)加 batch 就不產(chǎn)生收益了——這是大 batch 怪象的根源。所以曲線先升后平不是玄學(xué)是兩種瓶頸交接的必然結(jié)果。怎么調(diào)兩個參數(shù)三條條件規(guī)則規(guī)則一看序列長度形狀定 pack_gqa如果seqlen_q短明顯小于 2 × kBlockM或者不是 kBlockM 的整數(shù)倍開pack_gqaTrue。decode、增量解碼這類場景最常命中這條。如果seqlen_q長且接近 kBlockM 整數(shù)倍保持None自動或False。tile 本來就沒空行打包白付跳轉(zhuǎn)成本。拿不準(zhǔn)先留None。接口默認(rèn)值就是走啟發(fā)式自動決策hopper/flash_attn_interface.py 里pack_gqaNone它內(nèi)部就是按上面那句注釋做的判斷。規(guī)則二SM 吃不滿就用 num_splits 補并行num_splits是把每個塊的 KV 序列維切成幾段、各自獨立并行再合并。如果 batch 小、nvidia-smi 里 GPU-Util 明顯填不滿設(shè)num_splits0官方啟發(fā)式自動選段數(shù)或顯式給 24。塊不夠時切 KV 是人為制造并行度的最直接手段代價是多一次flash_attn_combine合并。如果 batch 已經(jīng)不小回到num_splits1。SM 已經(jīng)吃飽再切只是增加合并開銷還會引入 fp32 的累積緩沖區(qū)顯存和耗時雙輸。參考 hopper/flash_attn_interface.py 的 docstringnum_splits1不切、1按段數(shù)切、0走啟發(fā)式。規(guī)則三一個最小起手配置from flash_attn import flash_attn_func # Hopper (FA3) 入口見 hopper/ out flash_attn_func( q, k, v, causalTrue, pack_gqaTrue, # seqlen_q 短 / 非 kBlockM 整數(shù)倍時 num_splits0, # 0 啟發(fā)式自動; 1 不切; 1 切 N 段 )改動原則一次只動一個變量其余保持默認(rèn)測完再動下一個。自檢清單動手調(diào)之前過一遍調(diào)參前把這張表跑完能省掉大部分彎路正確性前提Q 頭數(shù)必須能被 KV 頭數(shù)整除接口 docstring 明確要求不滿足直接報錯先確認(rèn)配置合法。序列長度檢查seqlen_q對 kBlockM典型 128取余是否為 0不是 →pack_gqa優(yōu)先開。瓶頸定位邊跑邊看nvidia-smi或nvidia-smi dmon兩個數(shù)GPU-Util 長期低于 70% → 并行度不足 → 用num_splits補顯存帶寬Mem-Util已接近飽和 → 參數(shù)再調(diào)也快不過去了出路是降 batch、縮序列、換精度而不是繼續(xù)擰pack_gqa?;€先行先用默認(rèn)組合pack_gqaNone, num_splits1跑一遍當(dāng)基線之后每次只改一處。掃出你的拐點batch 按 8 → 16 → 32 → 64 → 128 → 256 掃一遍記吞吐峰值就是你的最優(yōu) batch而不是越大越好。對應(yīng)的決策路徑收斂成一句話版序列短或不成 tile 整數(shù)倍→pack_gqaTrue否則維持自動GPU-Util 填不滿→num_splits0或 24填得滿 →num_splits1帶寬已打滿→ 停止調(diào)參改 batch 或精度?? 最后提醒一句以上所有閾值128、0.7 等是方向性參考不是鐵律。同一張卡換 head dim、換 causal/非 causal、換 varlen 接口拐點都會挪。以實測為準(zhǔn)永遠以實測為準(zhǔn)。【免費下載鏈接】flash-attentionFast and memory-efficient exact attention項目地址: https://gitcode.com/GitHub_Trending/fl/flash-attention創(chuàng)作聲明:本文部分內(nèi)容由AI輔助生成(AIGC),僅供參考