torch.nn.DataParallel(单进程多卡)适合单机多卡、快速上手,但效率一般(有 GIL 和主卡瓶颈)。
import torch
import torch.nn as nn
model = MyModel()
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
if torch.cuda.device_count() > 1:
print(f"Using {torch.cuda.device_count()} GPUs")
model = nn.DataParallel(model)
model = model.to(device)outputs = model(inputs.to(device))
loss = criterion(outputs, labels.to(device))
loss.backward()
optimizer.step()batch 维度拆分model.module.state_dict()DistributedDataParallel(DDP,多进程多卡)生产环境、训练大模型首选,效率更高,支持单机多卡 & 多机多卡。
torchrun)torchrun \
--nproc_per_node=4 \
train.pyimport os
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data.distributed import DistributedSampler
def main(rank, world_size):
# 初始化进程组
dist.init_process_group(
backend="nccl",
init_method="env://",
rank=rank,
world_size=world_size
)
torch.cuda.set_device(rank)
model = MyModel().to(rank)
model = DDP(model, device_ids=[rank])
dataset = MyDataset()
sampler = DistributedSampler(dataset)
loader = torch.utils.data.DataLoader(
dataset, batch_size=32, sampler=sampler
)
for x, y in loader:
x, y = x.to(rank), y.to(rank)
loss = model(x, y).loss
loss.backward()
optimizer.step()
if __name__ == "__main__":
world_size = torch.cuda.device_count()
mp.spawn(main, args=(world_size,), nprocs=world_size)DistributedSampler 避免数据重复DataParallel 快很多| 对比项 | DataParallel | DistributedDataParallel |
|---|---|---|
| 多进程 | ❌ | ✅ |
| 效率 | 较低 | 高 |
| 多机支持 | ❌ | ✅ |
| 推荐程度 | 教学/调试 | 生产/训练 |
DataParallel 里直接 model.module 忘了写DistributedSamplerbatch_size 写成“总 batch”而不是“单卡 batch”torch.cuda.amp)如果你愿意,我可以: