使用 DistributedDataParallel 时:
model = nn.parallel.DistributedDataParallel(
model,
device_ids=[local_rank]
)torch.distributed.all_reduce(最常用)用于 对所有进程中的张量做归约(sum / mean / max 等)
import torch.distributed as dist
tensor = torch.tensor([1.0]).cuda()
dist.all_reduce(tensor, op=dist.ReduceOp.SUM)✅ 所有进程在 all_reduce 后 tensor 值相同
常见用法:
dist.all_reduce(tensor)
tensor /= dist.get_world_size()torch.distributed.broadcast把一个进程的数据 广播给所有进程
if dist.get_rank() == 0:
tensor = torch.tensor([1.0]).cuda()
else:
tensor = torch.zeros(1).cuda()
dist.broadcast(tensor, src=0)✅ 常用于:
torch.distributed.all_gather收集所有进程的张量
tensor = torch.tensor([dist.get_rank()]).cuda()
gathered = [torch.zeros_like(tensor) for _ in range(world_size)]
dist.all_gather(gathered, tensor)✅ 常用于:
reduce / gather(不常用)| 操作 | 说明 |
|---|---|
reduce | 归约到某一个进程 |
gather | 收集到某一个进程 |
普通 BatchNorm 只在 单卡 batch 上统计,可能不稳定。
model = torch.nn.SyncBatchNorm.convert_sync_batchnorm(model)
model = DDP(model, device_ids=[local_rank])DistributedSamplersampler = DistributedSampler(dataset)
loader = DataLoader(dataset, sampler=sampler)✅ 保证:
shuffle=Trueimport torch.distributed as dist
dist.init_process_group(
backend="nccl", # GPU 用 nccl
init_method="env://"
)或:
torchrun --nproc_per_node=8 train.py# 错误
dist.all_reduce(model.weight)if rank == 0:
do_something()✅ 会导致 hang
broadcast 或 all_reducetorch.manual_seed(42) # 只在一个进程设置✅ 正确做法:
seed = 42
torch.manual_seed(seed)| 场景 | 推荐方式 |
|---|---|
| 模型参数 | DDP 自动 |
| 梯度 | DDP 自动 |
| loss / metric | all_reduce |
| 超参数 | broadcast |
| 验证结果 | all_gather |
| BN 统计 | SyncBatchNorm |
| 数据采样 | DistributedSampler |
PyTorch 分布式中:
- 模型 & 梯度 → DDP 自动同步
- 自定义数据 →
all_reduce / broadcast / all_gather- BN → SyncBatchNorm
- 数据 → DistributedSampler
如果你有 具体代码场景(比如 loss 同步、参数平均、多机训练),我可以帮你直接改代码。