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