推荐:多机多卡训练首选
同步内容
适用
用于自定义同步逻辑
常用 API:
torch.distributed.all_reduce()
torch.distributed.broadcast()
torch.distributed.barrier()适用
❌ 不支持多机
import torch.distributed as dist
dist.init_process_group(
backend="nccl",
init_method="tcp://MASTER_IP:PORT",
rank=RANK,
world_size=WORLD_SIZE
)参数说明:
| 参数 | 含义 |
|---|---|
| MASTER_IP | 主节点 IP |
| PORT | 通信端口 |
| RANK | 全局进程编号 |
| WORLD_SIZE | 总进程数 |
model = model.to(local_rank)
model = torch.nn.parallel.DistributedDataParallel(
model,
device_ids=[local_rank]
)✅ DDP 会自动:
必须使用 DistributedSampler
sampler = torch.utils.data.distributed.DistributedSampler(dataset)
loader = DataLoader(dataset, sampler=sampler)loss.backward()
optimizer.step()DDP 内部:
# 节点0
torchrun \
--nnodes=2 \
--node_rank=0 \
--nproc_per_node=8 \
--master_addr=192.168.1.1 \
--master_port=29500 \
train.py
# 节点1
torchrun \
--nnodes=2 \
--node_rank=1 \
--nproc_per_node=8 \
--master_addr=192.168.1.1 \
--master_port=29500 \
train.py适用于 HPC 集群
srun python train.py原因:
barrier()✅ 解决:
dist.barrier()原因:
DistributedSampler✅ 解决:
torch.manual_seed(42)export NCCL_SOCKET_IFNAME=eth0| 方式 | 同步内容 | 是否自动 | 推荐度 |
|---|---|---|---|
| DDP | 梯度 | ✅ | ⭐⭐⭐⭐⭐ |
| 手动 dist | 任意 | ❌ | ⭐⭐⭐ |
| DP | 无 | ❌ | ❌ |
多机 PyTorch 训练 = torchrun + DistributedDataParallel + DistributedSampler + NCCL
如果你愿意,我可以: