1 / accum_stepsfor i, (x, y) in enumerate(loader):
loss = model(x, y) / accum_steps
loss.backward()
if (i+1) % accum_steps == 0:
optimizer.step()
optimizer.zero_grad()DDP(find_unused_parameters=False)all-reduce(如每隔 N 步同步一次)适合:
Fully Sharded Data Parallel
from torch.distributed.fsdp import FSDP
model = FSDP(model)# 伪代码
grad = quantize(grad) # int8 / fp16
all_reduce(grad)
grad = dequantize(grad)from torch.distributed.algorithms.ddp_comm_hooks import (
powerSGD_hook as powerSGD
)
model.register_comm_hook(state, powerSGD.powerSGD_hook)torch.distributed.all_reduce(
tensor, op=ReduceOp.SUM
)fp16 / bf16grad_scaler 控制limit_all_gathers=TrueFSDP(model, limit_all_gathers=True)export NCCL_IB_DISABLE=0
export NCCL_SOCKET_IFNAME=eth0推荐:
| 目标 | 推荐方案 |
|---|---|
| 快速见效 | 梯度累积 |
| 显存+通信 | FSDP |
| 超大模型 | FSDP + PowerSGD |
| 多机慢 | 减少节点 / IB |
| 通信频繁 | 降低同步频率 |
如果你能说明:
我可以给出更具体的配置和代码。