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