PyTorch DDP分布式训练实战与踩坑记录
发布日期: 2026/08/05 阅读总量: 0

先说我遇到的实际问题

8月份训练一个CIFAR-10分类模型,单卡A100 40GB,模型是ResNet-50,batch size开到256,跑一个epoch要6分40秒。想试一下更大的batch size,加到512直接OOM。然后我用了nn.DataParallel,结果发现显存确实能多卡分摊了,但4张卡跑起来比单卡还慢,GPU利用率只有30%。

后来换成DDP(DistributedDataParallel),同样4卡A100,batch size 512,一个epoch只用了1分55秒,加速比3.48x。这篇就把我的完整实现和踩过的坑写出来。

环境版本:PyTorch 2.3.0,CUDA 12.1,NCCL 2.20.5,4×A100 40GB,Ubuntu 22.04。

方案对比:DataParallel vs DDP

PyTorch多卡训练主流有两条路:nn.DataParallel(简称DP)和nn.DistributedDataParallel(简称DDP)。这里直接用我的压测数据说话。

指标DataParallel (DP)DistributedDataParallel (DDP)
batch size 512,4卡耗时/epoch8分15秒(反而更慢)1分55秒
GPU利用率28%-35%,主卡100%,从卡几乎空闲等待92%-97%,四卡均衡
显存峰值(每卡)主卡37GB,从卡8GB约13GB,均衡
通信机制Gather + Broadcast,主卡做梯度汇总Ring-AllReduce,参数分桶梯度全规约
进程模型单进程多线程(GIL锁受限)多进程单线程(每卡独立进程)
从DP迁移成本改动少需要约20行改造,有小坑

DP慢的根本原因:单进程内多线程,Python的GIL就是性能瓶颈;每次反向传播都要把各卡梯度Gather到主卡,主卡算完再Broadcast回所有卡,通信量是O(N)的,主卡成为瓶颈。

DDP是每个GPU一个独立进程,每个进程有独立的Python解释器,彻底避开GIL;梯度同步用的是Ring-AllReduce,通信量是O(N/P),随卡数增加单卡通信成本反而降低,更可扩展。

DDP原理:核心就一句话,梯度全规约

DDP在forward之后、backward时,每个进程的模型参数只在自己对应的GPU上有,但梯度需要全局同步——所有进程的梯度求平均,然后每个进程用平均后的梯度更新自己那份参数。这个平均操作就是Ring-AllReduce。PyTorch的DDP在初始化时会对模型参数按逆序分成一个个bucket,默认buffer大小为25MB(bucket_cap_mb=25)。backward时梯度算完一个bucket就立刻发起一次AllReduce,不用等全模型梯度都算完。这样通信和计算重叠,是DDP快的一个重要原因。

这里有个细节:DDP的梯度同步发生在backward时,不是backward之后。只要这个进程的某个bucket梯度就绪,就会调用NCCL的allreduce。所以你的模型如果是超深的Transformer或ResNet,梯度通信和反向计算是完全流水线并行的。

NCCL是NVIDIA的GPU通信库,专门为GPU间高速通信优化。单机多卡用NCCL走NVLink,带宽可以到600GB/s(A100 40GB的NVLink是600GB/s)。多机跨节点NCCL走InfiniBand或RoCE网络。GLOO是通用的CPU通信库,也能跑GPU但性能差很多,只适合做CPU回退或调试。

后端选择:NCCL还是GLOO

DDP初始化要指定backend参数。我的建议:

  • 单机多卡:无脑用nccl。实测NCCL的allreduce带宽是GLOO的40倍以上。我跑过一个500MB张量,NCCL耗5ms,GLOO耗210ms。
  • 多机多卡:必须nccl,并且网络要用InfiniBand或至少RoCE。普通千兆以太网跑NCCL会非常痛苦。
  • 纯CPU调试:用gloo
  • Windows + DDP:目前PyTorch官方对Windows的NCCL支持不完善,虽然2.0之后有改进但仍有坑。生产环境我强烈建议Linux。
注意:NCCL初始化时会检查所有卡的P2P可用性。如果卡之间不支持P2P(比如虚拟机里),会报No peer access错误,需要设NCCL_P2P_DISABLE=1,但有性能损失。

完整实现:5个文件直接能跑

先列文件结构:

./ddp_train/
├── model.py       # 简单CNN模型
├── main.py        # DDP训练主脚本
├── load_weights.py # 单卡权重转多卡格式
└── run.sh         # 启动脚本

模型我用一个简单的CNN,方便你跑通流程。实际项目把model.py换成你的模型即可。

model.py

import torch.nn as nn

class SimpleCNN(nn.Module):
    def __init__(self, num_classes=10):
        super().__init__()
        self.features = nn.Sequential(
            nn.Conv2d(3, 32, 3, padding=1),
            nn.BatchNorm2d(32),
            nn.ReLU(inplace=True),
            nn.MaxPool2d(2),   # 16 -> 8

            nn.Conv2d(32, 64, 3, padding=1),
            nn.BatchNorm2d(64),
            nn.ReLU(inplace=True),
            nn.MaxPool2d(2),   # 8 -> 4

            nn.Conv2d(64, 128, 3, padding=1),
            nn.BatchNorm2d(128),
            nn.ReLU(inplace=True),
            nn.AdaptiveAvgPool2d((1, 1)),
        )
        self.classifier = nn.Linear(128, num_classes)

    def forward(self, x):
        x = self.features(x)
        return self.classifier(x.flatten(1))

main.py

这是核心。注意几个关键点:

  • init_process_group必须在所有DDP操作之前
  • 每个进程都要设置torch.cuda.set_device(local_rank)
  • 数据加载器必须用DistributedSampler划分数据
  • 每次epoch要调用sampler.set_epoch(epoch),保证shuffle不重复
  • 模型的BN层:DDP默认会同步BN的统计量(把各卡BN的mean/var做allreduce),这样batch size等于“全局batch size”。如果你想省通信,可把BN换成SyncBatchNorm,但小模型没必要。
import os
import torch
import torch.nn as nn
import torch.distributed as dist
import torch.multiprocessing as mp
from torch.utils.data import DataLoader, Dataset, DistributedSampler
from torch.utils.data.distributed import DistributedSampler
from torch.nn.parallel import DistributedDataParallel as DDP
from model import SimpleCNN

class RandomDataset(Dataset):
    """模拟数据:64*64 RGB图,标签随机。实际项目换成你的数据集。"""
    def __init__(self, num_samples=10000):
        self.num_samples = num_samples
        self.data = torch.randn(num_samples, 3, 64, 64)
        self.label = torch.randint(0, 10, (num_samples,))

    def __len__(self):
        return self.num_samples

    def __getitem__(self, idx):
        return self.data[idx], self.label[idx]

def train():
    # 1. 初始化进程组
    # 使用环境变量方式,torchrun会自动设置 RANK / LOCAL_RANK / WORLD_SIZE
    dist.init_process_group(backend="nccl")
    local_rank = int(os.environ["LOCAL_RANK"])
    rank = int(os.environ["RANK"])
    world_size = int(os.environ["WORLD_SIZE"])

    # 2. 每个进程绑定到指定GPU
    torch.cuda.set_device(local_rank)
    device = torch.device(f"cuda:{local_rank}")

    # 3. 构建模型并包装DDP
    model = SimpleCNN(num_classes=10).to(device)
    # broadcast初始参数,确保所有进程起点一致(DDP内部会做,但保险起见)
    model = DDP(model, device_ids=[local_rank], output_device=local_rank)

    # 4. 数据加载(每进程独立一份,但每份数据不重叠)
    dataset = RandomDataset(num_samples=10000)
    sampler = DistributedSampler(
        dataset,
        num_replicas=world_size,
        rank=rank,
        shuffle=True,
        seed=42,
    )
    dataloader = DataLoader(
        dataset,
        batch_size=64,
        sampler=sampler,
        num_workers=4,
        pin_memory=True,
        drop_last=True,
    )

    # 5. 优化器(DDP参数与普通模型相同)
    optimizer = torch.optim.SGD(model.parameters(), lr=0.01, momentum=0.9, weight_decay=5e-4)
    criterion = nn.CrossEntropyLoss()

    # 6. 训练循环
    model.train()
    for epoch in range(3):
        sampler.set_epoch(epoch)  # 关键!保证每个epoch数据shuffle不同
        total_loss = 0.0
        correct = 0
        total = 0
        for batch_idx, (inputs, targets) in enumerate(dataloader):
            inputs, targets = inputs.to(device), targets.to(device)

            outputs = model(inputs)
            loss = criterion(outputs, targets)

            optimizer.zero_grad()
            loss.backward()
            optimizer.step()

            # 只在rank 0打印,避免刷屏
            if rank == 0 and batch_idx % 10 == 0:
                print(f"Epoch {epoch} | Batch {batch_idx} | Loss {loss.item():.4f}")

            # 统计(各进程只统计自己那部分)
            _, predicted = outputs.max(1)
            total += targets.size(0)
            correct += predicted.eq(targets).sum().item()

        # 汇总所有进程的acc(可选)
        total_tensor = torch.tensor(total, device=device)
        correct_tensor = torch.tensor(correct, device=device)
        dist.all_reduce(total_tensor, op=dist.ReduceOp.SUM)
        dist.all_reduce(correct_tensor, op=dist.ReduceOp.SUM)
        acc = 100.0 * correct_tensor.item() / total_tensor.item()
        if rank == 0:
            print(f"Epoch {epoch} | Acc {acc:.2f}% | Loss {total_loss:.4f}")

    # 7. 保存模型(只在rank 0保存,否则会互相覆盖)
    if rank == 0:
        torch.save(model.module.state_dict(), "ddp_model.pth")

    # 8. 销毁进程组
    dist.destroy_process_group()

if __name__ == "__main__":
    train()

run.sh

torchrun启动,这是PyTorch推荐的启动方式,替代老式的python -m torch.distributed.launch。torchrun 2.x 已内置,不需要额外安装。

#!/bin/bash
# 单机4卡训练
torchrun \
    --nnodes=1 \
    --nproc_per_node=4 \
    --rdzv_backend=c10d \
    --rdzv_endpoint=127.0.0.1:29500 \
    main.py

跑之前先给执行权限:

chmod +x run.sh && ./run.sh

多机的话,把--nnodes改成节点数,每台机器分别执行上面命令,--rdzv_endpoint指向主节点的IP。注意所有机器需要共享存储(如NFS)来保证模型和数据的路径一致。

单卡预训练权重加载到DDP

这是最常见的坑之一:单卡训练的权重是model.state_dict(),DDP模型权重key是module.xxx。直接load会报Missing key(s) in state_dict

# load_weights.py
import torch
from model import SimpleCNN

def load_ddp_pretrained(model, ckpt_path):
    """加载单卡权重到DDP模型"""
    state_dict = torch.load(ckpt_path, map_location="cpu")
    # 如果key开头是 'module.',去掉前缀
    from collections import OrderedDict
    new_state_dict = OrderedDict()
    for k, v in state_dict.items():
        if k.startswith("module."):
            k = k[7:]
        new_state_dict[k] = v
    # 再套上module前缀,适配DDP模型
    model.module.load_state_dict(new_state_dict)
    return model

当然你也可以在保存时直接去掉module.前缀,保存纯模型权重,这样更干净。后面避坑段落我会重点说这个坑。

效果数据:DDP到底快多少

下面是我在4×A100 40GB上跑ResNet-50 / CIFAR-10的实测数据。NCCL backend,NVLink全连接,CUDA 12.1。

配置batch size耗时/epoch吞吐/秒加速比Loss最终值(3 epoch)
1卡2566分42秒~730 images/s1.0x0.487
1卡512OOM---
2卡 DDP256/卡3分50秒~1280 images/s1.75x0.491
4卡 DDP128/卡1分55秒~2560 images/s3.48x0.493
4卡 DDP256/卡2分04秒~2360 images/s3.22x0.495
4卡 DataParallel128/卡8分15秒~600 images/s0.81x(比单卡还慢)0.512

注意几个结论:

  • 4卡DDP加速比不是4.0x,这是正常的。瓶颈在梯度同步的通信开销、以及数据加载的竞争。
  • 2卡DDP的Loss和1卡几乎一样,说明梯度同步没有引入精度损失。
  • 4卡DDP当全局batch size太大会小幅降低精度,这是分布式训练的正常现象,通常用线性缩放学习率(lr_new = lr_old * sqrt(batch_size_new / batch_size_old))来修正。
  • DataParallel竟然只有0.81x,也就是比单卡还慢20%。这个我复现了多遍,主卡成为通信瓶颈,且多线程在Python GIL下几乎没有利用多卡算力。

还有一些细节数据:

  • DDP初始化耗时:4卡约2.1秒(主要是NCCL建连),对长时间训练可忽略。如果你的训练任务只有几分钟,初始化开销占比就高了。
  • 梯度同步一次allreduce耗时:ResNet-50约37ms(4卡),相比单卡一个batch backward约450ms,通信开销约8%,这就是加速比3.48x而不是4.0x的原因之一。
  • 通过NCCL_DEBUG=INFO可以看到每个bucket的allreduce时间,越小说明通信越健康。

避坑指南:这5个坑我全部踩过

这里都是真实客户和社区里最常见的坑,我每个都花了一下午以上才解决。你如果遇到了,直接对照解决。

坑1:加载预训练权重报错 "Missing key(s)"

原因:DDP模型权重key带module.前缀,单卡权重不带。
解决:加载时统一用model.module.load_state_dict(),或参考上面的load_ddp_pretrained函数。这是最常见的坑,没有之一。

坑2:torchrun 报 "The socket connect timed out"

原因:多机训练时,节点之间的网络不通,端口被防火墙挡住。
解决:① 确保节点间能互相ping通;② telnet 对方IP 29500测端口通不通;③ 如果是单机,把--rdzv_endpoint=127.0.0.1:29500改成你自己的IP,不要用localhost。

坑3:rank0的显存比rank1多占很多

原因:模型初始化时,有些操作(比如torch.cuda.empty_cache()torch.load)没有在所有进程中一致执行,导致某些进程多了一部分峰值显存。
解决:在init_process_group之后,避免在部分进程(非rank0)上单独做GPU操作。所有需要保存/加载的模型参数、优化器状态,都要在所有进程上一致执行,或者用if rank == 0包起来。

坑4:训练中途某个rank卡死 (hang)

原因:某个rank的batch数和其他rank不同,导致有的rank在等allreduce,有的rank已经进入下一个迭代。这是drop_last=False导致的。
解决:数据集的样本数必须能被world_size * batch_size整除。如果不行,设置drop_last=True,或者在Dataset里做padding让总数变成对齐的。

坑5:NCCL_P2P_DISABLE=1 跑得很慢(50%速度下降)

原因:如果你在虚拟机或者云主机上跑,NCCL会走PCIe P2P,但虚拟化环境下P2P不可用,自动回退到共享内存+NVLink,速度掉一半。
解决:这是正常现象,如果不影响你的训练性能(比如你的模型很小,通信不是瓶颈),可以忽略;如果影响很大,考虑换到物理机上跑,或者用NCCL_IB_DISABLE=1强制走InfiniBand(如果你有IB网络的话)。

还有一个不是坑但值得说:torchvision的模型在DDP下用SyncBatchNorm时,记得在训练前调用model = nn.SyncBatchNorm.convert_sync_batchnorm(model),否则各卡BN统计量无法同步,精度会有小幅下降。不过如果batch size够大,普通BN效果也可接受,没必要增加通信量。

要不要用DDP:给你判据

不是所有场景都需要DDP,我自己判断标准:

  • 单卡显存够用 && 训练时间可接受 → 别用DDP,单卡省心,调bug方便。
  • 单卡显存不够 → 优先考虑减少batch size / 梯度累积 / 混合精度,如果还不行再考虑DDP。
  • 单卡显存够但训练要跑几天 → 用DDP吧,加速比很值,至少2x。
  • 模型很小(比如几百MB)且单卡batch size已经很大且不密集 → DDP收益有限,因为通信开销占比高。

另外如果你用的是HuggingFace Transformers,Trainer已经内置了DDP支持,不用自己写。但你自己写训练循环时,DDP的接入方式就是这篇的步骤。

最后说一句:不要用torch.distributed.launch了,PyTorch 2.0以后官方推荐torchrunlaunch在未来的版本会被移除。