推理场景下优先用dataparallel,因其开箱即用、无需进程管理;ddp为训练设计,启动复杂且易报错。需确保模型和输入均在cuda:0,启用model.eval()和torch.inference_mode(),并注意batchnorm稳定性问题。

PyTorch 多GPU推理该用 torch.nn.DataParallel 还是 torch.nn.parallel.DistributedDataParallel?
直接说结论:**推理场景下优先用 DataParallel,别碰 DistributedDataParallel(DDP)**。DDP 是为分布式训练设计的,启动成本高、必须用 torch.distributed.launch 或 torchrun,还要手动初始化进程组——对纯推理来说完全是杀鸡用牛刀,且容易因 rank/nccl 初始化失败直接卡死或报 RuntimeError: invalid argument。
DataParallel 虽然有 GIL 和主 GPU 内存瓶颈问题,但胜在开箱即用、无额外进程管理、不改模型结构就能跑通。实际测试中,2–4 卡推理吞吐提升基本符合线性预期(比如 2 卡 ≈ 1.8× 单卡),足够应对多数服务化场景。
如何正确包装模型并确保输入张量在主设备上?
DataParallel 要求所有输入张量必须先放到 device='cuda:0'(默认主卡),否则会报 Expected all tensors to be on the same device。它内部会自动把输入切分、复制到各卡,但不会帮你做跨设备搬运。
- 模型加载后立刻用
model = model.cuda(),不要指定具体卡号(如.cuda(0)) - 输入
tensor必须显式调用.to('cuda:0')或.cuda()(等价于 cuda:0) - 千万别写
model.to('cuda:1')或tensor.to('cuda:2')—— 主卡错位会导致 forward 直接崩溃 - 输出结果默认在 cuda:0,无需额外
.to()
示例:
model = MyModel().cuda() # 自动到 cuda:0
model = torch.nn.DataParallel(model)
<p>x = torch.randn(32, 3, 224, 224).cuda() # 必须 cuda(),不能 .to('cuda:1')
y = model(x) # y 也在 cuda:0 上
</p>
为什么 batch size 加倍后显存反而爆了?
根本原因是 DataParallel 的副本机制:模型参数和优化器状态(虽然推理不用优化器)会在每张卡上完整拷贝一份。假设单卡能跑 batch=64,4 卡并行时,**每张卡仍要承载完整模型 + 当前分配给它的子 batch**。所以总显存 ≈ 单卡模型显存 × GPU 数 + 单卡数据显存 × GPU 数。
- batch=64 在单卡占 3GB 模型 + 1GB 数据 → 总 4GB
- 4 卡并行、总 batch=256 时:每卡模型 3GB × 4 = 12GB + 每卡数据 64×1GB = 4GB → 单卡需 16GB,远超单卡容量
- 解决办法只有两个:降低总 batch(如 4 卡用 total_batch=128),或换
torch.compile+torch.inference_mode()压显存
务必在推理前加:torch.inference_mode()(比 torch.no_grad() 更轻量,禁用梯度+部分中间缓存),否则显存多占 15–20%。
使用 DataParallel 后输出 shape 不对或报错 Expected more than 1 value per channel?
这是 BatchNorm 层惹的祸。当子 batch 太小(比如 4 卡跑 total_batch=8,每卡只有 2 个样本),BatchNorm2d 在各卡上独立计算均值方差,会因样本不足导致数值不稳定甚至除零,触发上述错误。
- 最稳解法:把模型里所有
BatchNorm*替成nn.SyncBatchNorm.convert_sync_batchnorm(model)—— 但它只在 DDP 下生效,DataParallel无效 - 实用解法:推理时直接冻结 BN 统计,用训练好的 running_mean / running_var:
model.eval()(必须!)+ 确保没手动调train() - 更彻底:替换为
GroupNorm或LayerNorm,它们不依赖 batch size
漏掉 model.eval() 是高频翻车点,会导致 BN 层持续更新 running stats,不仅出错,还会污染模型状态。
多卡推理真正难的不是写几行 DataParallel,而是得盯着每张卡的显存水位、子 batch 分配是否均匀、BN 层是否真被 freeze —— 这些细节不验证,上线后只会随机 OOM 或精度漂移。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











