PyTorch 分布式训练主要依赖:
nccl(GPU,推荐)gloo(CPU / 部分 GPU)mpi(需要 MPI 环境)1 台机器 × 多张 GPUtorch.distributed.launch 或 torchrunnccl多台机器 × 多张 GPUMASTER_ADDRMASTER_PORTrank 不同import torch
import torch.distributed as dist
dist.init_process_group(
backend="nccl",
init_method="env://"
)local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)model = MyModel().cuda()
model = torch.nn.parallel.DDistributedDataParallel(
model,
device_ids=[local_rank]
)from torch.utils.data.distributed import DistributedSampler
sampler = DistributedSampler(dataset)
dataloader = DataLoader(
dataset,
batch_size=32,
sampler=sampler
)shuffle=Truetorchrun(PyTorch ≥ 1.9)torchrun \
--nproc_per_node=4 \
train.py# 节点 0
torchrun \
--nnodes=2 \
--node_rank=0 \
--nproc_per_node=4 \
--master_addr=192.168.1.1 \
--master_port=29500 \
train.py
# 节点 1
torchrun \
--nnodes=2 \
--node_rank=1 \
--nproc_per_node=4 \
--master_addr=192.168.1.1 \
--master_port=29500 \
train.pyimport os
import torch
import torch.distributed as dist
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
def main():
local_rank = int(os.environ["LOCAL_RANK"])
dist.init_process_group("nccl")
torch.cuda.set_device(local_rank)
dataset = MyDataset()
sampler = DistributedSampler(dataset)
loader = DataLoader(dataset, batch_size=32, sampler=sampler)
model = MyModel().cuda()
model = torch.nn.parallel.DistributedDataParallel(
model, device_ids=[local_rank]
)
for x, y in loader:
x, y = x.cuda(), y.cuda()
loss = model(x, y)
loss.backward()
optimizer.step()
if __name__ == "__main__":
main()✅ 必须使用 DistributedSampler
✅ batch_size 是 单卡 batch
✅ 保存模型只在 rank == 0
✅ 日志只在 rank == 0 打印
torch.distributed.fsdp如果你愿意,我可以:
你现在的训练场景是?