PyTorch 分布式训练:DDP 的原理与实战
·
PyTorch分布式训练:DDP原理与实战
一、DDP核心原理
分布式数据并行(Distributed Data Parallel)通过多进程实现模型并行训练:
- 数据分片
全局批次$B$被划分为$N$个子批次$B_i$,满足$B = \bigcup_{i=1}^{N} B_i$,其中$N$为GPU数量 - 模型复制
每个GPU持有完整的模型副本$M$ - 梯度同步
反向传播后,各GPU计算局部梯度$\nabla W_i$,通过Ring-AllReduce算法同步全局梯度: $$ \nabla W_{\text{global}} = \frac{1}{N} \sum_{i=1}^{N} \nabla W_i $$ - 参数更新
所有GPU使用$\nabla W_{\text{global}}$同步更新模型参数
二、DDP实战步骤
1. 环境初始化
import torch.distributed as dist
def setup(rank, world_size):
dist.init_process_group(
backend="nccl", # NVIDIA GPU推荐使用NCCL
init_method="env://",
rank=rank,
world_size=world_size
)
torch.cuda.set_device(rank)
2. 模型封装
from torch.nn.parallel import DistributedDataParallel as DDP
def build_model(rank):
model = ResNet50().to(rank)
ddp_model = DDP(model, device_ids=[rank])
return ddp_model
3. 数据加载器配置
from torch.utils.data.distributed import DistributedSampler
def prepare_dataloader(dataset, rank, world_size):
sampler = DistributedSampler(
dataset,
num_replicas=world_size,
rank=rank,
shuffle=True
)
return DataLoader(dataset, batch_size=64, sampler=sampler)
4. 训练循环模板
def train(rank, world_size):
setup(rank, world_size)
model = build_model(rank)
loader = prepare_dataloader(dataset, rank, world_size)
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)
for epoch in range(10):
sampler.set_epoch(epoch) # 确保每个epoch数据不同
for x, y in loader:
x, y = x.to(rank), y.to(rank)
pred = model(x)
loss = F.cross_entropy(pred, y)
optimizer.zero_grad()
loss.backward() # DDP自动同步梯度
optimizer.step()
dist.destroy_process_group()
三、启动训练脚本
使用torchrun启动多进程训练(示例启动4个GPU):
torchrun --nproc_per_node=4 --nnodes=1 train.py
四、性能优化技巧
- 梯度压缩
使用torch.distributed.algorithms.ddp_comm_hooks.default_hooks.allreduce_hook减少通信量 - 重叠计算与通信
设置broadcast_buffers=False避免同步BN层 - 混合精度训练
结合torch.cuda.amp减少显存占用:scaler = torch.cuda.amp.GradScaler() with torch.cuda.amp.autocast(): pred = model(x) scaler.scale(loss).backward() scaler.step(optimizer)
五、常见问题排查
- 死锁检测
使用NCCL_DEBUG=INFO环境变量输出通信日志 - 显存溢出
检查find_unused_parameters=True是否误启用 - 负载不均衡
验证DistributedSampler分片均匀性
注:完整代码需包含
if __name__ == "__main__"保护,使用spawn启动进程。实际训练时建议全局批次大小$B_{\text{global}} = N \times B_{\text{local}}$,其中$B_{\text{local}}$为单卡批次大小。
更多推荐
所有评论(0)