DDP分布式训练实战:从单卡到四卡的提升与坑
发布日期: 2026/08/11 阅读总量: 1

问题:单卡训练太慢,我决定上多卡

上个月我在 CIFAR-10 上训一个 ResNet-50,单卡 A100 跑一个 epoch 要 3.8 分钟,20 个 epoch 就是 76 分钟。调参几次,一天时间没了。换成 ImageNet 子集(5 万张图)后,单卡一个 epoch 要 40 分钟,20 个 epoch 要 13 小时。这没法干活。我决定用 4 卡并行。

一开始图省事,直接用了 torch.nn.DataParallel。结果跑起来后,GPU0 显存占用 38GB,其他三卡只有 10GB,GPU 利用率忽高忽低。4 卡训练时间只比单卡快了 1.75 倍,完全不像“4 卡并行”。后来换成 DistributedDataParallel(DDP),同样 4 卡,时间直接降到 21.8 分钟,加速比 3.49 倍。差距在哪?

两种多卡方案对比:DP vs DDP

先上结论:新项目一律用 DDP,DP 只适合“临时验证单机代码”。

对比项DataParallel (DP)DistributedDataParallel (DDP)
进程模型单进程多线程多进程,每进程一张卡
通信方式GPU0 作为 reducer,所有梯度汇总到 GPU0Ring AllReduce,梯度切块后点对点通信
负载均衡GPU0 同时负责前向和梯度归约,负载高每张卡只处理自己的 batch,梯度同步是异步的
可扩展性只支持单机,无法多机支持单机/多机多卡
速度(4卡实测)加速 1.75x加速 3.49x
推荐程度弃用官方推荐

为什么 DP 慢?它把全部梯度 copy 到 GPU0 再算平均,GPU0 的 PCIe 和显存带宽成了瓶颈。而且每个 step 都要同步,线程切换也有开销。DDP 则把梯度分成 N 个区间,每个进程只在自己负责的区间上做 AllReduce,然后拿回完整梯度更新参数。通信是异步的,能和反向传播重叠。

DDP 的核心机制

DDP 在 forward 时同步 buffer(如 BN 的 running_mean),在 backward 时通过 ReduceScatter 和 AllGather 完成梯度同步。具体到 PyTorch 2.1.0,默认后端是 nccl,支持 Ring AllReduce。通信只发生在相邻的两个进程之间,不会像 DP 那样是“一星辐射”结构。

这里有个容易误解的点:DDP 不是平均模型参数,而是平均梯度。每张卡独立计算自己 batch 的梯度,然后 AllReduce 求平均,再用平均后的梯度去更新每张卡上的模型副本。所以所有进程的模型参数最终严格相同。

Ring AllReduce 的通信量

假设模型大小为 M,卡数为 K。DP 的通信复杂度是 O(M),GPU0 要收 4 份完整梯度。DDP 的 Ring AllReduce 分两步:ReduceScatter 和 AllGather。每张卡只发送自己的梯度切片,总通信量是 2*(K-1)/K * M。K 越大,额外通信越接近 2M,但它是分布式传输,不是单点瓶颈。

完整代码实现

环境:Python 3.10、PyTorch 2.1.0、CUDA 11.8、4×A100 80G。

1. 启动脚本

最推荐的启动方式是 PyTorch 自带的 torchrun

# 单机4卡
torchrun --nproc_per_node=4 --master_port=29500 train.py

# 多机(2台机器,每台4卡)
# 节点0:
torchrun --nnodes=2 --nproc_per_node=4 \
  --master_addr=192.168.1.10 --master_port=29500 train.py
# 节点1:
torchrun --nnodes=2 --nproc_per_node=4 \
  --master_addr=192.168.1.10 --master_port=29500 train.py

2. 训练脚本核心骨架

# train.py
import os
import torch
import torch.nn as nn
import torchvision
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
import torch.distributed as dist

def cleanup():
    dist.destroy_process_group()

def main(local_rank):
    # 关键第一步:初始化进程组
    dist.init_process_group(
        backend='nccl',
        init_method='env://',
    )
    torch.cuda.set_device(local_rank)

    # 构造模型
    model = torchvision.models.resnet50(num_classes=10)
    model = model.to(local_rank)
    model = nn.parallel.DistributedDataParallel(
        model,
        device_ids=[local_rank],
        output_device=local_rank,
    )

    # 数据集
    transform = torchvision.transforms.Compose([
        torchvision.transforms.RandomCrop(32, padding=4),
        torchvision.transforms.RandomHorizontalFlip(),
        torchvision.transforms.ToTensor(),
        torchvision.transforms.Normalize((0.4914, 0.4822, 0.4465),
                                         (0.2023, 0.1994, 0.2010)),
    ])
    train_dataset = torchvision.datasets.CIFAR10(
        root='./data', train=True, transform=transform, download=True)
    train_sampler = DistributedSampler(
        train_dataset,
        num_replicas=dist.get_world_size(),
        rank=local_rank,
        shuffle=True,
    )
    train_loader = DataLoader(
        train_dataset,
        batch_size=64,
        sampler=train_sampler,
        num_workers=4,
        pin_memory=True,
    )

    criterion = nn.CrossEntropyLoss()
    # 学习率按世界卡数线性缩放:原LR=0.1,4卡时0.4
    optimizer = torch.optim.SGD(
        model.parameters(),
        lr=0.1 * dist.get_world_size(),
        momentum=0.9,
        weight_decay=1e-4,
    )

    for epoch in range(20):
        # 关键第二步:每个epoch需要设置sampler的epoch,保证数据shuffle顺序不同
        train_sampler.set_epoch(epoch)
        model.train()
        for images, labels in train_loader:
            images = images.to(local_rank)
            labels = labels.to(local_rank)
            outputs = model(images)
            loss = criterion(outputs, labels)
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()
        # 只在rank0上打印
        if local_rank == 0:
            print(f'epoch {epoch} loss {loss.item():.4f}')

    cleanup()

if __name__ == '__main__':
    local_rank = int(os.environ['LOCAL_RANK'])
    main(local_rank)

torchrun 会自动设置环境变量 WORLD_SIZELOCAL_RANKRANK。你也可以自己解析参数,但用环境变量更省事。

3. 保存与加载 checkpoint

保存模型最常见的错误是直接 torch.save(model.state_dict())。DDP 的 model 是包装类,它的 state_dict 键会带 module. 前缀。加载到原始模型时会报 key 不匹配。我的做法:

def save_checkpoint(model, optimizer, epoch, ckpt_path):
    # 保存原始模型状态字典,或者去掉module前缀
    if isinstance(model, nn.parallel.DistributedDataParallel):
        model_state = model.module.state_dict()
    else:
        model_state = model.state_dict()
    if dist.get_rank() == 0:
        torch.save({
            'epoch': epoch,
            'model_state_dict': model_state,
            'optimizer_state_dict': optimizer.state_dict(),
        }, ckpt_path)

def load_checkpoint(model, optimizer, ckpt_path, device):
    if os.path.exists(ckpt_path):
        # 每个进程都要加载,但设备要映射到当前rank的卡
        checkpoint = torch.load(ckpt_path, map_location=f'cuda:{device}')
        model.module.load_state_dict(checkpoint['model_state_dict'])
        optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
        return checkpoint['epoch']
    return 0

关键是 map_location 一定要用当前设备,否则会加载到 GPU0 导致显存溢出。

4. 多机多卡配置

多机时,除了启动参数,还要保证所有节点能互相访问。下面是我用的一个 Docker 内多机启动示例:

# 每台机器上执行(假设同一个docker网络)
docker run --gpus all --shm-size=64g --network=host \
  -v /data:/data \
  -e NCCL_DEBUG=INFO \
  -e NCCL_SOCKET_IFNAME=eth0 \
  -e NCCL_IB_DISABLE=1 \
  your_image \
  torchrun --nnodes=2 --nproc_per_node=4 \
    --master_addr=$MASTER_ADDR --master_port=29500 \
    train.py

5. 梯度累积实现“大 batch”

有时候 GPU 显存不够,每个进程只能放小 batch,但你又想保持和单卡相同的全局 batch size。可以用梯度累积:

from contextlib import nullcontext

accum_steps = 4  # 每卡batch=64,全局等效batch=64*4*4=1024
for i, (images, labels) in enumerate(train_loader):
    images = images.to(local_rank)
    labels = labels.to(local_rank)
    loss = criterion(model(images), labels) / accum_steps
    context = model.no_sync() if (i + 1) % accum_steps != 0 else nullcontext()
    with context:
        loss.backward()
    if (i + 1) % accum_steps == 0:
        optimizer.step()
        optimizer.zero_grad()

注意:DDP 默认每个 backward 都会触发梯度 AllReduce。如果你在累积步内每个小 step 都 backward,就会白白做多次通信。用 model.no_sync() 跳过前 accum_steps-1 次的同步,只在最后一个 step 同步。

6. 混合精度训练

混合精度不只能省显存,还能显著提速。在 DDP 环境下使用 AMP 很直接:

scaler = torch.cuda.amp.GradScaler()

for epoch in range(20):
    train_sampler.set_epoch(epoch)
    for images, labels in train_loader:
        images = images.to(local_rank)
        labels = labels.to(local_rank)
        optimizer.zero_grad()
        with torch.cuda.amp.autocast(dtype=torch.float16):
            outputs = model(images)
            loss = criterion(outputs, labels)
        scaler.scale(loss).backward()
        scaler.step(optimizer)
        scaler.update()

实测 4 卡 DDP 加上 AMP 后,21.8 分钟降到 16.4 分钟,显存从 12.8GB 降到 9.1GB。

实验数据:不同方案到底差多少

实验配置:ResNet-50,CIFAR-10 50000 张图,训练 20 epochs,SGD lr=0.1(单卡),多卡时按线性缩放为 0.1×world_size,momentum=0.9,weight_decay=1e-4。数据输入用 torchvision 默认 transform。测得的训练耗时如下(包括数据加载、前向、反向、更新,但不包括验证):

配置总耗时(分钟)相对加速比GPU峰值显存(每卡)
单卡 A10076.21.0x12.4GB
4卡 DP43.51.75xGPU0: 38.2 / 其他: 10.1GB
4卡 DDP21.83.49x12.8GB
4卡 DDP + AMP16.44.65x9.1GB
8卡 DDP12.66.04x12.9GB

8 卡为什么不是 8 倍增速?因为 batch 总数不变,每个进程负责的数据变少,每次 AllReduce 的通信开销占比变大。加上 A100 的算力对于这个 batch size 已经充裕,计算和通信的重叠效率有限。

我还用 PyTorch profiler 记录了 DDP 的资源利用率(4卡):

with torch.profiler.profile(
    activities=[torch.profiler.ProfilerActivity.CUDA],
    schedule=torch.profiler.schedule(wait=1, warmup=1, active=3)
) as prof:
    for i, (images, labels) in enumerate(train_loader):
        images = images.to(local_rank)
        labels = labels.to(local_rank)
        loss = criterion(model(images), labels)
        loss.backward()
        optimizer.step()
        prof.step()
print(prof.key_averages().table(
    sort_by="cuda_time_total", row_limit=10))

输出显示:nccl:AllReduce 平均耗时 4.8ms/step,约占整个 step 的 11.2%。说明 DDP 的通信开销并没有想象中那么大,原因是 PyTorch 2.1 的 DDP 会把梯度切块,让通信和反向计算重叠。

避坑指南

以下坑我全部踩过,按频率排序。

坑1:DistributedSampler 不 set_epoch,导致每个 epoch 数据顺序一样

我第一次写 DDP 时,只在 DataLoader 里放了 sampler,忘了在 epoch 循环里调用 train_sampler.set_epoch(epoch)。结果 20 个 epoch 的 shuffle 状态一直是初始化时的状态,模型每次看到的数据排列完全一样,训练损失波动大,最终精度低了 0.8%。官方文档明确要求:如果你在 sampler 里用 shuffle=True,必须在每个 epoch 调用 set_epoch。

坑2:保存模型时把 DDP 包装层也存了

torch.save(model.state_dict()) 保存的 key 是 module.xxx。之后在单卡环境加载,键名不匹配,报 Missing key(s) in state_dict: "xxx"。建议封装 save/load 函数,前面代码那样用 model.module.state_dict()

坑3:torch.load 没指定 map_location,导致所有进程都把模型加载到 GPU0

DDP 每个进程是独立的,但如果你写 torch.load(ckpt),它默认加载到当前进程的 cuda:0。多卡机器上,4 个进程都去读 GPU0,GPU0 直接被塞爆。同时你会发现其他卡显存只有几十 MB。必须改成 map_location=f'cuda:{local_rank}',或者干脆先加载到 CPU。

坑4:多机训练时 master 地址和端口问题

使用 --master_addr 时,两个节点的地址必须互通。有一次我把 master_addr 写成了内网 IP,但两个节点在同一个宿主机上的 Docker 中,网络模式是 bridge,导致节点1连不上节点0。换成 --network=host 解决。如果 NCCL 通信一直超时,设置环境变量 NCCL_DEBUG=INFO 查看日志,多半是 IB 或 socket 接口选择问题,可以试试 NCCL_IB_DISABLE=1。另外 --shm-size 别忘了,多进程数据加载会疯狂吃共享内存。

坑5:batch size 变了但学习率没调

DDP 中每卡 batch=64,4 卡全局 batch 就是 256。如果你还用单卡的 lr=0.1,模型更新步数变成了原来的 1/4,收敛速度会明显变慢。工程经验是把 lr 按全局 batch 比例线性放大:lr = lr × world_size。上面实验里,4卡 lr=0.4 的最终精度和单卡 lr=0.1 基本一致(62.3% vs 62.1%,CIFAR-10 无增强)。

坑6:BN 参数不同步

DDP 默认每个进程独立更新 BN 的 running_mean/running_var,因为每个进程只看到自己那一份数据,统计量会有偏差。如果模型对 BN 敏感,尤其是目标检测这类任务,可使用 torch.nn.SyncBatchNorm.convert_sync_batchnorm(model) 把 BN 转为同步 BN,但会增加通信量。我的建议:每卡 batch 大于 32 用普通 BN;每卡只有 8 或 16 时用同步 BN。

坑7:DataLoader 的 num_workers 设置不当

DDP 中每个进程都会启动一组 DataLoader worker。如果你的 num_workers=8 且 4 卡,总共 32 个 worker 同时读磁盘,IO 可能成为瓶颈。我在机械盘上试过,num_workers=8 时训练速度反而比 num_workers=2 慢 16%。建议 num_workers 按 CPU 核心数适当设置,并观察磁盘 IO 占用。另外把数据放到 SSD 上能解决 80% 的数据加载问题。

最后说一句

DDP 的源码和文档不难懂,真正难的是隐藏在“看起来正常的代码”里的那几个坑。把上面的坑记住,你的多卡训练至少能省两天调错时间。