練側(cè)PD分離:突破大規(guī)模分布式訓(xùn)練瓶頸的新路徑)
1. 從“訓(xùn)練也PD分離”這個反直覺說法說起第一次聽到“訓(xùn)練也PD分離”這個說法我的反應(yīng)是這不是推理側(cè)玩剩下的嗎Prefill 和 Decode 分離部署在推理服務(wù)里早就被聊爛了——Prefill 是計算密集型Decode 是訪存密集型兩者對硬件資源的訴求完全不同混在一起跑就是互相拖后腿。但把這套思路搬到訓(xùn)練側(cè)尤其是大規(guī)模分布式訓(xùn)練的場景下事情就變得有意思了。先把這個概念說清楚。所謂訓(xùn)練側(cè)的 PD 分離核心思路是把一次完整的訓(xùn)練迭代拆成兩個性質(zhì)截然不同的階段一個是前向與反向傳播階段對應(yīng)推理里的 Prefill計算密集、需要大算力另一個是參數(shù)更新與優(yōu)化器狀態(tài)處理階段對應(yīng)推理里的 Decode訪存密集、需要大帶寬和顯存容量。在傳統(tǒng)的數(shù)據(jù)并行或張量并行訓(xùn)練里這兩個階段是緊耦合的每一步都串行執(zhí)行硬件資源在階段切換時存在明顯的利用率波動。為什么現(xiàn)在要重新審視這件事因為 Scaling 的瓶頸變了。過去幾年大家拼的是算力堆疊卡越多、模型越大、效果越好這條路徑在千卡萬卡級別還能走但到了更大規(guī)模通信開銷、顯存墻、優(yōu)化器狀態(tài)膨脹這些問題開始吃掉大部分收益。單純堆硬件的邊際效益在遞減而結(jié)構(gòu)性的效率優(yōu)化——也就是讓每一份硬件資源都用在它最擅長的地方——反而成了更現(xiàn)實的 Scaling 路徑。這篇文章適合誰看如果你正在做大規(guī)模模型訓(xùn)練被顯存和通信問題折磨過或者你在關(guān)注 MoE 架構(gòu)的訓(xùn)練效率想搞清楚為什么有些團(tuán)隊能在同等硬件下訓(xùn)出更大的模型又或者你只是對分布式訓(xùn)練的系統(tǒng)設(shè)計感興趣想理解“分離”這個思路背后的邏輯——那這篇內(nèi)容應(yīng)該能給你一些可以直接參考的東西。我會盡量把原理講透把實操中容易踩的坑點出來同時補(bǔ)充一些基于常見工程實踐的合理推斷畢竟訓(xùn)練系統(tǒng)的很多細(xì)節(jié)不會寫在論文里。2. 訓(xùn)練側(cè) PD 分離到底在分離什么2.1 計算密集與訪存密集的天然矛盾要理解訓(xùn)練側(cè) PD 分離的價值得先看清楚訓(xùn)練迭代里兩個階段的資源畫像差異有多大。前向傳播和反向傳播階段核心操作是矩陣乘法、卷積、注意力計算這些全是計算密集型任務(wù)。GPU 的 Tensor Core 在這個階段基本能跑滿算力利用率可以做到很高。但這個階段對顯存帶寬的需求相對溫和因為數(shù)據(jù)在計算單元里流轉(zhuǎn)不需要頻繁地大塊讀寫。參數(shù)更新階段就完全是另一回事了。以 Adam 優(yōu)化器為例每個參數(shù)需要維護(hù)一階矩和二階矩兩個狀態(tài)加上梯度本身和參數(shù)副本顯存占用是參數(shù)量的四倍起步。這個階段的操作是逐元素的加減乘除計算量很小但需要把海量的優(yōu)化器狀態(tài)從顯存里讀出來、算完再寫回去。這時候瓶頸完全在顯存帶寬上算力單元反而閑著。傳統(tǒng)訓(xùn)練把這兩個階段綁在一起意味著你買的 GPU 算力在參數(shù)更新階段大量閑置而你買的顯存帶寬在前向反向階段又用不滿。這就像讓一個短跑運(yùn)動員和一個舉重運(yùn)動員共用一套訓(xùn)練設(shè)備誰都沒法發(fā)揮最佳水平。2.2 分離之后資源怎么分配PD 分離的核心操作是把這兩個階段放到不同的硬件池或者不同的并行策略下去執(zhí)行。一種常見的做法是異構(gòu)硬件分配計算密集階段用高算力卡參數(shù)更新階段用大顯存、高帶寬的卡。這樣每一類硬件都用在刀刃上整體吞吐能提升不少。另一種做法是同構(gòu)硬件但不同并行策略前向反向用張量并行加流水線并行參數(shù)更新用數(shù)據(jù)并行或者 ZeRO 式的分片策略。這樣通信模式也分開了不會互相干擾。這里有個關(guān)鍵點分離之后兩個階段之間的數(shù)據(jù)傳遞成了新的瓶頸。前向反向算完的梯度要傳給參數(shù)更新階段參數(shù)更新完的新權(quán)重要傳回前向反向階段。這個傳遞如果走 PCIe 或者網(wǎng)絡(luò)延遲和帶寬都是問題。所以實際工程里要么用 NVLink 這種高帶寬互聯(lián)把兩個池子連起來要么在同一個節(jié)點內(nèi)做分離靠共享顯存或者高速緩存來傳遞數(shù)據(jù)。提示分離的粒度很關(guān)鍵。粗粒度分離整個訓(xùn)練任務(wù)級別實現(xiàn)簡單但靈活性差細(xì)粒度分離每個 micro-batch 級別效率高但工程復(fù)雜度陡增。大多數(shù)團(tuán)隊會從粗粒度開始試跑通了再往細(xì)粒度走。2.3 和推理側(cè) PD 分離的本質(zhì)區(qū)別雖然都叫 PD 分離但訓(xùn)練側(cè)和推理側(cè)的目標(biāo)函數(shù)完全不同。推理側(cè)的 PD 分離追求的是延遲和吞吐的平衡。Prefill 階段要盡快算完Decode 階段要穩(wěn)定輸出 token兩者對 SLA 的要求不一樣所以分開部署、分別擴(kuò)縮容。訓(xùn)練側(cè)沒有延遲這個概念追求的是單位時間內(nèi)完成的迭代次數(shù)或者說達(dá)到目標(biāo)精度所需的總時間。所以訓(xùn)練側(cè)的分離核心考量是資源利用率和通信效率而不是響應(yīng)時間。另一個區(qū)別是狀態(tài)管理。推理側(cè) Decode 階段要維護(hù) KV Cache訓(xùn)練側(cè)參數(shù)更新階段要維護(hù)優(yōu)化器狀態(tài)兩者都是顯存大戶但訓(xùn)練側(cè)的優(yōu)化器狀態(tài)是跨迭代持久化的不能像 KV Cache 那樣用完就丟。這意味著訓(xùn)練側(cè)的分離方案必須考慮狀態(tài)的存儲和遷移工程上更復(fù)雜。3. 為什么說這是更優(yōu)的 Scaling 路徑3.1 傳統(tǒng) Scaling 撞上的三堵墻在聊 PD 分離為什么更優(yōu)之前先看看傳統(tǒng) Scaling 路徑現(xiàn)在撞上了什么。第一堵墻是顯存墻。模型參數(shù)、梯度、優(yōu)化器狀態(tài)、激活值這四樣?xùn)|西加起來讓單卡能承載的模型規(guī)模很快到頂。即使用上 ZeRO 或者 FSDP 做分片通信開銷又會成為新瓶頸。千億參數(shù)模型在千卡集群上訓(xùn)練顯存碎片和 OOM 是家常便飯。第二堵墻是通信墻。數(shù)據(jù)并行的 AllReduce、張量并行的 AllGather 和 ReduceScatter、流水線并行的點對點通信這些通信操作在大規(guī)模下會吃掉大量時間。有實測數(shù)據(jù)顯示在萬卡級別通信時間占比可以超過 40%意味著你花大價錢買的算力有將近一半在等數(shù)據(jù)。第三堵墻是利用率墻。前面說的計算密集和訪存密集的矛盾導(dǎo)致硬件利用率在訓(xùn)練過程中劇烈波動。前向反向階段算力利用率高但帶寬閑置參數(shù)更新階段反過來。整體算下來有效利用率可能只有峰值的 50% 到 60%。PD 分離對這三堵墻都有緩解作用。顯存墻方面參數(shù)更新階段可以獨立做分片和卸載不占用前向反向的顯存預(yù)算。通信墻方面兩個階段的通信模式分開優(yōu)化不會互相疊加。利用率墻方面各階段用最適合的硬件和并行策略整體利用率能往上提一截。3.2 和 MoE 架構(gòu)的協(xié)同效應(yīng)MoE 架構(gòu)現(xiàn)在很熱但 MoE 的訓(xùn)練效率問題一直是個痛點。MoE 的核心是稀疏激活每個 token 只走部分專家這導(dǎo)致計算負(fù)載天然不均衡。在傳統(tǒng)訓(xùn)練框架下專家并行和專家負(fù)載均衡會帶來額外的通信和同步開銷。PD 分離和 MoE 放在一起看有很有意思的協(xié)同點。MoE 的前向反向階段計算集中在被激活的專家上這個階段適合用高算力卡做專家并行。參數(shù)更新階段所有專家的參數(shù)都要更新但每個專家的參數(shù)量相對小適合用數(shù)據(jù)并行或者分片策略批量處理。分離之后MoE 的負(fù)載不均衡問題可以在前向反向階段通過動態(tài)路由和專家復(fù)制來緩解而參數(shù)更新階段則回歸到規(guī)整的稠密更新效率更高。YOCO 這類架構(gòu)的出現(xiàn)也印證了這個方向。YOCO 把解碼器的自注意力和交叉注意力解耦本質(zhì)上也是一種分離思路——讓不同性質(zhì)的注意力計算走不同的路徑。雖然 YOCO 主要面向推理優(yōu)化但它背后的“分離不同性質(zhì)的計算”這個理念和訓(xùn)練側(cè) PD 分離是一脈相承的。3.3 KITE 等方案帶來的啟發(fā)KITE 這個關(guān)鍵詞在熱詞里出現(xiàn)值得單獨聊幾句。KITE 代表的是一類通信優(yōu)化與計算重疊的方案核心思路是在訓(xùn)練過程中把通信操作和計算操作在時間上錯開讓網(wǎng)絡(luò)傳輸和 GPU 計算并行進(jìn)行。PD 分離和 KITE 的思路可以疊加。分離之后參數(shù)更新階段的通信比如 AllReduce 同步梯度可以和前向反向階段的計算重疊。因為兩個階段在不同的硬件池或者不同的并行組里執(zhí)行參數(shù)更新階段的通信不會阻塞前向反向的計算流。這種跨階段的通信計算重疊比單階段內(nèi)的重疊空間更大效果也更明顯。實際工程里可以用雙緩沖或者多緩沖的機(jī)制來實現(xiàn)這種重疊。前向反向階段算完一個 micro-batch 的梯度立刻異步傳給參數(shù)更新池同時繼續(xù)算下一個 micro-batch。參數(shù)更新池收到梯度后開始更新更新完的權(quán)重再異步傳回去。只要緩沖區(qū)夠深兩個池子就能一直保持忙碌狀態(tài)。4. 落地訓(xùn)練側(cè) PD 分離的工程細(xì)節(jié)4.1 硬件拓?fù)渑c互聯(lián)選擇PD 分離的硬件拓?fù)湓O(shè)計直接決定了方案能不能跑出效果。核心原則是計算密集階段和訪存密集階段之間的數(shù)據(jù)通路帶寬要足夠高延遲要足夠低。如果兩個階段在同一個節(jié)點內(nèi)用共享顯存或者 NVLink 互聯(lián)是最理想的。NVLink 的帶寬在幾百 GB/s 到 TB/s 級別傳梯度和權(quán)重的開銷可以忽略不計。但同節(jié)點意味著硬件配置是固定的沒法針對兩個階段做差異化選型。如果兩個階段跨節(jié)點那就得看網(wǎng)絡(luò)了。InfiniBand 的帶寬在幾百 Gb/s 級別比 NVLink 低一個數(shù)量級傳大塊數(shù)據(jù)時延遲會明顯。這時候要么壓縮梯度比如用 FP16 或者 BF16 傳輸要么增加緩沖區(qū)深度來掩蓋延遲。實際測試中跨節(jié)點 PD 分離的收益很大程度上取決于網(wǎng)絡(luò)帶寬和訓(xùn)練任務(wù)的通信量比例。注意不要為了分離而分離。如果訓(xùn)練任務(wù)本身通信量就很大分離之后跨池通信可能比原來更糟。建議先用小規(guī)模實驗測一下通信開銷占比再決定是否上分離方案。4.2 梯度與權(quán)重的傳遞機(jī)制梯度從計算池傳到更新池權(quán)重從更新池傳回計算池這個雙向傳遞是 PD 分離的命脈。傳遞機(jī)制的設(shè)計要考慮三個維度傳輸時機(jī)、傳輸格式、傳輸可靠性。傳輸時機(jī)上有兩種策略。一種是同步傳遞計算池算完所有 micro-batch 的梯度后一次性傳給更新池。這種方式實現(xiàn)簡單但計算池在等待更新池處理時會空閑。另一種是異步流水線傳遞每算完一個 micro-batch 就傳一次更新池邊收邊算。這種方式能保持兩個池子都忙碌但對緩沖管理和一致性要求更高。傳輸格式上梯度可以用 FP16 或 BF16 壓縮權(quán)重更新也可以用低精度傳輸接收端再做精度恢復(fù)。實測下來BF16 傳輸對最終精度的影響很小但帶寬節(jié)省接近一半。如果網(wǎng)絡(luò)實在緊張還可以做梯度稀疏化只傳絕對值大的梯度小的直接置零。傳輸可靠性上分布式訓(xùn)練里節(jié)點故障是常態(tài)傳遞機(jī)制要能處理丟包、超時、節(jié)點掉線這些情況。通常的做法是加校驗和重傳同時維護(hù)一個全局的版本號確保計算池和更新池的權(quán)重版本一致。4.3 優(yōu)化器狀態(tài)的分片與卸載參數(shù)更新階段最大的顯存開銷來自優(yōu)化器狀態(tài)。以 Adam 為例每個參數(shù)要存一階矩和二階矩加上梯度本身顯存占用是參數(shù)量的三倍。如果模型有千億參數(shù)優(yōu)化器狀態(tài)就是幾千億個浮點數(shù)單卡根本放不下。PD 分離之后優(yōu)化器狀態(tài)可以獨立做分片和卸載。分片就是把優(yōu)化器狀態(tài)切到多張卡上每張卡只負(fù)責(zé)一部分參數(shù)的更新。卸載就是把不常用的狀態(tài)放到 CPU 內(nèi)存甚至 NVMe 盤上需要時再加載回來。這兩種技術(shù)結(jié)合能讓參數(shù)更新池的顯存需求大幅下降。具體操作上可以用 ZeRO-Offload 或者類似的框架把優(yōu)化器狀態(tài)和梯度都放到 CPU 側(cè)GPU 只負(fù)責(zé)計算。CPU 和 GPU 之間的傳輸走 PCIe帶寬雖然不如 NVLink但參數(shù)更新階段的計算量小傳輸時間可以接受。實測中ZeRO-Offload 能讓單卡訓(xùn)練 10B 參數(shù)模型成為可能代價是訓(xùn)練速度下降 20% 到 30%。如果 PD 分離能把前向反向階段的效率提上去整體吞吐反而可能持平甚至更高。5. 實操中容易踩的坑與排查思路5.1 負(fù)載不均衡導(dǎo)致的池子空轉(zhuǎn)PD 分離之后兩個池子的負(fù)載很難天然均衡。前向反向階段的計算量取決于模型結(jié)構(gòu)和 batch size參數(shù)更新階段的計算量取決于參數(shù)量和優(yōu)化器類型。如果一邊快一邊慢快的那個池子就會空轉(zhuǎn)等數(shù)據(jù)。我見過的一個典型案例是計算池用 8 卡 A100更新池用 4 卡 A100結(jié)果計算池每輪迭代要 200ms更新池只要 80ms更新池大部分時間在等梯度。反過來如果更新池卡太少計算池又會等權(quán)重。解決這個問題要么調(diào)整兩個池子的卡數(shù)比例要么在快的池子里插入一些額外任務(wù)比如數(shù)據(jù)預(yù)處理、檢查點保存來填滿空閑時間。排查負(fù)載不均衡最直接的方法是打時間戳日志記錄每個階段開始和結(jié)束的時間算一下兩個池子的忙碌比例。如果忙碌比例差距超過 20%就說明需要調(diào)整資源配比了。5.2 版本不一致引發(fā)的梯度錯位異步流水線傳遞里最容易出的問題是版本不一致。計算池用版本 N 的權(quán)重算梯度傳給更新池后更新池可能已經(jīng)更新到版本 N1 了這時候梯度就對不上了。輕則訓(xùn)練震蕩重則直接發(fā)散。解決這個問題的標(biāo)準(zhǔn)做法是加版本號校驗。每次傳遞梯度時帶上計算時用的權(quán)重版本號更新池收到后檢查版本號是否匹配。如果不匹配要么丟棄這個梯度要么用舊版本的權(quán)重重新計算。更穩(wěn)妥的做法是維護(hù)一個全局的版本隊列計算池和更新池都從隊列里取任務(wù)確保順序一致。提示版本不一致的問題在小規(guī)模實驗里不容易暴露因為延遲低、隊列淺。一旦上到大規(guī)模網(wǎng)絡(luò)延遲和隊列深度都會增加這個問題就會頻繁出現(xiàn)。建議在實驗階段就加上版本校驗別等到大規(guī)模訓(xùn)練時再補(bǔ)。5.3 通信瓶頸的定位與緩解PD 分離引入了額外的跨池通信如果這部分通信成為瓶頸整體收益就會被吃掉。定位通信瓶頸可以用 NCCL 的調(diào)試日志或者 PyTorch Profiler 的通信算子時間線看看跨池通信占總時間的比例。如果跨池通信占比超過 15%就需要考慮緩解措施了。常見的緩解手段包括增加緩沖區(qū)深度讓通信和計算重疊得更充分壓縮傳輸數(shù)據(jù)用低精度或者稀疏化減少通信量優(yōu)化網(wǎng)絡(luò)拓?fù)浒褍蓚€池子放在同一個交換機(jī)下減少跳數(shù)。還有一個容易被忽略的點是通信和計算的優(yōu)先級。在 GPU 上通信操作和計算操作共享 PCIe 或者 NVLink 帶寬如果通信操作優(yōu)先級太高會搶占計算操作的帶寬??梢酝ㄟ^ CUDA Stream 的優(yōu)先級設(shè)置來調(diào)整讓計算操作優(yōu)先通信操作在計算間隙執(zhí)行。6. 從實驗到生產(chǎn)我的幾點經(jīng)驗體會6.1 先做小規(guī)模驗證再上大規(guī)模PD 分離的收益和規(guī)模強(qiáng)相關(guān)。小規(guī)模下通信開銷占比低分離帶來的收益可能不明顯甚至因為額外的協(xié)調(diào)開銷而變負(fù)。大規(guī)模下通信和顯存問題突出分離的收益才會顯現(xiàn)。所以建議先在 8 卡或者 16 卡的小集群上做驗證確認(rèn)分離的邏輯跑通、版本一致、負(fù)載均衡再往大集群上遷移。小規(guī)模驗證的重點不是看吞吐提升而是看正確性和穩(wěn)定性。正確性方面對比分離方案和傳統(tǒng)方案的 loss 曲線確保收斂行為一致。穩(wěn)定性方面跑夠足夠多的迭代步數(shù)觀察有沒有版本錯位、內(nèi)存泄漏、節(jié)點掉線這些問題。6.2 監(jiān)控指標(biāo)要覆蓋兩個池子傳統(tǒng)訓(xùn)練的監(jiān)控指標(biāo)loss、學(xué)習(xí)率、梯度范數(shù)在 PD 分離方案下不夠用因為兩個池子是獨立運(yùn)行的需要分別監(jiān)控。計算池要關(guān)注算力利用率、顯存占用、前向反向耗時更新池要關(guān)注顯存帶寬利用率、優(yōu)化器狀態(tài)大小、更新耗時??绯赝ㄐ乓P(guān)注帶寬、延遲、隊列深度。這些指標(biāo)最好能在一個面板上統(tǒng)一展示方便定位瓶頸。如果計算池的算力利用率低可能是更新池供權(quán)重太慢如果更新池的帶寬利用率低可能是計算池供梯度太慢。兩個池子的指標(biāo)要對照著看才能找到真正的瓶頸。6.3 什么時候不該用 PD 分離PD 分離不是銀彈有些場景下用了反而更糟。如果模型規(guī)模不大單卡或者少量卡就能放下那分離帶來的協(xié)調(diào)開銷可能超過收益。如果訓(xùn)練任務(wù)對延遲敏感比如在線學(xué)習(xí)分離引入的跨池通信延遲可能不可接受。如果團(tuán)隊工程能力有限維護(hù)兩套并行策略和通信機(jī)制的復(fù)雜度可能拖垮迭代速度。我的建議是先評估當(dāng)前訓(xùn)練的瓶頸在哪里。如果瓶頸是算力不足那加卡比分離更直接。如果瓶頸是顯存不夠或者通信占比過高那 PD 分離值得一試。如果瓶頸是數(shù)據(jù)加載或者預(yù)處理那分離幫不上忙得從數(shù)據(jù)管道入手。6.4 后續(xù)可以探索的方向PD 分離和 MoE 的結(jié)合還有很多可以挖的空間。MoE 的專家路由可以動態(tài)調(diào)整讓計算密集階段把負(fù)載集中到少數(shù)高算力卡上參數(shù)更新階段再均勻分布到所有卡上。這種動態(tài)分離策略理論上能進(jìn)一步提升利用率。另一個方向是和檢查點機(jī)制的協(xié)同。PD 分離之后檢查點可以只保存參數(shù)更新池的狀態(tài)計算池的狀態(tài)是臨時的不需要持久化。這樣檢查點的大小和保存時間都能減少對大規(guī)模訓(xùn)練來說是個不小的收益。還有一個值得關(guān)注的點是分離粒度的自適應(yīng)調(diào)整。訓(xùn)練初期梯度變化大可能需要更頻繁的參數(shù)更新訓(xùn)練后期梯度變化小可以降低更新頻率把更多資源給前向反向。這種自適應(yīng)策略需要根據(jù)訓(xùn)練動態(tài)來調(diào)整兩個池子的資源配比工程上比較復(fù)雜但潛力很大。我在實際項目里試過把更新池的卡數(shù)從固定值改成動態(tài)調(diào)整根據(jù)梯度范數(shù)的變化率來決定增減卡數(shù)。效果是有但調(diào)度邏輯寫起來很繞而且頻繁增減卡會導(dǎo)致通信組重建開銷不小。后來改成粗粒度的階段調(diào)整——比如每訓(xùn)練 10% 的步數(shù)重新評估一次資源配比——就穩(wěn)定多了。這個經(jīng)驗分享出來是希望大家在追求極致效率的時候也考慮一下工程復(fù)雜度的邊界。