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对比

特性NCCLGloo
最佳场景多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. 实战中的血泪教训

  1. 数据一致性:曾因为忘记设置sampler.set_epoch(),导致所有机器用相同数据顺序,验证集准确率死活上不去。

  2. OOM问题:在BERT训练中,发现当batch_size<8时,NCCL通信开销反而会使训练变慢。

  3. 死锁陷阱:某次在多机训练中,因为某个进程提前退出,导致其他进程一直卡在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")
Logo

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

更多推荐