:大模型分布式訓(xùn)練顯存優(yōu)化與DeepSpeed實戰(zhàn)解析)
做分布式訓(xùn)練的人2019年之后幾乎沒人繞得開一個詞——ZeRO。十年時間從2015年大家還在為單卡放不下ResNet發(fā)愁到2025年動輒千億萬億參數(shù)的基礎(chǔ)模型訓(xùn)練穩(wěn)定跑在多機多卡集群上中間最關(guān)鍵的那個存儲器優(yōu)化方案就是微軟DeepSpeed團(tuán)隊提出的ZeRO系列。這個技術(shù)解決的是大模型訓(xùn)練里最樸素也最致命的問題顯存不夠用。它不改變模型結(jié)構(gòu)不改數(shù)據(jù)集不犧牲精度只是把一張卡裝不下的權(quán)重、梯度、優(yōu)化器狀態(tài)拆開存到多張卡上就能讓你用更少的卡訓(xùn)練更大的模型。這篇文章我把ZeRO這十年怎么一步步走到今天的事情捋清楚同時把里面三個Stage的原理、offload到CPU甚至NVMe的玩法、和PyTorch FSDP怎么選、以及我自己用DeepSpeed踩過的一堆坑都寫出來。不管你是剛接觸大模型訓(xùn)練的新手還是想把自己的訓(xùn)練腳本再壓一壓顯存的老手這篇應(yīng)該能讓你少走不少彎路。1. 從一個顯存焦慮的時代說起1.1 2015年的起點分布式訓(xùn)練還停留在“數(shù)據(jù)并行”2015年做深度學(xué)習(xí)的人手里拿的模型基本就是AlexNet、VGG、ResNet這個量級幾千萬到上億參數(shù)。那時候大家說的“分布式訓(xùn)練”默認(rèn)就是數(shù)據(jù)并行每張卡復(fù)制一份完整模型喂不同的batch然后梯度求平均同步更新。這種方案在當(dāng)時沒有任何問題因為模型本身沒多大單卡放得下大家焦慮的是訓(xùn)練速度不是顯存。轉(zhuǎn)折來自兩個方向。一個是Transformer在2017年橫空出世模型規(guī)模開始一路狂飆BERT 3億參數(shù)GPT-2 15億參數(shù)GPT-3 1750億參數(shù)。另一個是模型越來越大之后單純加卡已經(jīng)失效了——因為每張卡都要放一份完整模型副本顯存瓶頸是單卡上限鎖死的你加到100張卡單卡放不下還是放不下。我記得2019年調(diào)GPT-2大小的模型時一臺8卡V100每卡32GB已經(jīng)要精打細(xì)算才能塞進(jìn)去等到175B的模型出來純靠數(shù)據(jù)并行是徹底沒戲了。這時候行業(yè)里有兩條路線開始分叉一條是模型并行/流水線并行把模型的不同層切到不同卡上比如Megatron-LM和GPipe另一條是微軟在2019年底放出的ZeRO走的是“狀態(tài)分區(qū)”路線。后來大家都知道了ZeRO及其衍生技術(shù)成了訓(xùn)練超大模型的事實標(biāo)準(zhǔn)之一。1.2 一張表看懂顯存到底被誰吃掉了先來算一筆賬。假設(shè)你在訓(xùn)練一個7B參數(shù)模型用混合精度FP16參數(shù)FP32優(yōu)化器跑AdamW。一張卡上要裝的東西分四塊數(shù)據(jù)對象每個參數(shù)占用7B模型總占用說明FP16參數(shù)副本2字節(jié)14GB前向和反傳用的權(quán)重FP16梯度2字節(jié)14GB反傳時累積的梯度FP32主權(quán)重4字節(jié)28GBAdam更新時用的精確權(quán)重Adam動量m4字節(jié)28GB一階矩估計Adam方差v4字節(jié)28GB二階矩估計合計112GB。這是不算激活值、不算臨時緩沖、不算通信緩沖的裸數(shù)據(jù)。一張A100 80GB根本放不下V100 32GB更是想都別想。再加激活值activation7B模型在seq len 2048、batch 4的場景下動輒還要額外幾十GB。這就是為什么“單卡放不下”不是一句空話而是明明白白的數(shù)學(xué)問題。ZeRO做的事情概括成一句話以前是每張卡把上面這堆東西全部復(fù)制一份現(xiàn)在是大家合起來平攤各存一份分片用的時候再互相拼起來。1.3 為什么“切一切”就能解決顯存問題這里面的直覺可以用倉庫來類比。原來每個倉庫管理員GPU都自己囤一整箱貨完整模型副本地方不夠了就堆不下。ZeRO的思路是一個模型的狀態(tài)不是每次都全量用到的尤其在數(shù)據(jù)并行場景下每張卡跑的都是同一個模型的同一層只是輸入數(shù)據(jù)不同。既然你們的計算模式完全一樣那這些狀態(tài)完全可以分區(qū)保存前向反傳到哪一層、需要哪個參數(shù)再臨時把那一片拉過來。這個“需要時再拼裝”的思路就是ZeRO-DP的核心。它的好處是幾乎不改變原有的數(shù)據(jù)并行訓(xùn)練流程不引入復(fù)雜的矩陣切分邏輯那是張量并行的事所以工程上落地特別快。2. ZeRO核心設(shè)計逐層拆解2.1 三個Stage從只分優(yōu)化器到連參數(shù)都分ZeRO論文里定義了三個遞進(jìn)的Stage對應(yīng)著分片覆蓋的范圍逐漸擴大Stage 1P_os只把優(yōu)化器狀態(tài)分片。Adam的三塊FP32狀態(tài)主權(quán)重、m、v每張卡只存1/N參數(shù)和梯度仍然每卡全量持有。這一步已經(jīng)能省大約75%的優(yōu)化器相關(guān)顯存因為Adam那12字節(jié)/參數(shù)是大頭。Stage 2P_osg在Stage 1基礎(chǔ)上把梯度也分片。反傳過程中梯度是逐步產(chǎn)生的每張卡只保留自己負(fù)責(zé)的那片梯度然后reduce-scatter匯總。省掉的顯存進(jìn)一步增加。Stage 3P_osgp連模型參數(shù)本身也分片。前向傳播到某一層時所有卡把這一層參數(shù)通過all-gather拼出來算完再丟棄。這是省顯存最徹底、通信開銷也最大的階段。拿7B模型在8卡上跑舉例用Stage 3每卡只需要存大約112GB 激活值/ 8 ≈ 14GB的靜態(tài)數(shù)據(jù)A100 80GB跑起來非常寬裕甚至可以塞更大的batch。這就是很多人對著A100說“8卡能訓(xùn)7B”的底氣來源。2.2 通信量為什么沒有爆炸一次數(shù)學(xué)上很劃算的交換有人會問參數(shù)、梯度、優(yōu)化器狀態(tài)都分開了每次都要拼裝通信成本是不是高到?jīng)]法用答案是沒有至少沒有高到離譜。DDP數(shù)據(jù)并行每個step要做一次全量梯度的all-reduce通信量大約等于全量梯度的2倍。ZeRO Stage 2用reduce-scatter做梯度歸約再用all-gather把更新后的權(quán)重廣播回去通信量和DDP基本持平但顯存占用大幅下降。Stage 3因為參數(shù)也要按需gather前向一次、反傳一次再加梯度歸約通信量大約是DDP的1.5倍左右。換句話說你是用增加的這0.5倍通信量換來了從“單卡放不下”到“隨便放”的質(zhì)變。在大規(guī)模集群上這0.5倍通信通??梢酝ㄟ^NVLink、IB、梯度重疊通信來消化絕對收益遠(yuǎn)大于代價。這也是為什么ZeRO初期最推薦Stage 2它幾乎不增加通信壓力但已經(jīng)把優(yōu)化器狀態(tài)和梯度這兩個最大頭解決了8卡訓(xùn)7B、64卡訓(xùn)30B這種場景完全夠用。只有模型大到連參數(shù)副本都放不下時才需要上Stage 3。2.3 ZeRO-R除了模型狀態(tài)還有三塊隱藏的顯存開銷ZeRO論文里除了ZeRO-DP之外還專門講了ZeRO-R負(fù)責(zé)處理另外三塊容易被忽略的開銷激活值、臨時緩沖、顯存碎片。激活值通過“分區(qū)激活”partition activations來處理每張卡只存自己負(fù)責(zé)那部分激活用的時候再跨卡gather。這和重計算activation checkpointing是互補的重計算以2倍前向計算換顯存分區(qū)激活以通信換顯存。兩者疊加效果更好也是為什么DeepSpeed配置里經(jīng)常同時開這兩項。臨時緩沖主要指all-gather和reduce-scatter的通信緩沖。DeepSpeed的做法是分配一個“恒定緩沖”constant buffer大小可以配置默認(rèn)可能幾百MB到1GB用來避免頻繁申請釋放造成的不穩(wěn)定和碎片。顯存碎片則通過內(nèi)存對齊和合并釋放來緩解。很多人在Stage 3下遇到莫名OOM其實不是模型太大而是碎片太嚴(yán)重。3. 從DDP到ZeRO的遷移實操3.1 工具選型DeepSpeed還是PyTorch FSDP2023年之后PyTorch原生集成了FSDPFully Sharded Data Parallel效果對標(biāo)ZeRO Stage 3社區(qū)適配也越來越好。所以現(xiàn)在做技術(shù)選型時很多人會糾結(jié)。我的經(jīng)驗是維度DeepSpeed ZeROPyTorch FSDP功能完整度支持Stage 1/2/3、CPU/NVMe offload、ZeRO主要對標(biāo)Stage 3offload到CPU較成熟配置復(fù)雜度需要寫JSON配置靈活但坑多參數(shù)化API和PyTorch生態(tài)貼合社區(qū)資料老牌訓(xùn)練大模型的案例最多更新快PyTorch 2.x之后已成主流與Megatron結(jié)合有DeepSpeed-Megatron組合方案需要自己搭橋上手速度中等較快我的建議如果從頭寫一個新項目PyTorch 2.0以上推薦直接用FSDP代碼侵入小如果是要跑成熟的大模型訓(xùn)練框架比如DeepSeek、LLaMA系微調(diào)、各類lora訓(xùn)練腳本大概率已經(jīng)內(nèi)置了DeepSpeed配置那就直接用DeepSpeed別折騰遷移。兩者原理相通很多經(jīng)驗可以互相套用。3.2 一份能直接跑的DeepSpeed配置逐項講清楚下面是我在單機多卡訓(xùn)7B模型時常用的一份配置按我的經(jīng)驗注釋一下{ train_batch_size: 64, gradient_accumulation_steps: 4, optimizer: { type: AdamW, params: { lr: 1e-5, betas: [0.9, 0.95], eps: 1e-8 } }, zero_optimization: { stage: 2, offload_optimizer: { device: cpu, pin_memory: true }, contiguous_gradients: true, overlap_comm: true, reduce_bucket_size: 5e8, stage3_prefetch_bucket_size: 5e8, stage3_param_persistence_threshold: 1e6, stage3_max_live_parameters: 1e9, stage3_max_reuse_distance: 1e9 }, activation_checkpointing: { partition_activations: true } }幾個容易踩坑的點stage 2還是37B模型在8卡A100上stage 2已經(jīng)足夠通信快顯存占用也穩(wěn)。只有當(dāng)你想把batch開得特別大或者卡特別少比如2卡跑7B時才需要stage 3。offload_optimizer到CPU這個配置在stage 2下也能開。它把優(yōu)化器狀態(tài)放到系統(tǒng)內(nèi)存顯存省得很猛但訓(xùn)練速度會明顯下降因為每次更新都要走PCIe。如果是純追求訓(xùn)練速度的集群建議別開如果是自己單機調(diào)試、卡顯存吃緊真香。reduce_bucket_size這個值過小會導(dǎo)致通信頻繁切小塊過大又占顯存。5e8500MB是我常用的平衡點。你可以盯一眼nvidia-smi和NCCL日志再調(diào)。activation_checkpointing.partition_activations對應(yīng)ZeRO-R的激活分區(qū)對長序列訓(xùn)練效果非常明顯能和重計算疊加。3.3 訓(xùn)練腳本改動清單如果你的腳本之前用的是PyTorch DDP遷移到DeepSpeed只需三步第一命令行啟動方式從python -m torch.distributed.launch換成deepspeed --num_gpus8 train.py --deepspeed ds_config.json。第二訓(xùn)練邏輯中把model DDP(model)換成model_engine, optimizer, _, _ deepspeed.initialize(modelmodel, model_parametersparams, configds_config)后續(xù)step時用model_engine.step()替代optimizer.step()。第三loss的scaling由DeepSpeed接管不需要你手工做混合精度相關(guān)操作除非你自己寫了AMP。這里有個很關(guān)鍵的經(jīng)驗DeepSpeed初始化之后原模型的forward/backward接口不變但backward需要傳入loss給model_engine.backward(loss)很多人第一次遷移會漏掉這一步導(dǎo)致梯度根本不算。還有如果你原來用torch.cuda.amp.GradScaler遷移后建議刪掉否則兩個scaler會互相打架出現(xiàn)莫名其妙的loss不下降。3.4 容量估算上集群前先算一筆賬上大規(guī)模集群前我強烈建議先做一次顯存估算。以7B為例6張A100 80GB、跑stage 2、不開offload每卡可用顯存里靜態(tài)占用大約30GB左右剩下的空間可以用來放激活值和更大的batch。反過來如果你的模型是30B8卡stage 3起步激活值一定要開重計算offload可以視帶寬情況決定。一個粗略公式經(jīng)驗值非精確Stage 2每卡占用約為FP16參數(shù) FP16梯度 Adam狀態(tài)/N再乘1.2的安全系數(shù)。Stage 3每卡占用約為FP16參數(shù) FP16梯度 Adam狀態(tài)/ N再乘1.3。以7B、8卡為例stage 2約141484/8×1.2 ≈ 42GBstage 3約112/8×1.3 ≈ 18GB。算出來的數(shù)如果已經(jīng)接近顯存上限先砍激活值、砍batch再考慮開offload不要一上來就把所有手段都堆上。4. 十年路線圖從Stage 1到ZeRO三代4.1 2015-2025關(guān)鍵時間線要理解ZeRO這十年不能只看它2019年橫空出世那一瞬間。我把這條線按年份拉出來2015-2017數(shù)據(jù)并行模型并行并行著走顯存問題開始顯現(xiàn)但還沒到不可收拾。2018-2019GPT-2/XLNet大規(guī)模模型出現(xiàn)Megatron-LM和張量并行成為熱點。微軟研究院啟動ZeRO項目目標(biāo)是“任何人用任何GPU都能訓(xùn)任意大的模型”。2019年底ZeRO論文上線提出ZeRO-DP和ZeRO-R首次在論文層面實現(xiàn)“train 100B模型”的理論可行性。2020年初DeepSpeed開源附帶ZeRO Stage 1/2實現(xiàn)。社區(qū)第一次可以在普通多卡機器上跑BERT-Large、GPT-2 1.5B成本驟降。2021年ZeRO-Offload發(fā)布把優(yōu)化器狀態(tài)和梯度卸載到CPU內(nèi)存隨后ZeRO-Infinity擴展到NVMe最大亮點是單卡也能訓(xùn)練超大模型。2023年ZeRO第一代發(fā)布引入分區(qū)通信partitioned communication、量化通信quantized communication和低精度參數(shù)解決Stage 3通信瓶頸。2024-2025年ZeRO第二代/第三代迭代低精度優(yōu)化器狀態(tài)、混合精度的通信、對推理場景的ZeRO-Inference支持逐步落地。同時PyTorch FSDP全面成熟ZeRO思想被吸收進(jìn)各家框架。“ZeRO十年演進(jìn)”嚴(yán)格說ZeRO本身是2019年才誕生的但2015年這個起點非常適合理解它解決的是什么時代的問題顯存從“不夠大的煩惱”變成了“決定模型規(guī)模的硬邊界”。4.2 ZeRO-Offload和ZeRO-Infinity把顯存邊界推到CPU和硬盤ZeRO-Offload的思路非常樸素既然GPU顯存不夠CPU內(nèi)存通常大得多服務(wù)器輕松512GB到1TB那就把不經(jīng)常用、但占地方的東西放到CPU上。它的最佳實踐是優(yōu)化器狀態(tài)放CPU梯度也放CPU參數(shù)留在GPU反向傳播和參數(shù)更新在CPU上異步執(zhí)行。實測下來對于幾十B級別模型CPU offload能讓人用少量GPU卡跑起來代價是速度明顯下降。后來ZeRO-Infinity把這一套擴展到了NVMe SSD允許優(yōu)化器狀態(tài)和梯度直接落盤這就把“單卡能訓(xùn)的模型上限”推到了接近無限大——理論上只要你有足夠的CPU內(nèi)存和硬盤單卡也能訓(xùn)170B參數(shù)模型只是速度慢到懷疑人生。這個方案適合調(diào)試、臨時跑通不是生產(chǎn)環(huán)境的首選。實際使用中offload到CPU的配置我之前已經(jīng)給了NVMe offload要再加一層offload_optimizer: { device: nvme, nvme_path: /mnt/nvme, buffer_count: 4, fast_readwrite: true }注意nvme_path必須是一個真實的本地NVMe設(shè)備路徑不能是網(wǎng)絡(luò)盤否則延遲會高到完全跑不動。fast_readwrite開啟后DeepSpeed會為每個buffer分配一塊固定內(nèi)存做DMA對這種場景提升很大。4.3 ZeRO三代通信量的正面硬剛ZeRO Stage 3最大的痛點是通信量是DDP的1.5倍所以在千卡集群上通信優(yōu)化比顯存優(yōu)化更能決定訓(xùn)練速度。ZeRO就是沖著這個去的。第一代ZeRO在2023年放出三個關(guān)鍵機制分區(qū)分組通信把大world size拆成小分組分層all-gather、量化通信把通信數(shù)據(jù)從FP16降到INT8傳輸量減半精度損失通過量化感知的梯度壓縮來彌補、低精度優(yōu)化器狀態(tài)直接砍掉一部分優(yōu)化器狀態(tài)的存儲需求。第二代和第三代繼續(xù)在通信分組策略、混合精度調(diào)度以及GPU/NVMe分層offload上做文章??雌饋韽?fù)雜落地其實很透明DeepSpeed的配置里打開comms_logger和zero_quantized_ gradients相關(guān)開關(guān)或者直接用啟動參數(shù)--zero-stage3 --zero-quantized-gradients。實測在100Gbps網(wǎng)絡(luò)的多機場景下量化通信能帶來20%-40%的端到端提速。但是注意如果你的網(wǎng)絡(luò)帶寬已經(jīng)很充裕比如全部NVLink互聯(lián)這個優(yōu)化收益不大反而可能因為量化反量化增加GPU計算。4.4 和Megatron、FSDP到底什么關(guān)系一張圖理清這里我經(jīng)常被人問ZeRO和Megatron不是沖突嗎不是。它們負(fù)責(zé)的維度完全不同方案切分維度適合場景ZeRO-DPstage 3數(shù)據(jù)并行中的模型狀態(tài)分片同層參數(shù)多模型大但單層能裝下集群通信條件好Zhang量并行Megatron按矩陣維度切分每層參數(shù)單層矩陣過大如超大hidden size流水線并行PP按層切分層數(shù)特別多機器間帶寬一般FSDP和ZeRO-DP類似PyTorch生態(tài)內(nèi)想快速實現(xiàn)實踐中訓(xùn)練一個175B模型典型組合是數(shù)據(jù)并行ZeRO Stage 2或3張量并行8路流水線并行若干階段激活重計算。單一方案解決不了所有問題ZeRO負(fù)責(zé)的是“讓每張卡不存冗余狀態(tài)”這部分和Megatron的張量切分是互補的。5. 實戰(zhàn)中的坑與排查速查5.1 六個高頻報錯/現(xiàn)象及處理建議OOM已經(jīng)在Stage 3了還是顯存爆了。優(yōu)先檢查三件事reduce_bucket_size是不是太大激活值有沒有開重計算和分區(qū)offload_param/offload_optimizer是否真的生效可以在DeepSpeed啟動日志里確認(rèn)。A100上直接看Deepspeed info輸出它會打印每塊顯存分配明細(xì)。訓(xùn)練卡死或異常慢NCCL超時。多機場景最常見通常和網(wǎng)絡(luò)環(huán)境有關(guān)容器里沒正確指定NCCL_SOCKET_IFNAME或者NCCL_P2P_DISABLE設(shè)置不對。我一般先設(shè)NCCL_DEBUGINFO看卡在哪一次集合通信上再逐項調(diào)。加載checkpoint后loss震蕩或完全亂掉。一般是optimizer state沒有正確加載或者resume_from_checkpoint路徑給錯DeepSpeed只會從checkpoint目錄里的zero_pp_rank_*.pt恢復(fù)確保路徑統(tǒng)一。推理或agent調(diào)用時報unknown model。一些模型服務(wù)或agent框架在啟動時會把模型名映射到實際權(quán)重文件如果填的模型名和注冊名不一致就會出現(xiàn)類似unknown model: xxx的報錯。這跟ZeRO本身關(guān)系不大多半是配置里模型的name、路徑、版本號沒對齊。查的時候先看配置里model_name_or_path和框架自帶的模型注冊表別一上來就懷疑顯存。開啟offload后訓(xùn)練變慢到無法接受。這是“顯存換速度”的必然代價。建議只offload優(yōu)化器不要連參數(shù)也offload開pin_memory確保數(shù)據(jù)加載不吃CPU有NVMe的話優(yōu)先用NVMe而不是SATA SSD。梯度更新前后loss完全不動。檢查你是不是把model_engine.backward(loss)寫成了loss.backward()以及有沒有在調(diào)用step之前手動調(diào)用了zero_grad。DeepSpeed的step內(nèi)部會處理梯度清空手動zero_grad有時會把累積的梯度清掉。5.2 調(diào)優(yōu)的幾個“土辦法”這些不是我編的是每次新上一個訓(xùn)練任務(wù)我都會在日志里用一組固定手段觀測先看nvidia-smi里顯存和功耗。顯存用了90%以上但功耗只有一半說明通信或數(shù)據(jù)加載瓶頸不是計算瓶頸??碞CCL日志里的帶寬數(shù)字。如果遠(yuǎn)低于預(yù)期比如100Gbps網(wǎng)絡(luò)上只有20Gbps先查是不是走錯了網(wǎng)卡。開overlap_comm后再看吞吐如果提升不明顯說明通信本來就不是瓶頸別再做量化通信了。用小batch跑通再逐漸放大找到顯存和吞吐的平衡點。這一步能幫你判斷是靜態(tài)顯存瓶頸還是激活值瓶頸。5.3 搜索時別把ZeRO和那些“zero”搞混了寫這篇文章時我順手搜了一下發(fā)現(xiàn)現(xiàn)在搜ZeRO出來的內(nèi)容會被各種其他項目稀釋嵌入式里的荔枝派Zero、Go生態(tài)里的go-zero框架、無人機里的Zero Omega、甚至一些AI agent工具報錯里的unknown model提示。它們和深度學(xué)習(xí)顯存優(yōu)化完全是兩碼事。想查技術(shù)資料時建議關(guān)鍵詞帶上下文集ZeRO DeepSpeed、ZeRO stage 3、ZeRO offload或者直接去DeepSpeed官方文檔和論文源碼里找能省不少時間。6. 我的一些個人體會回頭去看這十年ZeRO最值得學(xué)習(xí)的不是某個具體的顯存優(yōu)化技巧而是“發(fā)現(xiàn)問題、量化問題、系統(tǒng)解決問題”的思路。它沒有發(fā)明新的并行范式而是在已有數(shù)據(jù)并行框架里把“冗余存儲”這四個字摳到了極致然后把通信成本控制在一個可接受的范圍內(nèi)最終讓大規(guī)模訓(xùn)練從“大廠專用”變成了“工程師可及”的能力。我自己的實踐體會是拿到一個訓(xùn)練任務(wù)不要一上來就無腦Stage 3 offload先從Stage 2騎一遍用nvidia-smi看真實顯存分布再結(jié)合模型大小和卡數(shù)決定要不要升級Stage 3。多數(shù)時候Stage 2就夠用而且調(diào)參成本低得多。真正上到Stage 3時一定先把通信環(huán)境測好、NCCL設(shè)置調(diào)好否則你會被各種超時和卡死折磨到懷疑人生。最后再分享一個小技巧在DeepSpeed訓(xùn)練腳本里加一行torch.cuda.memory_summary()跑兩個step后輸出顯存分配詳情你能看到哪些buffer占了大頭。很多看起來玄乎的OOM一查就原形畢露。這比在網(wǎng)上盲搜報錯要靠譜得多。