化:從訓(xùn)練顯存到推理量化的工程實(shí)踐)
1. 從模型優(yōu)化器這個(gè)命名說起它到底在解決什么問題第一次看到 Model-Optimizer 這個(gè)名字很多人會(huì)下意識(shí)地把它歸類成又一個(gè)調(diào)參工具或者訓(xùn)練加速庫。但如果你真的在工程一線待過就會(huì)明白這個(gè)命名背后藏著一個(gè)非常具體的痛點(diǎn)模型從能跑到跑得好、跑得省、跑得穩(wěn)之間隔著一條巨大的鴻溝而這條鴻溝恰恰是絕大多數(shù)團(tuán)隊(duì)在項(xiàng)目落地階段最容易翻車的地方。我在過去幾年里接觸過不少模型落地的項(xiàng)目一個(gè)非常普遍的現(xiàn)象是算法同學(xué)在實(shí)驗(yàn)環(huán)境里把模型訓(xùn)到某個(gè)指標(biāo)興高采烈地交給工程同學(xué)部署結(jié)果一上生產(chǎn)環(huán)境就出問題——推理延遲翻了三倍、顯存直接爆掉、batch size 只能開到 1、量化之后精度掉得沒法看。這時(shí)候大家才回過頭來找原因發(fā)現(xiàn)根本不是模型本身不行而是整個(gè)優(yōu)化鏈路從來沒有被系統(tǒng)性地設(shè)計(jì)過。Model-Optimizer 這類工具要解決的正是這個(gè)優(yōu)化鏈路碎片化的問題。它不是一個(gè)單點(diǎn)工具而更像是一套圍繞模型全生命周期的優(yōu)化方法論 工具集從訓(xùn)練階段的顯存優(yōu)化、梯度累積策略到推理階段的量化、剪枝、算子融合再到部署階段的圖優(yōu)化、內(nèi)存復(fù)用、動(dòng)態(tài) shape 處理每一環(huán)都有對(duì)應(yīng)的手段而這些手段之間又需要協(xié)同不能各干各的。舉個(gè)很典型的例子。很多團(tuán)隊(duì)做量化的時(shí)候是拿一個(gè)已經(jīng)訓(xùn)練好的 FP32 模型直接上 INT8 量化結(jié)果精度崩了就得出這個(gè)模型不能量化的結(jié)論。但實(shí)際上如果訓(xùn)練階段就引入量化感知訓(xùn)練QAT或者在量化前先做一輪權(quán)重均衡weight equalization精度損失往往能控制在 1% 以內(nèi)。這就是優(yōu)化鏈路協(xié)同的價(jià)值——單點(diǎn)優(yōu)化做到極致也不如鏈路協(xié)同做到及格。所以這篇文章我想聊的不是某個(gè)具體 API 怎么調(diào)而是把 Model-Optimizer 這類工具背后的優(yōu)化思路拆開講清楚它為什么這么設(shè)計(jì)、每個(gè)優(yōu)化手段的適用邊界在哪、實(shí)際用的時(shí)候哪些坑最容易踩。適合已經(jīng)有一定模型訓(xùn)練和部署經(jīng)驗(yàn)、正在為模型跑不快/跑不動(dòng)/跑不穩(wěn)發(fā)愁的工程師也適合想系統(tǒng)建立模型優(yōu)化知識(shí)框架的同學(xué)。2. 訓(xùn)練階段的優(yōu)化顯存、吞吐與收斂的三方博弈2.1 顯存優(yōu)化不是省著用而是算著用訓(xùn)練階段最容易被忽視的優(yōu)化點(diǎn)就是顯存。很多人覺得顯存不夠就加卡、換大顯存機(jī)器這是最貴也最笨的做法。實(shí)際上顯存占用是可以被精確拆解和優(yōu)化的。一個(gè)標(biāo)準(zhǔn)的訓(xùn)練顯存占用大致分四塊模型參數(shù)、梯度、優(yōu)化器狀態(tài)、激活值。前三塊是相對(duì)固定的激活值則和 batch size、序列長(zhǎng)度強(qiáng)相關(guān)。以 Adam 優(yōu)化器為例它要為每個(gè)參數(shù)維護(hù)一階矩和二階矩兩個(gè)狀態(tài)所以優(yōu)化器狀態(tài)的顯存占用是參數(shù)量的兩倍。一個(gè) 7B 參數(shù)的模型FP16 存儲(chǔ)參數(shù)占 14GB梯度占 14GBAdam 狀態(tài)占 28GB光這三項(xiàng)就 56GB 了還沒算激活值。Model-Optimizer 這類工具在顯存優(yōu)化上通常提供幾個(gè)層次的方案梯度累積用時(shí)間換空間小 batch 多次前向后向再統(tǒng)一更新等效于大 batch。這個(gè)手段幾乎零成本但要注意 BatchNorm 層的統(tǒng)計(jì)量會(huì)受影響如果模型里有 BN要么改用 GroupNorm要么在累積時(shí)同步 BN 統(tǒng)計(jì)。梯度檢查點(diǎn)Gradient Checkpointing用計(jì)算換顯存前向時(shí)不保存中間激活反向時(shí)重新計(jì)算。實(shí)測(cè)下來通常能省 50%-70% 的激活顯存代價(jià)是訓(xùn)練速度慢 20%-30%。這里有個(gè)經(jīng)驗(yàn)不要對(duì)整個(gè)模型無腦開 checkpoint只對(duì)顯存占用最大的那幾個(gè) block 開性價(jià)比最高。優(yōu)化器狀態(tài)分片ZeRO 系列思路把優(yōu)化器狀態(tài)切分到不同設(shè)備上。Stage 1 切優(yōu)化器狀態(tài)Stage 2 再切梯度Stage 3 連參數(shù)都切。Stage 3 最省顯存但通信開銷最大實(shí)際選型要看你的卡間帶寬。提示顯存優(yōu)化有個(gè)反直覺的點(diǎn)——不是越省越好。過度優(yōu)化會(huì)引入大量重計(jì)算和通信訓(xùn)練吞吐掉得厲害最后總訓(xùn)練時(shí)間反而更長(zhǎng)。我的經(jīng)驗(yàn)是先把顯存壓到剛好能跑滿目標(biāo) batch size即可留 10%-15% 余量應(yīng)對(duì)波動(dòng)。2.2 吞吐優(yōu)化的核心是找到計(jì)算與通信的平衡點(diǎn)訓(xùn)練吞吐上不去八成是卡在通信上。尤其是多卡訓(xùn)練時(shí)數(shù)據(jù)并行下的梯度 AllReduce、模型并行下的層間通信都是吞吐殺手。一個(gè)實(shí)用的排查方法是先單卡跑看 GPU 利用率。如果單卡利用率就上不去比如只有 40%-60%那問題在數(shù)據(jù)加載或算子效率跟多卡通信無關(guān)。如果單卡能跑滿多卡掉利用率那才是通信瓶頸。針對(duì)通信瓶頸常見的優(yōu)化手段包括優(yōu)化手段適用場(chǎng)景預(yù)期收益注意事項(xiàng)梯度融合小梯度多的情況減少通信次數(shù) 30%融合后顯存略增通信重疊計(jì)算通信可并行隱藏 50% 通信時(shí)間需要框架支持混合并行超大模型突破單卡顯存限制調(diào)參復(fù)雜度高梯度壓縮帶寬受限通信量降 10 倍可能影響收斂這里重點(diǎn)說梯度融合。很多模型里大量參數(shù)是 bias、LayerNorm 的 scale/shift 這類小張量如果每個(gè)都單獨(dú)做一次 AllReduce通信次數(shù)會(huì)爆炸。把這些小梯度 concat 成一個(gè)大 buffer 再通信能顯著減少通信啟動(dòng)開銷。實(shí)測(cè)在一個(gè) Transformer 模型上梯度融合能把通信時(shí)間從 35% 降到 18% 左右。2.3 收斂性優(yōu)化別讓優(yōu)化手段把模型優(yōu)化壞了這是最容易被忽視的一點(diǎn)。很多優(yōu)化手段在提升效率的同時(shí)會(huì)悄悄改變訓(xùn)練的數(shù)值行為導(dǎo)致收斂變差甚至不收斂。比如混合精度訓(xùn)練FP16 的動(dòng)態(tài)范圍窄梯度容易下溢成 0。標(biāo)準(zhǔn)做法是配 loss scaling但 scale 值設(shè)多少有講究太小起不到防下溢作用太大又會(huì)導(dǎo)致梯度上溢成 inf?,F(xiàn)在主流框架都用動(dòng)態(tài) loss scaling每隔若干步檢查是否有 inf/nan有就跳過該步并降低 scale連續(xù)多步?jīng)]溢出就提高 scale。這個(gè)機(jī)制本身很穩(wěn)但如果你自己改了梯度裁剪的邏輯可能和動(dòng)態(tài) scale 打架導(dǎo)致 scale 一直上不去訓(xùn)練效率反而下降。再比如梯度累積它等效于大 batch但大 batch 通常需要調(diào)大學(xué)習(xí)率才能保持收斂速度。如果你累積了 8 步但學(xué)習(xí)率沒變收斂會(huì)明顯變慢。經(jīng)驗(yàn)公式是學(xué)習(xí)率隨等效 batch size 線性放大但放大到一定程度后要改用平方根縮放否則訓(xùn)練不穩(wěn)定。3. 推理階段的優(yōu)化量化、剪枝與算子融合的取舍3.1 量化精度與速度的蹺蹺板怎么壓量化是推理優(yōu)化里收益最直接的手段。FP32 到 INT8理論上顯存降 4 倍、速度提升 2-4 倍。但實(shí)際做下來能拿到 2 倍加速就算不錯(cuò)了因?yàn)楹芏嗨阕訉?duì) INT8 支持不好或者量化后要插入額外的反量化操作抵消了收益。量化的路線選擇上我一般這么判斷PTQ訓(xùn)練后量化零訓(xùn)練成本適合快速驗(yàn)證。但精度損失不可控尤其是激活值分布不均勻的模型比如有大量異常值的掉點(diǎn)可能很嚴(yán)重。QAT量化感知訓(xùn)練訓(xùn)練時(shí)模擬量化誤差精度保持好但需要重新訓(xùn)練成本高。動(dòng)態(tài)量化只量化權(quán)重激活值運(yùn)行時(shí)動(dòng)態(tài)量化。適合 NLP 類模型尤其是 LSTM、Transformer 的線性層。靜態(tài)量化權(quán)重和激活都提前量化需要校準(zhǔn)數(shù)據(jù)集。適合 CNN 類模型。這里有個(gè)實(shí)操經(jīng)驗(yàn)做 PTQ 之前先看激活值的分布。如果某一層的激活值最大值和均值差了幾個(gè)數(shù)量級(jí)那這層大概率是量化敏感層要么跳過不量化要么用 per-channel 量化而不是 per-tensor。per-channel 量化對(duì)權(quán)重的效果尤其明顯因?yàn)椴煌敵鐾ǖ赖臋?quán)重分布差異往往很大。還有一個(gè)坑是量化后的算子融合。比如 Conv BN ReLU 這種經(jīng)典組合量化時(shí)如果分別量化再融合中間的反量化/再量化操作會(huì)吃掉大部分收益。正確做法是先把 FP32 的算子融合好再整體量化。Model-Optimizer 這類工具通常會(huì)在量化前先做一輪圖優(yōu)化把能融的算子都融掉就是這個(gè)道理。3.2 剪枝結(jié)構(gòu)化與非結(jié)構(gòu)化的路線之爭(zhēng)剪枝的誘惑在于它理論上能直接減少參數(shù)量和計(jì)算量。但實(shí)際落地時(shí)非結(jié)構(gòu)化剪枝把單個(gè)權(quán)重置零在通用硬件上幾乎拿不到加速因?yàn)?GPU 是 SIMD 架構(gòu)一個(gè) warp 里只要有一個(gè)非零權(quán)重整個(gè)計(jì)算就得做。除非你有支持稀疏計(jì)算的專用硬件或庫。結(jié)構(gòu)化剪枝直接剪掉整個(gè)通道、注意力頭、層才能真正加速因?yàn)樗淖兞藦埩啃螤睢5Y(jié)構(gòu)化剪枝對(duì)精度的影響更大需要配合微調(diào)。我的經(jīng)驗(yàn)是剪枝比例從 10%-20% 開始試不要一上來就剪 50%。剪枝后一定要做一輪短時(shí)間的微調(diào)幾個(gè) epoch 即可精度通常能恢復(fù)大部分。優(yōu)先剪那些冗余度高的結(jié)構(gòu)比如注意力頭。很多研究表明 Transformer 里有大量注意力頭是可以剪掉的對(duì)精度影響很小。3.3 算子融合與圖優(yōu)化不改變數(shù)學(xué)等價(jià)性的加速算子融合是免費(fèi)的加速——它不改變計(jì)算結(jié)果只是減少 kernel 啟動(dòng)次數(shù)和中間張量的讀寫。典型的融合模式包括Element-wise 融合把連續(xù)的 add、mul、relu 等逐元素操作融成一個(gè) kernel減少顯存帶寬占用。Conv-BN 融合推理時(shí) BN 是線性變換可以直接折進(jìn) Conv 的權(quán)重和 bias 里。MatMul-Bias-Activation 融合Transformer 里的標(biāo)準(zhǔn)模式融合后能省一次中間張量寫回。圖優(yōu)化還包括常量折疊、死代碼消除、內(nèi)存復(fù)用等。內(nèi)存復(fù)用尤其重要推理時(shí)中間張量的生命周期往往不重疊可以復(fù)用同一塊顯存。一個(gè)優(yōu)化良好的推理圖峰值顯存能比樸素實(shí)現(xiàn)低 30%-50%。注意圖優(yōu)化要在量化之前做因?yàn)榱炕蟮膱D結(jié)構(gòu)會(huì)變復(fù)雜很多融合模式就匹配不上了。順序錯(cuò)了優(yōu)化效果大打折扣。4. 部署階段的工程細(xì)節(jié)那些文檔里不會(huì)寫的東西4.1 動(dòng)態(tài) shape 是推理性能的隱形殺手訓(xùn)練時(shí) shape 固定推理時(shí)輸入長(zhǎng)度千變?nèi)f化這是部署階段最頭疼的問題之一。動(dòng)態(tài) shape 會(huì)導(dǎo)致每次新 shape 都觸發(fā)一次 kernel 編譯或 autotune首次推理特別慢。顯存分配器無法復(fù)用內(nèi)存塊峰值顯存飆升。某些優(yōu)化如固定 tile size 的算子直接失效。應(yīng)對(duì)策略有幾個(gè)層次。最簡(jiǎn)單的是padding 到固定長(zhǎng)度比如把所有輸入 pad 到 128 的倍數(shù)。代價(jià)是短輸入浪費(fèi)算力但換來的是穩(wěn)定的性能。進(jìn)階做法是準(zhǔn)備幾檔預(yù)設(shè) shape如 128、256、512、1024輸入落到哪檔就 pad 到哪檔兼顧性能和浪費(fèi)。再進(jìn)階就是真正的動(dòng)態(tài) shape 支持但這需要推理引擎底層做大量工作不是應(yīng)用層能解決的。我的經(jīng)驗(yàn)是先統(tǒng)計(jì)線上輸入的長(zhǎng)度分布。如果 90% 的請(qǐng)求都在 200 以內(nèi)那就沒必要為那 10% 的長(zhǎng)尾做動(dòng)態(tài)優(yōu)化直接 pad 到 256 最劃算。4.2 批處理策略吞吐與延遲的權(quán)衡推理服務(wù)里batch size 直接決定吞吐和延遲。batch 越大吞吐越高但單個(gè)請(qǐng)求的等待時(shí)間也越長(zhǎng)要等湊夠一批。這是典型的吞吐-延遲權(quán)衡。常見的策略包括靜態(tài) batching固定 batch size湊夠才發(fā)。延遲不可控適合離線任務(wù)。動(dòng)態(tài) batching設(shè)一個(gè)最大等待時(shí)間窗口如 10ms窗口內(nèi)到的請(qǐng)求湊一批。這是在線服務(wù)的主流做法。連續(xù) batching請(qǐng)求完成一個(gè)就補(bǔ)一個(gè)不等整批結(jié)束。適合生成式任務(wù)能顯著提升 GPU 利用率。連續(xù) batching 是這兩年生成式模型推理的標(biāo)配。傳統(tǒng) batching 下一批里最長(zhǎng)的序列決定了整批的完成時(shí)間短序列的算力全浪費(fèi)了。連續(xù) batching 讓每個(gè)序列獨(dú)立推進(jìn)GPU 利用率能從 30% 提到 70% 以上。但它的實(shí)現(xiàn)復(fù)雜度高需要管理每個(gè)序列的 KV cache 和狀態(tài)自己寫很容易出 bug建議直接用成熟推理框架。4.3 顯存池化與碎片治理推理服務(wù)跑久了顯存碎片是個(gè)繞不開的問題。表現(xiàn)是明明總顯存夠但就是分配不出連續(xù)的大塊導(dǎo)致 OOM。根因在于不同 shape 的請(qǐng)求交替到來顯存分配器反復(fù)申請(qǐng)釋放不同大小的塊久而久之就碎了。解決辦法預(yù)分配顯存池啟動(dòng)時(shí)就把顯存按檔位切好運(yùn)行時(shí)只從池里取不向系統(tǒng)申請(qǐng)。統(tǒng)一 shape 檔位配合前面的 padding 策略讓所有請(qǐng)求都落到有限的幾檔 shape 上分配器就能高效復(fù)用。定期重啟最土但最有效。如果碎片問題嚴(yán)重到影響穩(wěn)定性設(shè)置一個(gè)低峰期自動(dòng)重啟比任何優(yōu)化都省心。5. 優(yōu)化效果的度量別用感覺用數(shù)據(jù)說話5.1 建立基線是一切優(yōu)化的前提我見過太多團(tuán)隊(duì)優(yōu)化了半天問他優(yōu)化前多少、優(yōu)化后多少答不上來。沒有基線所有優(yōu)化都是玄學(xué)。建立基線要固定幾個(gè)變量硬件型號(hào)、軟件版本、輸入 shape 分布、batch size、并發(fā)數(shù)。然后測(cè)這幾個(gè)指標(biāo)延遲P50、P90、P99 都要看。P99 才是用戶體驗(yàn)的真實(shí)反映。吞吐每秒處理請(qǐng)求數(shù)QPS或每秒處理 token 數(shù)。顯存峰值決定你能開多大 batch、能并發(fā)多少請(qǐng)求。GPU 利用率低于 50% 說明有優(yōu)化空間高于 90% 說明可能已經(jīng)到瓶頸。測(cè)的時(shí)候要注意預(yù)熱。GPU 首次運(yùn)行會(huì)有各種初始化開銷前幾十次推理的數(shù)據(jù)不能算。一般跑 100 次預(yù)熱再測(cè) 1000 次取統(tǒng)計(jì)值。5.2 優(yōu)化收益的歸因分析當(dāng)你做了一堆優(yōu)化性能提升了 2 倍怎么知道是哪項(xiàng)優(yōu)化貢獻(xiàn)的答案是控制變量逐項(xiàng)測(cè)。我的做法是維護(hù)一個(gè)優(yōu)化清單每加一項(xiàng)就測(cè)一次記錄增量收益。這樣最后能清楚地知道每項(xiàng)優(yōu)化的 ROI。有些優(yōu)化可能只貢獻(xiàn) 5% 的提升但引入了大量復(fù)雜度那就該砍掉。一個(gè)常見的誤區(qū)是過早優(yōu)化。比如模型還沒跑通就上量化結(jié)果精度崩了回頭排查發(fā)現(xiàn)是模型本身就有問題。正確的順序是先保證正確性再優(yōu)化性能先優(yōu)化大頭如量化、batching再摳細(xì)節(jié)如算子融合。5.3 精度回歸測(cè)試不能省任何優(yōu)化手段都可能影響精度所以每次優(yōu)化后都要跑精度回歸。測(cè)試集要覆蓋線上真實(shí)分布不能只用公開 benchmark。量化的精度回歸尤其重要。建議準(zhǔn)備一個(gè)小規(guī)模但高覆蓋的校準(zhǔn)集包含各種邊界 case超長(zhǎng)輸入、空輸入、特殊字符等。量化后如果某個(gè) case 的輸出和 FP32 差異超過閾值就要重點(diǎn)排查是不是那層的量化參數(shù)有問題。6. 我在實(shí)際項(xiàng)目中踩過的幾個(gè)坑第一個(gè)坑是盲目追求低精度。有次為了壓顯存把模型從 FP16 降到 INT8顯存確實(shí)降了一半但推理速度只快了 20%因?yàn)槟P屠镉写罅?LayerNorm 和 Softmax這些算子在 INT8 下要么不支持、要么要插反量化收益被吃掉了。后來改成只量化線性層速度反而提升了 1.8 倍。教訓(xùn)是量化要看算子構(gòu)成不是所有模型都適合全量化。第二個(gè)坑是忽略 warmup 對(duì)性能測(cè)試的影響。有次測(cè)出來優(yōu)化后 P99 延遲反而變高了排查半天發(fā)現(xiàn)是新引入的某個(gè)算子在首次遇到新 shape 時(shí)要編譯而測(cè)試流量里恰好有少量長(zhǎng)尾 shape。后來加了 shape 預(yù)熱把常見 shape 在啟動(dòng)時(shí)都跑一遍P99 就正常了。第三個(gè)坑是優(yōu)化手段之間互相打架。梯度檢查點(diǎn)省顯存但它改變了計(jì)算圖導(dǎo)致某些算子融合失效量化感知訓(xùn)練提升量化精度但它要求訓(xùn)練時(shí)插入偽量化節(jié)點(diǎn)又增加了訓(xùn)練顯存。所以優(yōu)化不是簡(jiǎn)單疊加要按優(yōu)先級(jí)排序先做收益大且無副作用的再做有取舍的。最后一個(gè)心得優(yōu)化是個(gè)持續(xù)過程不是一次性任務(wù)。模型在迭代、流量在變化、硬件在更新今天的優(yōu)化方案明天可能就不適用了。建議把性能測(cè)試做成 CI 的一部分每次模型更新都自動(dòng)跑一遍性能回退能第一時(shí)間發(fā)現(xiàn)。這比事后救火省心得多。