【PyTorch】torch.distributed 实战指南:从单机多卡到多机分布式训练
1. 为什么需要分布式训练?
想象一下你正在训练一个超大的视觉模型,数据集有1000万张图片,单卡训练要跑一个月。这时候分布式训练就像召唤了一群帮手,把任务拆解后并行处理,训练时间可能缩短到几天甚至几小时。
PyTorch的torch.distributed模块就是这样的神器,它能让你:
- 把数据拆分到多个GPU上并行计算(数据并行)
- 把超大模型拆解到不同设备上(模型并行)
- 在多个物理机器上组建计算集群
我去年在训练一个3D点云检测模型时,单卡训练每个epoch要6小时,改用4卡DDP后直接降到1.5小时,效果立竿见影。
2. 单机多卡实战
2.1 环境初始化
先来看最简单的单机多卡场景。假设你有台8卡服务器:
import torch.distributed as dist
def setup(rank, world_size):
# 关键配置(torchrun会自动设置这些环境变量)
os.environ['MASTER_ADDR'] = 'localhost'
os.environ['MASTER_PORT'] = '12355'
# 初始化进程组
dist.init_process_group(
backend="nccl", # NVIDIA显卡必选
rank=rank,
world_size=world_size
)
torch.cuda.set_device(rank) # 每个进程绑定不同GPU
这里有个坑我踩过:MASTER_PORT如果被占用会报错,建议选20000以上的端口。
2.2 数据并行改造
数据并行的核心是让每个GPU处理不同的数据批次:
from torch.utils.data.distributed import DistributedSampler
def get_dataloader(dataset, batch_size):
sampler = DistributedSampler(
dataset,
num_replicas=world_size,
rank=rank,
shuffle=True
)
return DataLoader(
dataset,
batch_size=batch_size,
sampler=sampler,
pin_memory=True # 加速数据加载
)
注意这个**sampler.set_epoch(epoch)**操作必须加在训练循环里,否则每个epoch的数据顺序会完全一样:
for epoch in range(epochs):
dataloader.sampler.set_epoch(epoch) # 保证shuffle有效
for batch in dataloader:
...
2.3 模型包装
用DDP包装模型是核心操作:
model = YourAwesomeModel().to(rank)
model = DDP(model, device_ids=[rank])
这里有个性能优化点:如果模型有些层不参与计算(比如某些条件分支),可以设置find_unused_parameters=True,但会增加约10%的开销。
3. 多机分布式实战
3.1 跨机器通信配置
多机训练需要指定主节点IP。比如有两台机器,IP分别是192.168.1.101和192.168.1.102:
# 在第一台机器上(rank=0)
os.environ['MASTER_ADDR'] = '192.168.1.101'
os.environ['MASTER_PORT'] = '29500'
# 在第二台机器上(rank=1)
os.environ['MASTER_ADDR'] = '192.168.1.101' # 指向主节点
os.environ['MASTER_PORT'] = '29500' # 必须相同
3.2 启动方式对比
推荐使用torchrun替代旧的启动方式:
# 单机8卡
torchrun --nproc_per_node=8 train.py
# 多机示例(每台机器8卡)
torchrun \
--nnodes=2 \
--nproc_per_node=8 \
--rdzv_id=123456 \ # 唯一任务ID
--rdzv_backend=c10d \
--rdzv_endpoint=192.168.1.101:29500 \
train.py
我在AWS上实测过16台机器的集群训练,用Elastic Launch(带自动容错)特别稳,即使有节点宕机也能继续训练。
4. 通信后端选型指南
4.1 NCCL vs Gloo对比
| 特性 | NCCL | Gloo |
|---|---|---|
| 最佳场景 | 多GPU训练 | CPU集群或混合设备 |
| 跨机器 | 支持 | 支持 |
| RDMA支持 | 是 | 否 |
| 集合通信优化 | 极致优化 | 基础实现 |
| 调试难度 | 较难 | 简单 |
4.2 常见问题排查
问题1:NCCL报错"unhandled system error"
- 解决方案:添加环境变量
export NCCL_DEBUG=INFO export NCCL_SOCKET_IFNAME=eth0 # 指定网卡
问题2:多机训练连接超时
- 检查防火墙设置
- 测试节点间网络:
nc -zv <ip> <port>
5. 高级优化技巧
5.1 梯度累积实现
当显存不足时,可以用梯度累积模拟更大batch:
model = DDP(...)
for i, batch in enumerate(dataloader):
with model.no_sync(): # 前N-1次不同步梯度
if i % 3 != 0:
loss = compute_loss(batch)
loss.backward()
continue
# 第N次同步梯度
loss = compute_loss(batch)
loss.backward()
optimizer.step()
optimizer.zero_grad()
5.2 混合精度训练
结合AMP使用效果更好:
from torch.cuda.amp import autocast
scaler = GradScaler()
with autocast():
output = model(input)
loss = criterion(output, target)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
6. 实战中的血泪教训
-
数据一致性:曾因为忘记设置
sampler.set_epoch(),导致所有机器用相同数据顺序,验证集准确率死活上不去。 -
OOM问题:在BERT训练中,发现当batch_size<8时,NCCL通信开销反而会使训练变慢。
-
死锁陷阱:某次在多机训练中,因为某个进程提前退出,导致其他进程一直卡在
barrier()。
建议大家在正式训练前,先用小数据跑通整个流程。分布式调试就像在迷宫找出口,有日志才能不迷路:
# 每个进程打印日志
if dist.get_rank() == 0:
print(f"[Rank 0] Epoch {epoch} completed")
dist.barrier() # 同步点
if dist.get_rank() == 1:
print(f"[Rank 1] Passed barrier")
更多推荐
所有评论(0)