PyTorch分布式训练如何同步梯度

GPU
小华
2026-07-20

在 PyTorch 分布式训练中,梯度同步主要依赖 DistributedDataParallel(DDP),其核心机制是 反向传播时的梯度 AllReduce。下面从原理、流程、代码示例和常见注意点来说明。

一、梯度同步的核心机制

1️⃣ 基本思想

  • 每个 GPU 上有一个 模型副本
  • 每张卡计算 本地 batch 的梯度
  • 在反向传播时,自动对所有卡的梯度做 AllReduce(求和或平均)
  • 然后每个卡用同步后的梯度更新参数

不需要手动写梯度同步代码

二、DDP 中梯度同步流程

前向传播

input -> model -> loss

反向传播(关键点)

loss.backward()

loss.backward() 时:

  • DDP hook 自动触发
  • 对每个参数的梯度执行:
grad = all_reduce(grad) / world_size

参数更新

optimizer.step()

每个进程使用 相同的梯度,因此参数保持一致。

三、标准 DDP 使用示例

1️⃣ 初始化进程组

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.py

2️⃣ 包装模型

local_rank = int(os.environ["LOCAL_RANK"])
model = MyModel().to(local_rank)
model = torch.nn.parallel.DistributedDataParallel(
model,
device_ids=[local_rank]
)

3️⃣ 训练循环

optimizer.zero_grad()
output = model(data)
loss = criterion(output, target)
loss.backward()     # ✅ 梯度自动同步
optimizer.step()

四、梯度同步的关键细节

✅ AllReduce 的默认行为

  • 求和 + 平均
  • 等价于:
grad = all_reduce(grad) / world_size

✅ 同步粒度

  • 参数级别
  • 每个参数反向传播时立即同步(不是等所有参数算完)

五、梯度同步的控制方式

1️⃣ 是否同步梯度(梯度累积)

# 不更新参数,只累积梯度
optimizer.zero_grad()
for i in range(accum_steps):
loss = model(data)
loss = loss / accum_steps
loss.backward()   # 梯度仍然同步
optimizer.step()

⚠️ DDP 中 loss.backward() 依然会同步梯度

2️⃣ 手动控制同步(高级)

  • 使用 梯度累积 + 自定义 hook
  • no_sync() 上下文管理器
with model.no_sync():
loss.backward()

✅ 在 no_sync() 内:

  • 不会进行梯度 AllReduce
  • 常用于梯度累积
for i, (data, target) in enumerate(dataloader):
with model.no_sync() if i % accum_step != 0 else nullcontext():
loss.backward()

六、常见误区

❌ 误区 1:以为 optimizer 会同步参数

  • ❌ optimizer 只负责更新
  • ✅ 参数同步靠 DDP 初始化时 broadcast

❌ 误区 2:以为 loss 会同步

  • ❌ loss 是 本地计算的
  • ✅ 梯度才是同步的

❌ 误区 3:DDP + DataParallel 混用

  • ❌ 不推荐
  • ✅ 只用 DDP

七、DDP vs FSDP 对比

特性DDPFSDP
梯度同步AllReduceAllGather + ReduceScatter
显存占用
适合模型中等超大模型

八、一句话总结

PyTorch 分布式训练中,梯度同步由 DDP 在 loss.backward() 时自动完成,通过 AllReduce 将各卡梯度求平均,无需手动干预。

如果你愿意,我可以:

  • 画一张 DDP 梯度同步流程图
  • DDP 和 FSDP 的梯度同步差异
  • 结合 你当前的训练代码 帮你检查是否正确同步
亿速云提供售前/售后服务

售前业务咨询

售后技术保障

400-100-2938

7*24小时售后电话

官方微信小程序