練實(shí)戰(zhàn):從單卡到分布式提速3.7倍)
先說(shuō)結(jié)論P(yáng)yTorch DDP 這個(gè)坑我踩了不少但一旦把原理和調(diào)用方式理清實(shí)測(cè)下來(lái)是真的快。我手頭一個(gè) ResNet 在單卡上要跑接近 10 小時(shí)的任務(wù)改成 DDPDistributedDataParallel后4 張卡只用了不到 3 小時(shí)加速比接近 3.7 倍。這個(gè)增速不是玄學(xué)靠的是 DDP 的梯度同步機(jī)制和正確的超參數(shù)配置。今天就把我從單卡腳本改造成多卡訓(xùn)練的完整思路、代碼細(xì)節(jié)和踩坑過(guò)程整理出來(lái)給正在折騰 DDP 的朋友一個(gè)可以直接抄作業(yè)的參考。1. DDP 為什么快先弄清楚它解決的核心問(wèn)題1.1 單卡訓(xùn)練的真正瓶頸在哪兒很多朋友問(wèn)我訓(xùn)練慢是不是因?yàn)轱@卡不行其實(shí)單卡訓(xùn)練時(shí)GPU 的算力通常沒(méi)有被榨干。我見過(guò)不少項(xiàng)目是模型不算大、數(shù)據(jù)也不算多但訓(xùn)練時(shí)間就是上不去。核心瓶頸往往不在計(jì)算而在數(shù)據(jù)流水線、CPU 預(yù)處理、以及單卡顯存對(duì) batch size 的限制。當(dāng)你把 batch size 壓小去適配顯存時(shí)每個(gè) step 的梯度噪聲會(huì)變大收斂反而更慢當(dāng)你把數(shù)據(jù)增強(qiáng)、解碼這類操作放到 CPU 上時(shí)GPU 又經(jīng)常空轉(zhuǎn)等數(shù)據(jù)。這就是典型的算不快、喂不飽問(wèn)題。多卡分布式訓(xùn)練解決的就是兩件事一是把數(shù)據(jù)分到多張卡上并行處理攤薄單卡的壓力二是把梯度同步的開銷壓到足夠低讓多卡協(xié)作接近單卡效率的線性疊加。DDP 之所以能成為 PyTorch 實(shí)測(cè)最快的分布式方案不是因?yàn)樗鼤?huì)什么魔法而是它在設(shè)計(jì)上把所有能省的通信量都省了把數(shù)據(jù)并行和梯度同步的配合做到了很干凈。1.2 Ring AllReduce梯度聚合是怎么做到低開銷的要理解 DDP 為什么快必須先理解梯度是怎么在多卡之間同步的。數(shù)據(jù)并行模式下每張卡都持有完整模型副本各自用一部分?jǐn)?shù)據(jù)做前向和反向算出來(lái)的梯度是局部梯度。要讓所有卡保持一致的模型參數(shù)就必須把所有局部梯度相加取平均再讓每張卡用自己的優(yōu)化器更新參數(shù)。最笨的做法是搞一個(gè)主節(jié)點(diǎn)收集所有梯度、求和、廣播回去這就是中心化 AllReduce。通信量是 O(2N)N 是卡數(shù)卡越多主節(jié)點(diǎn)瓶頸越嚴(yán)重。DDP 用的是 Ring AllReduce所有 GPU 首尾相連成一個(gè)環(huán)把梯度切成 N 份每一輪每張卡只和自己相鄰的節(jié)點(diǎn)交換一份數(shù)據(jù)N-1 輪之后所有節(jié)點(diǎn)就持有了全局平均梯度。通信量是 O(2(N-1)/N)當(dāng)卡數(shù)很多時(shí)這個(gè)方案對(duì)帶寬的利用率高得多也不會(huì)被某一張卡拖死。我自己的理解是中心化方案像辦公室所有人把文件都交給一個(gè)前臺(tái)妹子再由她分發(fā)前臺(tái)再快也是瓶頸Ring AllReduce 像同事們圍成圈傳文件每個(gè)人只和左右鄰居交接總量一樣但分?jǐn)偟矫總€(gè)人頭上就很輕松。這也是為什么 DDP 在 8 卡、16 卡甚至跨機(jī)場(chǎng)景下提速依然能保持接近線性的核心原因。1.3 DDP 和 DataParallel 的區(qū)別直接決定了速度上限很多人把 DDP 誤以為是 DataParallelDP的改良版其實(shí)二者在設(shè)計(jì)上有本質(zhì)區(qū)別。DP 是單進(jìn)程多線程模型有一個(gè)主 GPU 負(fù)責(zé)匯總梯度并廣播而且 Python 的 GIL 還會(huì)讓多個(gè)線程爭(zhēng)搶解釋器資源多張卡很難真正跑滿。DDP 是真正的多進(jìn)程模型每個(gè)進(jìn)程綁定一張卡擁有獨(dú)立的 Python 解釋器、獨(dú)立的模型副本進(jìn)程間只通過(guò)梯度 AllReduce 通信完全繞開了 GIL 的干擾。我用一個(gè)實(shí)際對(duì)比說(shuō)明差距。同樣在 4 卡機(jī)器上訓(xùn)練同一個(gè)模型DP 的加速比大概只有 2.8 到 3.0 倍而且主卡的顯存明顯偏高、其他卡利用率參差不齊。換成 DDP 后四張卡的利用率非常均勻加速比直接到了 3.6 倍以上。如果你的環(huán)境允許直接用 DDP 就好DP 只適合臨時(shí)驗(yàn)證小模型生產(chǎn)級(jí)訓(xùn)練請(qǐng)無(wú)條件選擇 DDP。維度DataParallel (DP)DistributedDataParallel (DDP)進(jìn)程模型單進(jìn)程多線程多進(jìn)程每進(jìn)程綁定一張卡梯度同步主卡匯總再?gòu)V播Ring AllReduce 對(duì)等聚合GIL 影響有無(wú)負(fù)載均衡主卡容易成瓶頸多卡天然均衡適用場(chǎng)景小模型、臨時(shí)驗(yàn)證多卡/多機(jī)、生產(chǎn)訓(xùn)練2. 手把手把單卡訓(xùn)練腳本改成 DDP2.1 用 torchrun 做標(biāo)準(zhǔn)啟動(dòng)別再手動(dòng)傳參了DDP 改造的第一步是啟動(dòng)方式。PyTorch 官方推薦的啟動(dòng)工具是torchrun它會(huì)自動(dòng)幫我們注入一系列環(huán)境變量包括全局進(jìn)程編號(hào) RANK、當(dāng)前節(jié)點(diǎn)上的進(jìn)程編號(hào) LOCAL_RANK、總進(jìn)程數(shù) WORLD_SIZE 等。你只需要在命令行里指定用幾張卡torchrun --nproc_per_node4 --master_port29500 train.pytorchrun做的事情非常多包括進(jìn)程拉起、失敗重啟、多機(jī)統(tǒng)一入口協(xié)調(diào)等。早期不少人是在代碼里手動(dòng)mp.spawn()或者自己設(shè)置環(huán)境變量再subprocess.Popen啟動(dòng)問(wèn)題非常多。如果你是在單機(jī)多卡上跑直接用torchrun就對(duì)了多機(jī)場(chǎng)景下再額外加--nnodes、--node_rank和--master_addr這類參數(shù)。我在第一次改造時(shí)犯過(guò)一個(gè)典型錯(cuò)誤啟動(dòng)命令寫了torchrun --nproc_per_node4結(jié)果每張卡上都跑了完整的數(shù)據(jù)集相當(dāng)于每張卡數(shù)據(jù)沒(méi)分只是獨(dú)立訓(xùn)練了四遍。問(wèn)題就出在缺少 DistributedSampler 上后面會(huì)詳細(xì)說(shuō)。2.2 rank、local_rank、world_size 這些參數(shù)到底代表什么新手第一次看到這些英文參數(shù)基本都會(huì)懵。我試著用最簡(jiǎn)單的方式解釋world_size參與并行訓(xùn)練的進(jìn)程總數(shù)也就是 GPU 數(shù)量。單機(jī) 4 卡就是 4兩機(jī)各 4 卡就是 8。rank全局進(jìn)程編號(hào)從 0 到 world_size-1。它可以理解為你是第幾個(gè)到達(dá)訓(xùn)練室的人在保存模型、打印日志、判主節(jié)點(diǎn)時(shí)特別有用。local_rank當(dāng)前機(jī)器內(nèi)部的進(jìn)程編號(hào)。在兩機(jī)各 4 卡場(chǎng)景下節(jié)點(diǎn) 0 上的進(jìn)程 local_rank 是 0-3節(jié)點(diǎn) 1 上的進(jìn)程 local_rank 也是 0-3。這個(gè)參數(shù)的值會(huì)直接決定進(jìn)程綁定到哪張物理 GPU。MASTER_ADDR和MASTER_PORT分布式通信中的導(dǎo)演rank 0 進(jìn)程負(fù)責(zé)協(xié)調(diào)其他進(jìn)程之間的連接關(guān)系其他進(jìn)程需要知道它的地址和端口才能完成握手。在代碼里我習(xí)慣這樣獲取這些值import os import torch import torch.distributed as dist def init_process_group(): dist.init_process_group(backendnccl, init_methodenv://) rank dist.get_rank() world_size dist.get_world_size() local_rank int(os.environ[LOCAL_RANK]) torch.cuda.set_device(local_rank) return rank, world_size, local_rank注意在沒(méi)有額外指定時(shí)init_process_group 會(huì)用env://方式自動(dòng)讀取環(huán)境變量中的 RANK、WORLD_SIZE、MASTER_ADDR、MASTER_PORT所以在 torchrun 的配合下這四行代碼就夠了。2.3 DistributedSampler多卡分?jǐn)?shù)據(jù)最容易出錯(cuò)的一步如果只做 init 和模型包裝不做數(shù)據(jù)切分你訓(xùn)練時(shí)的表現(xiàn)就是四張卡各看各的數(shù)據(jù)梯度各算各的模型永遠(yuǎn)不會(huì)收斂到一致狀態(tài)。正確做法是給 DataLoader 掛一個(gè)DistributedSampler由它負(fù)責(zé)把數(shù)據(jù)集按進(jìn)程數(shù)均勻切分每個(gè)進(jìn)程只拿到屬于自己的那部分。from torch.utils.data import DataLoader, DistributedSampler sampler DistributedSampler(dataset, num_replicasworld_size, rankrank, shuffleTrue) loader DataLoader(dataset, batch_size64, samplersampler, num_workers4, pin_memoryTrue)這里有兩個(gè)細(xì)節(jié)極容易踩坑。第一batch_size是每個(gè)進(jìn)程的 batch size全局 batch size 實(shí)際上是per_process_batch_size * world_size。所以如果你原來(lái)單卡跑 64切到 4 卡后想保持全局 64就要把每卡 batch size 改成 16否則相當(dāng)于全局變成 256模型收斂行為會(huì)完全不同。第二每個(gè) epoch 開始前必須調(diào)用sampler.set_epoch(epoch)否則 DistributedSampler 內(nèi)部的隨機(jī)打亂順序不會(huì)改變每個(gè) epoch 的數(shù)據(jù)劃分都是同一個(gè)順序模型訓(xùn)練會(huì)退化。2.4 一份可以直接跑的完整示例代碼下面這份代碼我盡量保持了最小化適合拿來(lái)做改造的起點(diǎn)。它用 MNIST 做演示實(shí)際項(xiàng)目中你只需要替換模型、數(shù)據(jù)和訓(xùn)練邏輯即可。import os import torch import torch.nn as nn import torch.distributed as dist from torch.utils.data import DataLoader, DistributedSampler from torch.nn.parallel import DistributedDataParallel from torchvision import datasets, transforms def train(): dist.init_process_group(backendnccl, init_methodenv://) rank dist.get_rank() world_size dist.get_world_size() local_rank int(os.environ[LOCAL_RANK]) torch.cuda.set_device(local_rank) transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) dataset datasets.MNIST(./data, trainTrue, downloadTrue, transformtransform) sampler DistributedSampler(dataset, num_replicasworld_size, rankrank, shuffleTrue) loader DataLoader(dataset, batch_size64, samplersampler, num_workers4, pin_memoryTrue) model nn.Sequential( nn.Flatten(), nn.Linear(784, 128), nn.ReLU(), nn.Linear(128, 10) ).cuda() model DistributedDataParallel(model, device_ids[local_rank]) criterion nn.CrossEntropyLoss() optimizer torch.optim.SGD(model.parameters(), lr0.01, momentum0.9) for epoch in range(5): sampler.set_epoch(epoch) total_loss 0.0 for images, labels in loader: images, labels images.cuda(local_rank), labels.cuda(local_rank) out model(images) loss criterion(out, labels) optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() dist.barrier() if rank 0: print(fepoch {epoch} loss {total_loss / len(loader):.4f}) dist.destroy_process_group() if __name__ __main__: train()啟動(dòng)命令就一行torchrun --nproc_per_node4 train.py如果你想讓這份代碼跑兩個(gè)節(jié)點(diǎn)假設(shè)節(jié)點(diǎn) 0 的 IP 是 192.168.1.10就在節(jié)點(diǎn) 0 上執(zhí)行torchrun --nnodes2 --nproc_per_node4 --node_rank0 --master_addr192.168.1.10 --master_port29500 train.py節(jié)點(diǎn) 1 上執(zhí)行同樣命令只把--node_rank改成 1。注意所有節(jié)點(diǎn)的代碼、數(shù)據(jù)集路徑和 Python 環(huán)境最好保持一致否則分布式的報(bào)錯(cuò)會(huì)讓人崩潰。3. 實(shí)戰(zhàn)提速的幾個(gè)關(guān)鍵配置3.1 混合精度配合 DDP顯存和時(shí)間一起省如果 DDP 是分布式訓(xùn)練的第一個(gè)加速器那么混合精度AMP就是第二個(gè)。AMP 的核心思路是讓模型的大部分計(jì)算用 float16 進(jìn)行同時(shí)保留一部分操作比如損失計(jì)算、梯度更新用 float32 保證數(shù)值穩(wěn)定性再配合梯度縮放GradScaler防止 float16 下浮點(diǎn)下溢。它帶來(lái)的好處是顯存占用幾乎減半、速度常有 20%-50% 的提升。與 DDP 配套使用時(shí)邏輯上并不復(fù)雜只需把訓(xùn)練循環(huán)稍微改造from torch.cuda.amp import autocast, GradScaler scaler GradScaler() for images, labels in loader: images, labels images.cuda(local_rank), labels.cuda(local_rank) optimizer.zero_grad() with autocast(): out model(images) loss criterion(out, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()注意AMP 的 autocast 生效范圍要盡量覆蓋模型的前向計(jì)算不要只包一小部分。另外所有進(jìn)模型的張量都已經(jīng)是 CUDA float32 的話autocast 會(huì)自動(dòng)選擇合適精度不需要你手動(dòng)轉(zhuǎn)成 half。DDP 和 AMP 有一個(gè)配合點(diǎn)值得留意DDP 的梯度同步發(fā)生在 backward 階段也就是scaler.scale(loss).backward()這一步?;旌暇认碌奶荻缺旧砭褪?float16 的NCCL 傳輸時(shí)會(huì)按照 float16 進(jìn)行通信通信量直接減半。這也是為什么 AMPDDP 在帶寬受限的多機(jī)場(chǎng)景下加速效果比單機(jī)更明顯。3.2 學(xué)習(xí)率、全局 batch size 和梯度累積怎么配合多卡并行時(shí)全局 batch size 會(huì)成倍變大如果你還沿用原來(lái)的學(xué)習(xí)率訓(xùn)練大概率會(huì)不穩(wěn)定甚至直接發(fā)散。業(yè)界比較常用的經(jīng)驗(yàn)法則是linear scaling rulebatch size 變成原來(lái)的 k 倍時(shí)學(xué)習(xí)率也可以近似乘以 k但為了穩(wěn)妥更常見的做法是乘以 sqrt(k)或者給優(yōu)化器加一個(gè) warmup 階段讓學(xué)習(xí)率從一個(gè)小值線性爬升到目標(biāo)值。我個(gè)人的實(shí)操習(xí)慣是先保持學(xué)習(xí)率不變用一個(gè)小數(shù)據(jù)集跑幾步看看 loss 是否正常下降如果正常再嘗試按 sqrt(k) 放大學(xué)習(xí)率觀察幾個(gè) epoch 的曲線如果曲線比原來(lái)抖得更厲害就降低到原始學(xué)習(xí)率或增加 warmup 步數(shù)。不要盲目信global batch 變大就 lr 乘 k這條規(guī)則模型結(jié)構(gòu)、數(shù)據(jù)分布都會(huì)影響最終結(jié)果。說(shuō)到梯度累積很多人會(huì)把累積當(dāng)成減小全局 batch 的替代方案。梯度累積確實(shí)可以在不增加顯存的情況下模擬更大的 batch size但它在 DDP 下的實(shí)現(xiàn)有一個(gè)隱藏坑如果你累積了 4 個(gè) batch 再 backward 一次梯度會(huì)被放大 4 倍。正確做法是每次累積時(shí)手動(dòng)對(duì) loss 除以累積步數(shù)或者等 DDP 的梯度同步完成后再 average。我在代碼里一般這樣處理accum_steps 4 optimizer.zero_grad() for i, (images, labels) in enumerate(loader): with autocast(): out model(images) loss criterion(out, labels) / accum_steps scaler.scale(loss).backward() if (i 1) % accum_steps 0: scaler.step(optimizer) scaler.update() optimizer.zero_grad()這個(gè)寫法的邏輯是讓每次 loss 先除以累積步數(shù)backward 時(shí) DDP 鏡像出來(lái)的梯度就是單步梯度的平均值最后幾次累積得到的梯度相當(dāng)于減小了 batch 的梯度噪聲不會(huì)出現(xiàn) loss 數(shù)值被無(wú)意義放大的問(wèn)題。3.3 多機(jī)多卡的網(wǎng)絡(luò)配置與 NCCL 優(yōu)化多機(jī) DDP 和單機(jī)最大的不同在于進(jìn)程間的通信從本機(jī) GPU 的 NVLink 或 PCIe 變成了跨機(jī)器的以太網(wǎng)或者 InfiniBand。NCCL 是 PyTorch 默認(rèn)的 GPU 通信后端它對(duì)跨機(jī)通信的實(shí)現(xiàn)直接決定了多機(jī)的效率。想要多機(jī)跑得順暢以下三個(gè)點(diǎn)值得優(yōu)先檢查。第一個(gè)是確認(rèn)所有節(jié)點(diǎn)的MASTER_ADDR和MASTER_PORT設(shè)置正確。MASTER_ADDR必須填 rank 0 那個(gè)節(jié)點(diǎn)所有網(wǎng)卡都能訪問(wèn)到的 IP不要填回環(huán)地址 127.0.0.1。端口盡量選擇一個(gè)不太可能被占用的高位端口比如 29500 或 29501并在防火墻規(guī)則里放行 TCP 和 UDP 對(duì)應(yīng)端口。第二個(gè)是 NCCL 的調(diào)試開關(guān)。如果出現(xiàn)連接不上、初始化失敗等問(wèn)題我建議在啟動(dòng)命令前加上NCCL_DEBUGINFO讓 NCCL 把每一步通信日志打印出來(lái)。日志會(huì)告訴你進(jìn)程在嘗試連接哪個(gè) IP 的哪個(gè)端口哪里失敗一目了然。生產(chǎn)環(huán)境排查完后可以關(guān)掉因?yàn)?DEBUG 日志對(duì)性能有少量影響。第三個(gè)是針對(duì)不同網(wǎng)絡(luò)環(huán)境的 NCCL 開關(guān)。在 IB 網(wǎng)不可用時(shí)偶爾會(huì)出現(xiàn)某張卡連接不上或者連接超時(shí)的問(wèn)題這時(shí)可以試試在啟動(dòng)命令前加NCCL_P2P_DISABLE1強(qiáng)制走共享內(nèi)存或 TCP 通道雖然會(huì)降低一些通信效率但至少能讓程序跑起來(lái)。如果是多機(jī)場(chǎng)景還可以設(shè)置NCCL_SOCKET_IFNAME指定使用哪塊網(wǎng)卡比如NCCL_SOCKET_IFNAMEeth0。NCCL_DEBUGINFO NCCL_SOCKET_IFNAMEeth0 torchrun --nnodes2 --nproc_per_node4 --node_rank0 --master_addr192.168.1.10 --master_port29500 train.py3.4 隨機(jī)種子和訓(xùn)練結(jié)果的可復(fù)現(xiàn)性DDP 多進(jìn)程并行時(shí)隨機(jī)種子處理不好會(huì)帶來(lái)兩個(gè)問(wèn)題一是每個(gè)進(jìn)程的數(shù)據(jù)順序不同導(dǎo)致最終模型有差異二是調(diào)試 bug 時(shí)每次結(jié)果都不一樣很難判斷問(wèn)題到底出在哪。PyTorch 官方推薦的做法是在每個(gè)進(jìn)程內(nèi)設(shè)置一個(gè)基礎(chǔ)種子加 rank 偏移量的種子import random import numpy as np def setup_seed(seed_value, rank): seed seed_value rank random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed)每個(gè)進(jìn)程拿到的隨機(jī)序列既不同整體又可控既保證了打亂數(shù)據(jù)的多樣性又讓整個(gè)訓(xùn)練過(guò)程可以復(fù)現(xiàn)。需要注意DistributedSampler內(nèi)部已經(jīng)自帶了一套基于 epoch 和 seed 的確定性邏輯所以它不需要額外做說(shuō)明但你要確保它的shuffleTrue時(shí)每個(gè) epoch 都調(diào)用set_epoch否則隨機(jī)性不強(qiáng)。模型初始權(quán)重也需要同步。DDP 的構(gòu)造函數(shù)雖然會(huì)默認(rèn)做一次參數(shù)的 broadcast把所有進(jìn)程的模型初始參數(shù)拉齊但如果你是先從 checkpoint 加載權(quán)重再做 DDP 包裝就一定要保證每個(gè)進(jìn)程加載的 checkpoint 路徑一致、加載后的參數(shù)一致否則 DDP 會(huì)在訓(xùn)練過(guò)程中檢測(cè)到參數(shù)不一致并報(bào)錯(cuò)。4. 常見問(wèn)題與排查技巧實(shí)錄4.1 init_process_group 失敗NCCL 初始化報(bào)錯(cuò)這是我被問(wèn)得最多的一類問(wèn)題。最常見的原因有三個(gè)NCCL 版本和 CUDA 版本不匹配、網(wǎng)絡(luò)端口不通、PYTHON 環(huán)境不一致。處理順序我建議先看報(bào)錯(cuò)日志如果是連接超時(shí)優(yōu)先檢查多機(jī)場(chǎng)景的防火墻和MASTER_ADDR如果日志里出現(xiàn) CUDA driver version is insufficient 或 NCCL version mismatch優(yōu)先升級(jí)或?qū)R PyTorch、CUDA 和 nccl 的版本。有一個(gè)小技巧在正式訓(xùn)練腳本之前寫一個(gè)只有 init_process_group、打印 rank 和 world_size 的最小腳本先把通信鏈路驗(yàn)證通。我?guī)缀趺看尾鹊椒植际较嚓P(guān)的坑都會(huì)先用這種方式把環(huán)境問(wèn)題隔離掉再去看業(yè)務(wù)代碼問(wèn)題。這樣可以節(jié)省大量排查時(shí)間。4.2 梯度不同步、loss 忽大忽小如果你發(fā)現(xiàn)訓(xùn)練過(guò)程中 loss 在幾個(gè)進(jìn)程之間明顯不一致或者模型結(jié)果時(shí)好時(shí)壞第一步檢查是不是DistributedSampler忘了加。如果沒(méi)加各個(gè)進(jìn)程拿到的就是全部數(shù)據(jù)梯度方向混沌loss 波動(dòng)會(huì)非常大。第二個(gè)容易出問(wèn)題的點(diǎn)是模型里有部分參數(shù)沒(méi)有參與 loss 計(jì)算。DDP 默認(rèn)會(huì)檢查參數(shù)梯度的同步情況如果某些參數(shù)沒(méi)有梯度它會(huì)等待所有進(jìn)程都產(chǎn)生梯度再統(tǒng)一同步導(dǎo)致阻塞甚至死鎖。此時(shí)你需要在構(gòu)造 DDP 時(shí)設(shè)置find_unused_parametersTruemodel DistributedDataParallel(model, device_ids[local_rank], find_unused_parametersTrue)但這會(huì)讓性能稍微下降所以只在你確實(shí)存在未使用參數(shù)時(shí)開啟不要在一切正常時(shí)盲目加。這是我在一個(gè)帶輔助損失頭的模型上踩過(guò)的坑找了好幾天才發(fā)現(xiàn)是 unused parameter 的問(wèn)題。4.3 多卡后模型效果反而變差多卡訓(xùn)練后 loss 數(shù)值比單卡高、收斂變慢或者最終精度低于單卡這大概率不是 DDP 本身的問(wèn)題而是全局 batch size 變大后學(xué)習(xí)率沒(méi)有同步調(diào)整。我曾經(jīng)把一個(gè) batch64 的模型改成 4 卡 DDP沒(méi)調(diào)學(xué)習(xí)率結(jié)果訓(xùn)練 3 個(gè) epoch 后 loss 依然在初值附近晃悠。后來(lái)把全局 batch 從 256 降低到 128并加上 warmup收斂就正常了。另外也要關(guān)注數(shù)據(jù)集的 BatchNormBN層在 DDP 下的行為。DDP 默認(rèn)每個(gè)進(jìn)程獨(dú)立計(jì)算 BN 的均值和方差因?yàn)槊總€(gè)進(jìn)程只看到自己的數(shù)據(jù)子集。如果你的 batch size 較小BN 統(tǒng)計(jì)量會(huì)很不穩(wěn)定這時(shí)可以考慮用同步 BN 模塊torch.nn.SyncBatchNorm.convert_sync_batchnorm(model)讓 BN 的統(tǒng)計(jì)量跨進(jìn)程同步。注意這個(gè)方法應(yīng)該在 DDP 包裝之前調(diào)用否則無(wú)法正確替換。4.4 顯存不均衡和反復(fù) OOMDDP 多卡訓(xùn)練時(shí)顯存通常比較均衡但如果某張卡 OOM 的次數(shù)特別頻繁而另外幾張卡顯存還很富余問(wèn)題往往出在數(shù)據(jù)不均衡或者模型初始化不均衡上。先確認(rèn)你的DataLoader使用的是DistributedSampler而不是普通 sampler再看 Pin Memory 和 num_workers 是否設(shè)置合理。還有一種常見情況是某個(gè)進(jìn)程里加載了額外的數(shù)據(jù)或臨時(shí)變量比如 rank 0 負(fù)責(zé)日志打印時(shí)把最后一個(gè) batch 的輸出圖像存到了本地這部分顯存占用量沒(méi)有及時(shí)釋放導(dǎo)致該進(jìn)程率先 OOM。我的經(jīng)驗(yàn)是所有和訓(xùn)練無(wú)關(guān)的保存操作盡量都放在with torch.no_grad()或 CPU 端完成避免額外占用顯存。如果實(shí)在壓縮不下來(lái)可以先做梯度檢查點(diǎn)gradient checkpointing降低顯存也可以用torch.cuda.empty_cache()在每輪 epoch 后釋放顯存碎片但記住它是治標(biāo)不治本的。下面把最常見的幾個(gè)問(wèn)題和排查點(diǎn)整理成一個(gè)速查表方便大家直接對(duì)照?,F(xiàn)象可能原因排查/解決NCCL 初始化失敗/超時(shí)網(wǎng)絡(luò)不通、防火墻、MASTER_ADDR 錯(cuò)誤先跑最小 init 腳本驗(yàn)證通信放行端口檢查 IPloss 在兩個(gè)進(jìn)程間不一致缺少 DistributedSampler給 DataLoader 掛 DistributedSampler 并 set_epochloss 發(fā)散或收斂慢全局 batch size 變大、學(xué)習(xí)率未調(diào)整降低每卡 batch size 或調(diào)學(xué)習(xí)率加 warmup訓(xùn)練卡死無(wú)響應(yīng)find_unused_parameters 未設(shè)置檢查是否有參數(shù)未參與 loss設(shè)置該選項(xiàng)模型 BN 統(tǒng)計(jì)量抖動(dòng)每卡 batch 太小用 SyncBatchNorm 替代普通 BN某一進(jìn)程 OOM該進(jìn)程做了額外顯存操作保存/日志操作移到 CPU 或 no_grad 下執(zhí)行多機(jī)連接不穩(wěn)定多網(wǎng)卡 IP 模式不匹配設(shè)置 NCCL_SOCKET_IFNAME / NCCL_P2P_DISABLE我在實(shí)際處理這些問(wèn)題的過(guò)程中最大的體會(huì)是DDP 本身并不復(fù)雜復(fù)雜的是訓(xùn)練流程里的各種隱式假設(shè)。單卡腳本能跑通不代表它在多進(jìn)程場(chǎng)景下語(yǔ)義依然正確。每次排查都要問(wèn)自己三個(gè)問(wèn)題數(shù)據(jù)是不是分開了、模型參數(shù)是不是同步了、梯度是不是平均了。只要這三個(gè)點(diǎn)穩(wěn)了剩下的性能優(yōu)化都是錦上添花。最后分享一個(gè)我一直在用的習(xí)慣任何 DDP 改造都先從兩卡、小數(shù)據(jù)集、5 個(gè) epoch 開始跑通再逐步放大到全量數(shù)據(jù)和多機(jī)環(huán)境。別一上來(lái)就追求最大規(guī)模分布式訓(xùn)練的錯(cuò)誤往往在小規(guī)模下更容易暴露。這樣積累幾輪之后你會(huì)發(fā)現(xiàn) PyTorch DDP 其實(shí)是一個(gè)非常成熟且省心的工具真正難的從來(lái)不是它而是你對(duì)整個(gè)訓(xùn)練管線有沒(méi)有足夠的掌控力。