# 推荐 Python >= 3.8
conda create -n torch_dist python=3.10
conda activate torch_dist
# CUDA 11.8 示例
pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118验证:
import torch
print(torch.cuda.device_count())
print(torch.distributed.is_available())# 查看 NCCL 版本
python -c "import torch; print(torch.cuda.nccl.version())"| 模式 | 说明 | 适用场景 |
|---|---|---|
torch.distributed + NCCL | 最常用 | 多 GPU / 多机 |
torch.nn.DataParallel | 简单但已不推荐 | 单机 |
torch.distributed.fsdp | 超大模型 | 显存不足 |
torch.distributed.elastic | 容错训练 | 集群 |
torch.distributed + NCCLimport os
import torch
import torch.distributed as dist
def setup():
dist.init_process_group(
backend="nccl",
init_method="env://"
)需要设置环境变量:
export MASTER_ADDR="127.0.0.1"
export MASTER_PORT=29500
export WORLD_SIZE=4
export RANK=0torchrun(强烈推荐 ✅)不需要手动设置 RANK / WORLD_SIZE
torchrun \
--nproc_per_node=4 \
--nnodes=1 \
--node_rank=0 \
--master_addr="127.0.0.1" \
--master_port=29500 \
train.pyimport torch
import torch.distributed as dist
import torch.multiprocessing as mp
import torch.nn as nn
from torch.nn.parallel import DistributedDataParallel as DDPtrain.pyimport os
import torch
import torch.distributed as dist
import torch.nn as nn
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
def setup():
dist.init_process_group(backend="nccl")
def cleanup():
dist.destroy_process_group()
def main():
setup()
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
model = nn.Linear(128, 10).cuda(local_rank)
model = DDP(model, device_ids=[local_rank])
dataset = torch.randn(1000, 128)
sampler = DistributedSampler(dataset)
loader = DataLoader(dataset, batch_size=32, sampler=sampler)
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)
for epoch in range(3):
sampler.set_epoch(epoch)
for x in loader:
x = x.cuda(local_rank)
loss = model(x).sum()
loss.backward()
optimizer.step()
optimizer.zero_grad()
cleanup()
if __name__ == "__main__":
main()torchrun --nproc_per_node=4 train.py假设:
torchrun \
--nproc_per_node=4 \
--nnodes=2 \
--node_rank=0 \
--master_addr="192.168.1.10" \
--master_port=29500 \
train.pytorchrun \
--nproc_per_node=4 \
--nnodes=2 \
--node_rank=1 \
--master_addr="192.168.1.10" \
--master_port=29500 \
train.pyinit_process_group 没执行完WORLD_SIZE 不一致export NCCL_DEBUG=INFO
export NCCL_SOCKET_IFNAME=eth0✅ 必须用 DistributedSampler
shuffle=True如果你愿意,我可以:
直接告诉我你的 GPU 数量 + 单机/多机 + PyTorch 版本 即可。