nn.DataParallel导致cuda:0显存暴涨的根本原因是其主从式单进程设计:loss计算、梯度汇聚、参数更新及优化器状态均集中于cuda:0,同时需缓存所有卡的前向输出与梯度缓冲区。

根本原因不是模型或数据本身,而是并行策略设计导致0号卡承担额外内存负担。
nn.DataParallel 为什么让 cuda:0 显存暴涨
它不是真并行,而是一个进程内“主从式”调度:所有子 batch 前向计算分散到各卡,但 loss 计算、梯度 gather、参数更新全在 cuda:0 完成。这意味着 cuda:0 必须额外缓存:
- 所有 GPU 的前向输出(比如 logits)拼接后的完整 tensor
- 反向传播时所有卡的梯度汇总缓冲区
- 优化器状态(如 Adam 的
m和v向量)只存一份,在 cuda:0 上维护
实测 ResNet-50 + batch_size=256(4卡)下,反向传播阶段 cuda:0 显存比其他卡高 47%,参数更新时达 61%。
DDP 加载模型时为什么所有进程都往 cuda:0 写显存
问题出在 torch.load 默认行为:不指定 map_location 时,它会把 checkpoint 中的 tensor 先加载到 cuda:0,再由各进程自行搬运——结果是 4 个进程同时往 cuda:0 申请临时空间,哪怕它们最终要挪到 cuda:1/2/3。
立即学习“Python免费学习笔记(深入)”;
典型错误写法:
checkpoint = torch.load("model.pth") # ← 这行就把所有权重先塞进 cuda:0
model.load_state_dict(checkpoint["state_dict"])
model = model.cuda(local_rank)
正确做法必须加 map_location:
- 推荐:
torch.load("model.pth", map_location=torch.device("cpu")),再.cuda(local_rank) - 若 CPU 内存紧张:用
map_location=f"cuda:{local_rank}",直接加载到目标卡,避免跨卡拷贝
batch 分配看似均匀,实际可能隐含偏差
DP 按顺序切分 batch,但没考虑样本计算复杂度差异。NLP 中长序列和短序列混在一个 batch 里,DP 仍均分 token 数,导致某张卡实际计算量翻倍。
DDP 虽无此问题,但若 DistributedSampler 使用不当也会失衡:
- 数据集长度不能被 GPU 数整除时,
indices[self.rank:self.total_size:self.num_replicas]会让部分卡多拿 1 个样本 - 没调用
sampler.set_epoch(epoch)→ shuffle 失效 → 某些卡反复拿到难样本 -
DataLoader的num_workers > 2在 NFS 或某些文件系统上易触发句柄泄漏,间接导致某进程卡住、显存滞留
最常被忽略的一点:显存“不均”未必真是 bug。比如你用 3 卡跑 batch_size=100,DP 会分 34/33/33,而 DDP 下若没配 DistributedSampler,三张卡可能都在处理全部 100 个样本——这时看到的“不均”其实是逻辑错误,不是显存分配问题。


















