DistributedDataParallel(DDP),其核心机制是 反向传播时的梯度 AllReduce。下面从原理、流程、代码示例和常见注意点来说明。input -> model -> lossloss.backward()在 loss.backward() 时:
grad = all_reduce(grad) / world_sizeoptimizer.step()import torch
import torch.distributed as dist
dist.init_process_group(
backend="nccl", # GPU 用 nccl,CPU 用 gloo
init_method="env://"
)通常通过 torchrun 启动:
torchrun --nproc_per_node=4 train.pylocal_rank = int(os.environ["LOCAL_RANK"])
model = MyModel().to(local_rank)
model = torch.nn.parallel.DistributedDataParallel(
model,
device_ids=[local_rank]
)optimizer.zero_grad()
output = model(data)
loss = criterion(output, target)
loss.backward() # ✅ 梯度自动同步
optimizer.step()grad = all_reduce(grad) / world_size# 不更新参数,只累积梯度
optimizer.zero_grad()
for i in range(accum_steps):
loss = model(data)
loss = loss / accum_steps
loss.backward() # 梯度仍然同步
optimizer.step()loss.backward() 依然会同步梯度no_sync() 上下文管理器with model.no_sync():
loss.backward()✅ 在 no_sync() 内:
for i, (data, target) in enumerate(dataloader):
with model.no_sync() if i % accum_step != 0 else nullcontext():
loss.backward()| 特性 | DDP | FSDP |
|---|---|---|
| 梯度同步 | AllReduce | AllGather + ReduceScatter |
| 显存占用 | 高 | 低 |
| 适合模型 | 中等 | 超大模型 |
PyTorch 分布式训练中,梯度同步由 DDP 在 loss.backward() 时自动完成,通过 AllReduce 将各卡梯度求平均,无需手动干预。如果你愿意,我可以: