練實(shí)戰(zhàn):顯存優(yōu)化與通信調(diào)優(yōu))
大模型訓(xùn)練走到今天單卡單機(jī)的時代早就過去了。但凡參數(shù)規(guī)模上到百億、千億甚至只是想在有限顯存里塞下一個稍微像樣的模型你都會撞上同一堵墻顯存不夠。DeepSpeed 的 ZeRO 系列就是為解決這堵墻而生的而 ZeRO-3 是其中把顯存優(yōu)化做到最激進(jìn)的一檔。與此同時MoE混合專家架構(gòu)又給訓(xùn)練帶來了另一套完全不同的挑戰(zhàn)——它不追求把每個參數(shù)都用上而是讓不同的 token 走不同的專家路徑。把 ZeRO-3 和 MoE 放在一起訓(xùn)練是很多團(tuán)隊在做的組合但這兩者的脾氣并不完全對付。這篇內(nèi)容就圍繞 DeepSpeed ZeRO-3 與 MoE 訓(xùn)練展開把顯存到底省在哪、MoE 的參數(shù)到底要不要全進(jìn)顯存、兩者結(jié)合時哪些地方容易翻車一條條拆開講清楚。適合已經(jīng)跑過單機(jī)訓(xùn)練、準(zhǔn)備上分布式、或者正在被顯存和通信問題折磨的從業(yè)者參考。1. 先把顯存這筆賬算明白才知道 ZeRO-3 省的是什么很多人一上來就背 ZeRO 的三個 stage但真到調(diào)參的時候還是懵根本原因是沒搞清楚訓(xùn)練時顯存到底被誰吃掉了。我們先把這筆賬攤開。1.1 訓(xùn)練顯存的四大塊開銷在標(biāo)準(zhǔn)的混合精度訓(xùn)練里顯存主要被四部分占據(jù)模型參數(shù)FP16/BF16 權(quán)重本身占 2 字節(jié)每參數(shù)。梯度和參數(shù)同精度也是 2 字節(jié)每參數(shù)。優(yōu)化器狀態(tài)如果用 AdamFP32 的 master weight、一階動量、二階動量加起來是 12 字節(jié)每參數(shù)444。激活值前向傳播中間結(jié)果跟 batch size、序列長度、模型深度強(qiáng)相關(guān)這塊最不可控。拿一個 10B 參數(shù)的模型舉例光參數(shù)梯度優(yōu)化器狀態(tài)就是 221216 字節(jié)每參數(shù)也就是 160GB。這還沒算激活值。單張 80GB 的卡連靜態(tài)開銷都放不下更別提激活了。這就是為什么必須做切分。1.2 ZeRO 三個 stage 到底切了什么ZeRO 的核心思路是既然數(shù)據(jù)并行里每張卡都存了一份完整的模型狀態(tài)那這份冗余就是浪費(fèi)。它把模型狀態(tài)切成若干份每張卡只存一份需要的時候再通信拿回來。Stage切分對象單卡顯存相對基線通信量ZeRO-1優(yōu)化器狀態(tài)約 1/4低ZeRO-2優(yōu)化器狀態(tài) 梯度約 1/8中ZeRO-3優(yōu)化器狀態(tài) 梯度 參數(shù)約 1/NN 為卡數(shù)高ZeRO-1 只切優(yōu)化器狀態(tài)省得有限但通信開銷小適合顯存壓力不大的場景。ZeRO-2 把梯度也切了性價比通常最高是很多團(tuán)隊的主力選擇。ZeRO-3 把參數(shù)也切了每張卡只保留自己負(fù)責(zé)的那一片參數(shù)前向和反向時按需從其他卡 gather 過來用完就釋放。注意ZeRO-3 省顯存最狠但代價是通信量顯著上升。參數(shù)量越大、卡越多這個通信開銷越明顯。如果你的瓶頸是通信而不是顯存盲目上 ZeRO-3 反而會更慢。1.3 為什么 ZeRO-3 的通信開銷最大ZeRO-3 在每一層的前向和反向都要做一次參數(shù)的 all-gather。前向時把完整參數(shù)拼出來算算完釋放反向時再拼一次算梯度。這意味著參數(shù)量被反復(fù)搬運(yùn)。相比之下 ZeRO-2 只在梯度歸約時通信頻率低得多。所以選 stage 的邏輯很清晰顯存夠用就別上 ZeRO-3。只有當(dāng)參數(shù)本身大到單卡放不下、或者激活值已經(jīng)把顯存擠爆時ZeRO-3 才是必要的。我見過不少團(tuán)隊一上來就 ZeRO-3結(jié)果訓(xùn)練速度掉了一半其實(shí) ZeRO-2 加梯度檢查點(diǎn)就能解決。2. MoE 的參數(shù)到底要不要全部進(jìn)顯存這是被問得最多的問題之一也是 MoE 訓(xùn)練里最容易產(chǎn)生誤解的地方。答案要分情況說不能一句話概括。2.1 MoE 的結(jié)構(gòu)決定了它的顯存特性MoE 層里有一組專家expert每個 token 經(jīng)過路由router后只激活其中 top-k 個專家。比如 8 個專家、top-2 激活那每個 token 實(shí)際只用到 2 個專家的計算。關(guān)鍵點(diǎn)在于激活的專家少不代表參數(shù)少。整個 MoE 層的所有專家參數(shù)加起來可能非常龐大但每次前向只有一部分參與計算。這就帶來一個矛盾——計算量按激活比例算但顯存要按全部參數(shù)算。2.2 稠密訓(xùn)練 vs 稀疏訓(xùn)練的顯存賬假設(shè)一個 MoE 模型總參數(shù) 100B但每次激活只有 20B如果按稠密方式訓(xùn)練所有 100B 參數(shù)都要參與梯度更新優(yōu)化器狀態(tài)、梯度全都要存顯存按 100B 算。如果按稀疏方式訓(xùn)練只有被激活的專家參與更新理論上顯存可以按激活部分算。但現(xiàn)實(shí)是絕大多數(shù)框架在訓(xùn)練時仍然需要把全部專家參數(shù)加載到顯存里因為路由是動態(tài)的你無法預(yù)知下一個 batch 會激活哪些專家。所以MoE 訓(xùn)練時參數(shù)通常還是要全部進(jìn)顯存除非你用了專家并行expert parallelism把不同專家放到不同卡上。2.3 專家并行怎么解決這個問題專家并行的思路很直接既然專家之間是獨(dú)立的那就把不同專家分到不同設(shè)備上。每張卡只負(fù)責(zé)一部分專家token 通過 all-to-all 通信被送到對應(yīng)專家所在的卡上計算算完再送回來。這樣一來單卡顯存只需要裝下自己負(fù)責(zé)的那幾個專家而不是全部。代價是引入了 all-to-all 通信這個通信在跨節(jié)點(diǎn)時開銷很大是 MoE 訓(xùn)練的主要瓶頸之一。提示專家并行和 ZeRO-3 可以疊加使用。ZeRO-3 切分的是每張卡內(nèi)部的參數(shù)、梯度、優(yōu)化器狀態(tài)專家并行切分的是專家在不同卡之間的分布。兩者維度不同組合起來能進(jìn)一步壓低單卡顯存。3. ZeRO-3 和 MoE 結(jié)合時那些配置項到底怎么設(shè)理論講完落到實(shí)操。這一節(jié)把關(guān)鍵配置項和它們背后的邏輯講透避免你照抄配置卻不知道在改什么。3.1 DeepSpeed 配置文件的骨架一個典型的 ZeRO-3 MoE 配置大概長這樣{ train_batch_size: 64, gradient_accumulation_steps: 4, fp16: { enabled: true }, zero_optimization: { stage: 3, offload_optimizer: { device: cpu, pin_memory: true }, offload_param: { device: cpu, pin_memory: true }, overlap_comm: true, contiguous_gradients: true, stage3_gather_16bit_weights_on_model_save: true }, moe: { enabled: true, ep_size: 8, use_tutel: false } }這里幾個參數(shù)值得單獨(dú)說。3.2 offload 到底該不該開offload_optimizer和offload_param把優(yōu)化器狀態(tài)和參數(shù)卸載到 CPU 內(nèi)存。開了之后顯存壓力驟降但訓(xùn)練速度會明顯變慢因為 CPU 和 GPU 之間的 PCIe 帶寬遠(yuǎn)低于顯存帶寬。我的經(jīng)驗是顯存實(shí)在不夠再開 offload。如果開了 offload 之后訓(xùn)練慢到無法接受那說明你的并行策略有問題應(yīng)該優(yōu)先考慮加卡或者調(diào)整專家并行度而不是硬扛 offload。offload 是最后的救命稻草不是常規(guī)配置。3.3 overlap_comm 和 contiguous_gradients 的作用overlap_comm讓通信和計算重疊進(jìn)行能有效掩蓋一部分 all-gather 的延遲。contiguous_gradients把梯度整理成連續(xù)內(nèi)存減少內(nèi)存碎片對 ZeRO-3 這種頻繁分配釋放的場景很有幫助。這兩個參數(shù)在 ZeRO-3 下基本是默認(rèn)要開的除非你遇到特定的穩(wěn)定性問題。實(shí)測下來開了 overlap_comm 之后ZeRO-3 的吞吐能提升 10% 到 20%具體取決于模型結(jié)構(gòu)和網(wǎng)絡(luò)拓?fù)洹?.4 ep_size 怎么定ep_size是專家并行度表示專家被切分到多少張卡上。它的取值要和總卡數(shù)、專家數(shù)量配合。假設(shè)你有 64 張卡、64 個專家那 ep_size 設(shè)成 8 意味著每 8 張卡一組每組負(fù)責(zé) 8 個專家。設(shè)成 64 就是每個專家獨(dú)占一張卡。ep_size 越大單卡專家越少顯存越省但 all-to-all 通信范圍越大。注意ep_size 必須能整除總卡數(shù)也必須能整除專家總數(shù)否則會報錯。配置前先算清楚這兩個數(shù)。4. 訓(xùn)練跑起來之后那些讓人抓狂的報錯怎么排配置寫對了不代表就能順利跑。MoE ZeRO-3 的組合在實(shí)戰(zhàn)里有一批高頻報錯這一節(jié)按排查鏈路來講。4.1 安裝 deepspeed 就報錯先別懷疑代碼很多人第一步就卡在pip install deepspeed。這個包編譯依賴比較重常見報錯集中在幾個方向CUDA 版本不匹配deepspeed 編譯時會檢測本機(jī) CUDA如果 PyTorch 編譯用的 CUDA 版本和系統(tǒng) CUDA 不一致就會報錯。解決辦法是先確認(rèn)torch.version.cuda再裝對應(yīng)版本的 deepspeed。缺少編譯工具需要 gcc、g 等。報錯信息里如果有g(shù)cc: command not found裝一下 build-essential 就行。JIT 編譯超時deepspeed 有些算子會在首次運(yùn)行時 JIT 編譯如果環(huán)境變量沒設(shè)好會卡住??梢栽O(shè)置DS_BUILD_OPS0先跳過自定義算子編譯跑通流程再補(bǔ)。我一般建議先用pip install deepspeed --no-build-isolation試能省掉不少隔離環(huán)境帶來的坑。4.2 專家負(fù)載不均衡導(dǎo)致的訓(xùn)練崩潰MoE 最經(jīng)典的問題就是路由塌縮——所有 token 都涌向少數(shù)幾個專家其他專家餓死。表現(xiàn)是 loss 突然飆升或者直接 NaN。排查思路是這樣的先打印每個專家被選中的頻率。如果某個專家占比超過 50%基本可以確認(rèn)是負(fù)載不均衡。檢查有沒有加負(fù)載均衡損失load balancing loss。這是 MoE 訓(xùn)練的標(biāo)準(zhǔn)配置通常是一個輔助 loss鼓勵 token 均勻分配到各專家。檢查 router 的初始化。router 權(quán)重初始化太極端會導(dǎo)致早期就塌縮。負(fù)載均衡損失的代碼邏輯大致是統(tǒng)計每個專家被選中的比例和理想均勻分布做對比算一個輔助損失加進(jìn)總 loss。系數(shù)一般設(shè)在 0.01 量級太大影響主任務(wù)太小起不到作用。4.3 ZeRO-3 下的參數(shù) gather 超時ZeRO-3 在保存 checkpoint 或者做某些操作時需要把所有分片的參數(shù) gather 回來。如果模型很大、卡很多這個 gather 可能超時。stage3_gather_16bit_weights_on_model_save這個參數(shù)就是控制保存時是否 gather 的。如果保存經(jīng)常超時可以設(shè)成 false保存分片權(quán)重后續(xù)再合并。另外通信超時時間可以通過NCCL_TIMEOUT環(huán)境變量調(diào)大。4.4 顯存明明夠卻 OOM 的詭異情況有時候顯存監(jiān)控顯示還有余量但就是 OOM。這種情況在 ZeRO-3 MoE 下通常是內(nèi)存碎片導(dǎo)致的。contiguous_gradients能緩解一部分但更徹底的辦法是調(diào)整round_robin_gradients讓梯度分配更均勻。還有一種可能是激活值峰值。MoE 的 all-to-all 通信會產(chǎn)生額外的緩沖區(qū)這部分顯存容易被忽略??梢蚤_梯度檢查點(diǎn)gradient checkpointing把激活值壓下來用計算換顯存。5. 讓訓(xùn)練真正跑快的幾個調(diào)優(yōu)方向跑通只是第一步跑快才是本事。這一節(jié)講幾個實(shí)測有效的調(diào)優(yōu)方向。5.1 通信和計算的重疊程度ZeRO-3 的性能瓶頸幾乎全在通信上。overlap_comm能重疊一部分但重疊效果取決于模型結(jié)構(gòu)。如果某一層的計算量太小通信還沒藏完計算就結(jié)束了那這層就是瓶頸。一個實(shí)用的做法是調(diào)整 bucket size。DeepSpeed 會把參數(shù)按 bucket 分組做 all-gatherbucket 太小通信次數(shù)多太大則重疊粒度粗。默認(rèn)值通常夠用但在 MoE 場景下可以適當(dāng)調(diào)大因為專家層的參數(shù)塊比較大。5.2 專家并行的拓?fù)溥x擇all-to-all 通信對網(wǎng)絡(luò)拓?fù)浜苊舾?。如果專家并行組跨了節(jié)點(diǎn)通信要走慢速網(wǎng)絡(luò)開銷會暴漲。所以盡量讓專家并行組落在同一個節(jié)點(diǎn)內(nèi)用 NVLink 或高速互聯(lián)。假設(shè)一個節(jié)點(diǎn) 8 張卡那 ep_size 設(shè)成 8 就能保證專家并行不跨節(jié)點(diǎn)。如果非要設(shè)成 16就要接受跨節(jié)點(diǎn)通信的代價。這個取舍在配置階段就要想清楚。5.3 batch size 和梯度累積的配合MoE 訓(xùn)練對 batch size 比較敏感因為 batch 太小會導(dǎo)致路由統(tǒng)計不穩(wěn)定負(fù)載均衡損失波動大。但 batch 太大又吃顯存。折中方案是用梯度累積。train_batch_size設(shè)成實(shí)際想要的大 batchgradient_accumulation_steps設(shè)成累積步數(shù)兩者相除就是單步的 micro batch。這樣既保證了路由統(tǒng)計的穩(wěn)定性又控制了單步顯存。提示梯度累積步數(shù)增加會拉長單次迭代時間但不會增加顯存。如果顯存是瓶頸而時間不是這是最劃算的調(diào)法。5.4 混合精度的選擇BF16 相比 FP16 動態(tài)范圍更大不容易溢出在 MoE 訓(xùn)練里更穩(wěn)。如果硬件支持 BF16優(yōu)先用它。FP16 需要配合 loss scaling而 MoE 的輔助損失會讓 loss 尺度變化更復(fù)雜loss scaling 調(diào)起來更麻煩。實(shí)測下來同樣的配置換成 BF16訓(xùn)練穩(wěn)定性提升明顯尤其是訓(xùn)練后期 loss 已經(jīng)很小的時候FP16 更容易出現(xiàn)梯度下溢。6. 幾個容易被忽略的實(shí)戰(zhàn)細(xì)節(jié)最后聊幾個文檔里不太寫、但實(shí)際會踩的細(xì)節(jié)。6.1 checkpoint 的兼容性ZeRO-3 保存的 checkpoint 是分片的和普通 checkpoint 格式不一樣。如果你中途想換并行策略比如從 ZeRO-3 換到 ZeRO-2checkpoint 需要先合并轉(zhuǎn)換。這個轉(zhuǎn)換過程對 MoE 模型更麻煩因為專家參數(shù)還要考慮并行度的映射。建議在訓(xùn)練早期就把并行策略定下來別中途大改。如果非要改預(yù)留出轉(zhuǎn)換和驗證的時間。6.2 學(xué)習(xí)率 warmup 要更保守MoE 模型因為路由的存在早期訓(xùn)練更不穩(wěn)定。warmup 步數(shù)建議比同等規(guī)模的稠密模型更長一些讓 router 有時間找到合理的分配。我一般會把 warmup 比例從 1% 提到 2% 到 3%。6.3 監(jiān)控指標(biāo)要盯緊專家利用率除了常規(guī)的 loss、學(xué)習(xí)率、吞吐MoE 訓(xùn)練一定要盯專家利用率。如果發(fā)現(xiàn)某些專家長期不被激活要么是負(fù)載均衡沒做好要么是專家數(shù)量設(shè)多了。專家數(shù)量不是越多越好要和數(shù)據(jù)規(guī)模、任務(wù)復(fù)雜度匹配。6.4 別忽視數(shù)據(jù)質(zhì)量對路由的影響MoE 的路由是根據(jù) token 表示來決定的如果訓(xùn)練數(shù)據(jù)分布很偏路由也會偏。數(shù)據(jù)里如果某類樣本特別多對應(yīng)的專家就會被過度激活。所以數(shù)據(jù)配比在 MoE 訓(xùn)練里比稠密模型更關(guān)鍵預(yù)處理階段就要把分布調(diào)勻。我個人在幾次 MoE 項目里最大的體會是ZeRO-3 解決的是顯存問題MoE 解決的是容量和計算效率問題但兩者疊加之后真正的難點(diǎn)從能不能跑變成了跑得穩(wěn)不穩(wěn)、快不快。配置只是起點(diǎn)負(fù)載均衡、通信拓?fù)洹?shù)據(jù)分布這些軟性的東西才是決定訓(xùn)練成敗的關(guān)鍵。如果一開始就沖著省顯存去堆配置很容易在穩(wěn)定性上栽跟頭。先把小規(guī)模跑通、把專家利用率盯住再逐步放大這條路走得最穩(wěn)。