问题:单卡训练太慢,我决定上多卡
上个月我在 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,所有梯度汇总到 GPU0 | Ring 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_SIZE、LOCAL_RANK、RANK。你也可以自己解析参数,但用环境变量更省事。
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峰值显存(每卡) |
|---|---|---|---|
| 单卡 A100 | 76.2 | 1.0x | 12.4GB |
| 4卡 DP | 43.5 | 1.75x | GPU0: 38.2 / 其他: 10.1GB |
| 4卡 DDP | 21.8 | 3.49x | 12.8GB |
| 4卡 DDP + AMP | 16.4 | 4.65x | 9.1GB |
| 8卡 DDP | 12.6 | 6.04x | 12.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 的源码和文档不难懂,真正难的是隐藏在“看起来正常的代码”里的那几个坑。把上面的坑记住,你的多卡训练至少能省两天调错时间。