✅ 首选:torch.distributed + NCCL + DDP
torch.distributed.init_process_group(backend="nccl")
model = DistributedDataParallel(model, device_ids=[local_rank])❌ 不推荐:
DataParallel(慢、GIL 限制)gloo(CPU 还可以,GPU 慢)DistributedDataParallel (DDP)DataParallel 快很多减少同步频率:
accum_steps = 4
for i, batch in enumerate(dataloader):
loss = model(batch) / accum_steps
loss.backward()
if (i + 1) % accum_steps == 0:
optimizer.step()
optimizer.zero_grad()✅ 适合:
from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
for batch in dataloader:
optimizer.zero_grad()
with autocast():
loss = model(batch)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()✅ 收益:
DataLoader(
dataset,
batch_size=batch_size,
num_workers=8,
pin_memory=True,
prefetch_factor=2,
persistent_workers=True
)⚠️ 注意:
num_workers 不是越大越好(一般 4~16)torch.cuda.synchronize()
print(loss.item()) # 会隐式同步✅ 改进:
loss.detach()torch.compile(PyTorch 2.x)model = torch.compile(model)✅ 适合:
⚠️ 注意:
torch.distributed.algorithms.ddp_comm_hooksfrom torch.distributed.algorithms.ddp_comm_hooks import default_hooks
model.register_comm_hook(
state=None,
hook=default_hooks.fp16_compress_hook
)DDP(model, bucket_cap_mb=25)✅ 小模型 ↓ bucket
| 方法 | 适合 |
|---|---|
| DDP | 常规模型 |
| FSDP | 超大模型 |
| DeepSpeed ZeRO | 超大规模 |
from torch.distributed.fsdp import FullyShardedDataParallel as FSDPfrom torch.utils.checkpoint import checkpoint
x = checkpoint(layer, x)torchrun --nproc_per_node=8 train.py而不是:
python -m torch.distributed.launch| 问题 | 排查 |
|---|---|
| GPU 利用率低 | dataloader / 同步 |
| 卡在 all_reduce | 网络 / NCCL |
| 显存爆炸 | batch size / 激活 |
| 多机慢 | IB / MTU / NCCL |
✅ 单机能跑满 GPU
DDP + AMP + 梯度累积 + torch.compile✅ 大模型
FSDP + ZeRO + Activation Checkpoint你可以直接贴: