PyTorch分布式训练:DDP原理与实战

一、DDP核心原理

分布式数据并行(Distributed Data Parallel)通过多进程实现模型并行训练:

  1. 数据分片
    全局批次$B$被划分为$N$个子批次$B_i$,满足$B = \bigcup_{i=1}^{N} B_i$,其中$N$为GPU数量
  2. 模型复制
    每个GPU持有完整的模型副本$M$
  3. 梯度同步
    反向传播后,各GPU计算局部梯度$\nabla W_i$,通过Ring-AllReduce算法同步全局梯度: $$ \nabla W_{\text{global}} = \frac{1}{N} \sum_{i=1}^{N} \nabla W_i $$
  4. 参数更新
    所有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

四、性能优化技巧
  1. 梯度压缩
    使用torch.distributed.algorithms.ddp_comm_hooks.default_hooks.allreduce_hook减少通信量
  2. 重叠计算与通信
    设置broadcast_buffers=False避免同步BN层
  3. 混合精度训练
    结合torch.cuda.amp减少显存占用:
    scaler = torch.cuda.amp.GradScaler()
    with torch.cuda.amp.autocast():
        pred = model(x)
    scaler.scale(loss).backward()
    scaler.step(optimizer)
    

五、常见问题排查
  1. 死锁检测
    使用NCCL_DEBUG=INFO环境变量输出通信日志
  2. 显存溢出
    检查find_unused_parameters=True是否误启用
  3. 负载不均衡
    验证DistributedSampler分片均匀性

注:完整代码需包含if __name__ == "__main__"保护,使用spawn启动进程。实际训练时建议全局批次大小$B_{\text{global}} = N \times B_{\text{local}}$,其中$B_{\text{local}}$为单卡批次大小。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐