顯存測量與優(yōu)化:從賬單拆解到預算決策實戰(zhàn))
1. 訓練側(cè)顯存測量到底在測什么顯存優(yōu)化這件事很多人一上來就想著怎么省結果省了半天發(fā)現(xiàn)根本沒省到點子上。問題出在哪出在沒搞清楚顯存到底被誰吃掉了。訓練側(cè)的顯存測量核心目標就一個把顯存賬單拆開看清楚每一筆開銷的去向然后才能做預算決策。我見過太多人拿著nvidia-smi看一眼顯存占用發(fā)現(xiàn)快滿了就開始慌然后盲目上梯度檢查點、盲目降batch size最后訓練速度掉了一半顯存也沒省下多少。這種做法的問題在于nvidia-smi看到的只是一個總數(shù)它不會告訴你這8GB里有3GB是模型參數(shù)、2GB是優(yōu)化器狀態(tài)、1.5GB是激活值、剩下的是碎片和臨時緩沖區(qū)。你不知道錢花在哪就沒辦法做預算。訓練側(cè)顯存測量要回答的問題很具體模型參數(shù)占多少、梯度占多少、優(yōu)化器狀態(tài)占多少、激活值占多少、臨時緩沖區(qū)占多少。這五塊加起來才是真實賬單。而且這五塊的比例關系會隨著模型規(guī)模、batch size、序列長度、精度策略的變化而劇烈變化。一個7B模型在FP16下參數(shù)占14GB但如果用Adam優(yōu)化器優(yōu)化器狀態(tài)就要占28GBFP32的一階動量和二階動量各14GB再加上梯度14GB光這三項就56GB了還沒算激活值。所以為什么大家說7B模型全量微調(diào)至少要80GB顯存賬就是這么算出來的。測量方法上最直接的是用PyTorch的顯存分析工具。torch.cuda.memory_allocated()能拿到當前分配的顯存torch.cuda.max_memory_allocated()能拿到峰值。但這兩個數(shù)字只反映PyTorch分配器層面的情況不包括CUDA上下文本身占用的那幾百MB。更細的拆解需要用torch.cuda.memory_summary()它會按分配塊大小分類列出。不過這個輸出比較原始我一般會自己寫個hook在模型forward前后、backward前后分別打點算出每個階段的增量。還有一個容易被忽略的點顯存碎片。PyTorch的緩存分配器會預留一些顯存不還給系統(tǒng)導致nvidia-smi看到的占用比memory_allocated()高不少。這個差值在長時間訓練中會逐漸增大尤其是當你有動態(tài)shape的輸入時。測量的時候要把這個差值也記錄下來否則你按memory_allocated()做的預算到了實際訓練時就會OOM。注意測量一定要在真實訓練循環(huán)里做不能只跑一個forward就完事。因為激活值的峰值出現(xiàn)在backward階段只測forward會嚴重低估。2. 顯存賬單的五大部分與計算邏輯2.1 模型參數(shù)與梯度的顯存占用模型參數(shù)的顯存占用最好算參數(shù)量乘以每個參數(shù)的字節(jié)數(shù)。FP32是4字節(jié)FP16和BF16是2字節(jié)INT8是1字節(jié)。一個7B模型在FP16下就是7B乘2等于14GB。梯度占用的字節(jié)數(shù)和參數(shù)一致因為梯度需要和參數(shù)同樣的精度來保證更新時不丟信息。所以FP16訓練時參數(shù)加梯度就是28GB。但這里有個坑很多人以為用FP16訓練參數(shù)就是FP16。實際上PyTorch的AMP自動混合精度會保留一份FP32的master weight。也就是說參數(shù)實際上占了兩份一份FP16用于前向和反向計算一份FP32用于優(yōu)化器更新。這樣算下來7B模型的參數(shù)占用是14GBFP16加28GBFP32 master總共42GB。梯度也是FP16一份但優(yōu)化器更新時需要FP32的梯度所以梯度實際占用也是14GBFP16加28GBFP32又是42GB。這就是為什么AMP能省顯存但省不了太多——它省的是激活值和部分計算緩沖區(qū)參數(shù)和優(yōu)化器狀態(tài)的大頭省不掉。2.2 優(yōu)化器狀態(tài)的顯存黑洞優(yōu)化器狀態(tài)是顯存占用里最容易被低估的部分。以Adam為例它需要為每個參數(shù)維護一階動量m和二階動量v都是FP32。所以優(yōu)化器狀態(tài)的顯存等于參數(shù)量乘以4字節(jié)乘以2也就是參數(shù)量乘以8字節(jié)。7B模型就是56GB。加上前面的參數(shù)和梯度已經(jīng)98GB了。這就是為什么全量微調(diào)7B模型至少需要8張80GB的卡來做數(shù)據(jù)并行——單卡根本放不下。AdamW稍微好一點它把權重衰減和梯度更新解耦了但動量部分還是一樣的。Adafactor通過分解二階動量矩陣來省顯存能把優(yōu)化器狀態(tài)降到參數(shù)量乘以4字節(jié)左右但收斂性會受一些影響。Sophia優(yōu)化器用對角Hessian估計來替代二階動量也能省不少但實現(xiàn)復雜度和調(diào)參難度都上去了。實際做預算的時候我一般按這個公式估算總顯存等于參數(shù)量乘以2加2加8加激活值加緩沖區(qū)。前面的2是FP16參數(shù)第二個2是FP16梯度8是Adam的優(yōu)化器狀態(tài)。這個公式在AMP加AdamW的場景下比較準誤差在10%以內(nèi)。2.3 激活值的動態(tài)波動激活值是顯存占用里最動態(tài)的部分它跟batch size、序列長度、模型層數(shù)、隱藏維度都相關。粗略估算的話激活值顯存約等于batch size乘以序列長度乘以隱藏維度乘以層數(shù)乘以一個系數(shù)。這個系數(shù)取決于具體的網(wǎng)絡結構Transformer里主要是注意力矩陣和FFN的中間激活。注意力矩陣的顯存是序列長度的平方乘以batch size乘以頭數(shù)乘以2字節(jié)。序列長度2048時這個平方項就是4M乘以batch size和頭數(shù)后很容易上GB。序列長度拉到8192平方項變成64M直接爆炸。這就是長上下文訓練顯存吃緊的核心原因。FFN的中間激活是batch size乘以序列長度乘以4倍隱藏維度乘以2字節(jié)。這個和序列長度是線性關系比注意力矩陣溫和一些。但層數(shù)一多累積起來也很可觀。梯度檢查點gradient checkpointing就是針對激活值的優(yōu)化手段。它不保存中間激活而是在backward時重新計算。代價是計算量增加約30%但激活值顯存能降到原來的平方根級別。對于層數(shù)很深的模型這個 trade-off 非常劃算。2.4 臨時緩沖區(qū)與碎片臨時緩沖區(qū)包括CUDA kernel執(zhí)行時的workspace、通信操作的緩沖區(qū)、以及PyTorch分配器的預留空間。這部分很難精確測量但可以通過對比nvidia-smi和memory_allocated()的差值來估算。一般來說這個差值在500MB到2GB之間模型越大、并行策略越復雜差值越大。碎片問題在動態(tài)shape場景下特別嚴重。比如你訓練時序列長度不固定PyTorch的緩存分配器會按最大shape預留塊導致實際占用遠高于理論值。解決辦法是設置PYTORCH_CUDA_ALLOC_CONF環(huán)境變量啟用expandable_segments讓分配器能合并碎片。這個設置在新版PyTorch里效果很明顯我實測能把碎片率從20%降到5%以下。3. 預算決策從測量結果到優(yōu)化方案3.1 顯存預算的分配原則測完顯存賬單后下一步是做預算決策。預算的核心原則是先保訓練穩(wěn)定性再保吞吐量最后才考慮省顯存。很多人搞反了順序為了省顯存把batch size降到1結果訓練不穩(wěn)定收斂慢得要命省下來的顯存也沒換來什么好處。我的預算分配一般是這樣的參數(shù)和優(yōu)化器狀態(tài)是剛性支出沒法省必須留足。激活值是彈性支出可以通過梯度檢查點、序列并行、FlashAttention等手段壓縮。臨時緩沖區(qū)留10%到15%的余量防止峰值OOM。具體到數(shù)字上假設你有80GB顯存7B模型AMP加AdamW參數(shù)加梯度加優(yōu)化器狀態(tài)大約98GB單卡放不下。這時候你有幾個選擇一是上ZeRO Stage 2把優(yōu)化器狀態(tài)和梯度分片到多卡每卡只存一部分二是上LoRA只訓練低秩適配器參數(shù)量降到原來的百分之一三是上QLoRA把基座模型量化到4bit進一步壓縮。3.2 ZeRO與FSDP的顯存賬ZeRO Stage 1只分片優(yōu)化器狀態(tài)每卡顯存等于參數(shù)量乘以2加梯度乘以2加優(yōu)化器狀態(tài)除以N。7B模型8卡每卡優(yōu)化器狀態(tài)7GB加上參數(shù)14GB和梯度14GB總共35GB80GB卡能放下。Stage 2再分片梯度每卡梯度降到1.75GB總共約23GB。Stage 3分片參數(shù)每卡參數(shù)1.75GB總共約10GB但通信開銷大幅增加。FSDP本質(zhì)上是ZeRO Stage 3的PyTorch原生實現(xiàn)它把參數(shù)、梯度、優(yōu)化器狀態(tài)都分片每卡只存1/N。但FSDP在forward和backward時需要all-gather參數(shù)通信量很大。實際用下來FSDP在8卡A100上訓練7B模型每卡顯存約12GB但吞吐量比DDP低20%左右。這個 trade-off 要看你的瓶頸是顯存還是算力。3.3 AMP的顯存收益與代價AMP自動混合精度是顯存優(yōu)化的第一板斧但它的收益經(jīng)常被高估。AMP省的主要是激活值和部分計算緩沖區(qū)參數(shù)和優(yōu)化器狀態(tài)的大頭省不掉。實測下來AMP能把激活值顯存降到FP32的60%左右總體顯存節(jié)省約20%到30%。AMP的代價是數(shù)值穩(wěn)定性。FP16的動態(tài)范圍窄梯度容易下溢或上溢。PyTorch的GradScaler通過動態(tài)調(diào)整loss scale來緩解這個問題但在某些模型結構上仍然會出現(xiàn)NaN。BF16的動態(tài)范圍和FP32一樣不需要loss scaling但精度比FP16低。A100及以上支持BF16V100只支持FP16。選哪個取決于你的硬件和模型對精度的敏感度。實操心得用AMP時把LayerNorm和softmax強制保留在FP32能顯著提升訓練穩(wěn)定性。PyTorch的torch.cuda.amp.autocast默認已經(jīng)這么做了但如果你自己寫kernel要注意手動指定。4. 實操測量流程與工具鏈4.1 測量腳本的編寫要點我一般會寫一個獨立的測量腳本不摻在訓練代碼里這樣干凈、可復現(xiàn)。腳本的核心結構是初始化模型和優(yōu)化器構造一個真實batch的輸入然后分階段打點。import torch from torch.cuda import memory_allocated, max_memory_allocated, reset_peak_memory_stats def measure(model, optimizer, input_ids, labels): reset_peak_memory_stats() base memory_allocated() # forward outputs model(input_ids, labelslabels) loss outputs.loss after_forward memory_allocated() # backward loss.backward() after_backward memory_allocated() # optimizer step optimizer.step() optimizer.zero_grad() after_step memory_allocated() return { base: base, forward_delta: after_forward - base, backward_delta: after_backward - after_forward, step_delta: after_step - after_backward, peak: max_memory_allocated() }這個腳本跑一次就能拿到四個關鍵數(shù)字。forward_delta主要是激活值backward_delta是梯度加激活值的峰值step_delta是優(yōu)化器狀態(tài)的增量。peak是整個過程的最大值用來做OOM判斷。4.2 不同配置的對比測量測量不能只測一個配置要測一組配置做對比。我一般會測這幾組FP32基線、AMP、AMP加梯度檢查點、AMP加梯度檢查點加ZeRO。每組跑三次取平均排除冷啟動的影響。對比的時候重點看兩個指標峰值顯存和吞吐量。峰值顯存決定你能不能跑起來吞吐量決定你跑得多快。有時候省顯存的方案會把吞吐量砍半這時候就要算一筆賬省下來的顯存能不能換來更大的batch size如果能吞吐量可能反而更高。舉個例子7B模型FP32訓練峰值顯存60GB吞吐量100 samples/s。AMP后峰值45GB吞吐量130 samples/s。AMP加梯度檢查點后峰值30GB吞吐量90 samples/s。雖然梯度檢查點讓吞吐量降了但省下的15GB顯存可以讓你把batch size翻倍實際吞吐量變成180 samples/s。這就是預算決策的價值。4.3 顯存碎片的手動清理測量過程中如果發(fā)現(xiàn)nvidia-smi和memory_allocated()差值越來越大說明碎片在累積。這時候可以手動調(diào)torch.cuda.empty_cache()但注意這個操作會釋放緩存分配器預留的塊可能導致后續(xù)分配變慢。更好的辦法是設置PYTORCH_CUDA_ALLOC_CONFexpandable_segments:True讓分配器自己管理碎片。還有一個技巧是在訓練循環(huán)里定期調(diào)用torch.cuda.reset_peak_memory_stats()把峰值統(tǒng)計重置這樣能更準確地看到每個step的顯存波動。如果不重置峰值會一直保留歷史最大值掩蓋了后期的顯存增長。5. 常見問題與排查技巧實錄5.1 測量結果與nvidia-smi對不上這是最常見的問題。memory_allocated()只統(tǒng)計PyTorch分配器管理的顯存不包括CUDA上下文、cuDNN workspace、NCCL通信緩沖區(qū)。這些加起來可能有1到2GB。所以nvidia-smi看到的數(shù)字總是比memory_allocated()大。解決辦法是用torch.cuda.memory_snapshot()拿到完整的內(nèi)存快照它會列出所有分配塊包括那些不在PyTorch管理范圍內(nèi)的。不過這個輸出很大一般只在排查問題時用。5.2 激活值峰值出現(xiàn)在意想不到的地方有時候你會發(fā)現(xiàn)backward階段的顯存峰值比forward高很多甚至高出一倍。這通常是因為某些操作的backward需要保存forward的中間結果而這些結果在forward結束后并沒有釋放。比如attention的softmax輸出forward時算完就丟了但backward需要它來計算梯度所以PyTorch會把它保留到backward。排查方法是給每個模塊注冊forward hook和backward hook記錄每個模塊的輸入輸出顯存。這樣能精確定位到哪個模塊的backward顯存開銷異常。5.3 OOM發(fā)生在optimizer.step()optimizer.step()本身不分配大塊顯存但它會觸發(fā)參數(shù)更新而參數(shù)更新需要讀取梯度、寫入?yún)?shù)。如果梯度是FP16而參數(shù)是FP32這里會有一個隱式的類型轉(zhuǎn)換產(chǎn)生臨時緩沖區(qū)。Adam的動量更新還會產(chǎn)生中間變量。這些加起來可能幾百MB在顯存已經(jīng)接近上限時就是壓死駱駝的最后一根稻草。解決辦法是在step之前手動torch.cuda.empty_cache()或者把優(yōu)化器狀態(tài)分片到多卡。另一個辦法是用foreach實現(xiàn)的優(yōu)化器它把多個參數(shù)的更新合并成一個kernel減少臨時緩沖區(qū)的分配次數(shù)。5.4 梯度檢查點與AMP的兼容問題梯度檢查點在backward時會重新計算forward如果和AMP一起用重計算時的精度策略要和原forward一致否則梯度會對不上。PyTorch的torch.utils.checkpoint默認會保留AMP的autocast狀態(tài)但如果你自己實現(xiàn)了checkpoint邏輯要注意手動傳遞autocast上下文。還有一個坑是梯度檢查點不能和某些自定義autograd函數(shù)一起用因為重計算時這些函數(shù)的forward可能不是確定性的。遇到這種情況要么把自定義函數(shù)排除在checkpoint范圍外要么確保它是確定性的。5.5 多卡訓練時的顯存不均衡數(shù)據(jù)并行時每卡的顯存占用應該基本一致。如果發(fā)現(xiàn)某張卡顯存特別高通常是數(shù)據(jù)加載不均衡或者通信操作卡住了。檢查DataLoader的num_workers和pin_memory設置確保每個進程拿到的batch大小一致。通信方面NCCL的all-reduce是同步操作如果某張卡算得慢其他卡會在通信點等待顯存占用會暫時升高。排查方法是打印每張卡的memory_allocated()看差異是否超過5%。如果超過檢查數(shù)據(jù)分片邏輯和通信配置。問題現(xiàn)象可能原因排查方法解決手段nvidia-smi比memory_allocated高2GBCUDA上下文和通信緩沖區(qū)memory_snapshot預留2GB余量backward峰值遠高于forward中間激活未釋放模塊級hook梯度檢查點step時OOM類型轉(zhuǎn)換臨時緩沖區(qū)逐步打點foreach優(yōu)化器多卡顯存不均衡數(shù)據(jù)或通信不均衡逐卡打印調(diào)整DataLoader碎片率持續(xù)增長動態(tài)shape監(jiān)控差值expandable_segments6. 預算決策的實戰(zhàn)案例6.1 單卡24GB訓練7B模型的可行性分析有人問單卡24GB能不能訓7B模型。按前面的公式算7B模型AMP加AdamW參數(shù)14GBFP16加梯度14GBFP16加優(yōu)化器狀態(tài)56GBFP32總共84GB遠超24GB。所以全量微調(diào)不可能。但LoRA可以。LoRA只訓練低秩矩陣參數(shù)量通常是原模型的0.1%到1%。7B模型的LoRA參數(shù)量約7M到70MFP16下占14MB到140MB。優(yōu)化器狀態(tài)按Adam算是參數(shù)量乘以8也就是112MB到1.1GB。加上基座模型的14GBFP16總共約15GB到16GB。24GB卡能放下還能留出8GB給激活值和緩沖區(qū)。QLoRA更進一步把基座模型量化到4bit7B模型只占3.5GB。加上LoRA參數(shù)和優(yōu)化器狀態(tài)總共約5GB。24GB卡能跑得很寬裕甚至能上更大的batch size。6.2 多卡場景下的并行策略選擇如果你有4張24GB卡總共96GB顯存想訓7B模型全量微調(diào)。DDP每卡都要存完整的參數(shù)、梯度、優(yōu)化器狀態(tài)每卡84GB放不下。ZeRO Stage 2把優(yōu)化器狀態(tài)和梯度分片到4卡每卡參數(shù)14GB加梯度3.5GB加優(yōu)化器狀態(tài)14GB總共31.5GB還是放不下。ZeRO Stage 3把參數(shù)也分片每卡參數(shù)3.5GB加梯度3.5GB加優(yōu)化器狀態(tài)14GB總共21GB勉強能放下但激活值還沒算。這時候要么上梯度檢查點把激活值壓到2GB以內(nèi)要么上CPU offload把優(yōu)化器狀態(tài)放到內(nèi)存。CPU offload的代價是step速度慢3到5倍但能省下14GB顯存。實測下來4卡24GB加ZeRO Stage 3加梯度檢查點加CPU offload能訓7B模型但吞吐量只有DDP的十分之一。所以如果追求速度還是建議上80GB卡。6.3 預算決策的決策樹我把預算決策整理成一個簡單的決策樹方便快速判斷第一步算剛性支出參數(shù)量乘以2加2加8除以并行度。如果這個數(shù)字超過單卡顯存的80%考慮LoRA或QLoRA。第二步算激活值batch size乘以序列長度乘以隱藏維度乘以層數(shù)乘以系數(shù)。如果超過剩余顯存的50%上梯度檢查點。第三步算碎片余量留10%到15%的顯存給臨時緩沖區(qū)和碎片。如果不夠調(diào)expandable_segments或減小batch size。第四步測吞吐量在滿足顯存約束的前提下找吞吐量最大的配置。有時候大batch加梯度檢查點比小batch不加檢查點更快。這個決策樹不是絕對的但能幫你快速縮小選擇范圍。實際做的時候還是要跑測量腳本驗證因為不同模型結構、不同框架版本的顯存行為差異很大。最后分享一個小技巧在訓練腳本里加一個顯存監(jiān)控回調(diào)每個step記錄峰值顯存和吞吐量輸出到TensorBoard。這樣你能看到顯存隨訓練進程的變化趨勢及時發(fā)現(xiàn)碎片累積或激活值增長的問題。我靠這個回調(diào)抓到過好幾次隱蔽的顯存泄漏都是自定義層里緩存了不該緩存的東西。