nn.DataParallel常报错或不加速,因其将模型参数默认置于cuda:0,导致所有GPU需从该卡拉取权重,形成IO瓶颈;且不支持非均匀batch、torch.compile或FSDP。

PyTorch里nn.DataParallel为什么常报错或不加速?
多数人一上来就用nn.DataParallel,结果发现显存没省、速度反而更慢,甚至报RuntimeError: Expected all tensors to be on the same device。根本原因是它只在前向/反向时做输入分发和梯度收集,但模型参数仍默认在cuda:0,所有GPU都要从cuda:0拉权重,成了IO瓶颈;且不支持非均匀batch(比如某些GPU显存小),也不兼容torch.compile或FSDP。
实操建议:
- 确认你真需要单机多GPU——如果模型能塞进一张卡,优先用
torch.compile+ 单卡mixed precision,通常比DataParallel快2–3倍 - 若必须多卡,改用
nn.DistributedDataParallel(DDP),哪怕单机也要启动多进程,这是当前PyTorch官方唯一推荐的多GPU训练方式 - 别在Jupyter或交互式环境里试DDP——它依赖
torch.distributed.init_process_group,必须走独立Python脚本+torchrun
单机DDP训练必须写的三段核心代码
DDP不是“加一行就跑”,它要求数据、模型、优化器三者严格对齐到各自进程。漏掉任一环节,就会出现梯度不同步、loss震荡或NCCL timeout。
关键步骤:
立即学习“Python免费学习笔记(深入)”;
- 启动时用
torchrun --nproc_per_node=4 train.py(不是python train.py),它会自动设置MASTER_ADDR、RANK等环境变量 - 在
train.py开头初始化:import torch.distributed as dist dist.init_process_group(backend="nccl") # 必须在创建模型前
- 模型包装必须放在每个进程自己的GPU上:
model = MyModel().to(rank) # rank来自dist.get_rank() model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[rank])
- 数据加载器要用
DistributedSampler:train_sampler = torch.utils.data.DistributedSampler(dataset, shuffle=True) loader = DataLoader(dataset, sampler=train_sampler, batch_size=32)
torch.nn.parallel.DistributedDataParallel的常见陷阱
DDP看着简单,但几个细节不处理,轻则性能打折,重则训练失效。
容易踩的坑:
-
device_ids和output_device别乱设——单卡每进程只用一个GPU,写device_ids=[rank]即可;设output_device反而触发跨卡拷贝,拖慢速度 - 验证阶段记得
model.eval()后调用torch.no_grad(),否则DDP会尝试同步梯度(而eval时不该有) - 保存模型只能由
rank == 0的进程执行:if dist.get_rank() == 0: torch.save(...),否则所有进程都写同一个文件,内容错乱 - 自定义
loss函数里如果用了.mean(),要确保是全局mean——因为每个GPU只算自己那部分batch,得先all_reduce再除总样本数
单机多GPU下显存不够怎么办?
即使开了DDP,模型本身太大,单卡显存还是爆——这时候不能靠“增加GPU数量”硬扛,得换策略。
可选路径:
- 启用
torch.compile(PyTorch 2.0+):model = torch.compile(model),它能自动融合op、减少中间tensor,显存常降20%–40% - 开
gradient_checkpointing(以时间换空间):model.gradient_checkpointing_enable(),适合Transformer类模型,但训练变慢约15% - 用
FSDP(Fully Sharded Data Parallel)替代DDP:它把模型参数、梯度、优化器状态都分片到各GPU,显存占用接近线性下降,但配置稍复杂,需替换DistributedDataParallel为torch.distributed.fsdp.FSDP - 避免
pin_memory=True+num_workers>0组合在小数据集上——这反而引发内存泄漏,显存虚高
DDP本身不解决模型过大问题,它只是让多卡协同工作;真正压显存的是模型结构和batch size。定位时先用torch.cuda.memory_summary()看每卡实际分配量,再决定该剪枝、该分片,还是该降batch。


















