如何用PyTorch做数据并行

GPU
小华
2026-09-18

在 PyTorch 中,数据并行(Data Parallelism)通常指把同一个模型复制到多个 GPU 上,把一批数据拆分到不同 GPU 上并行计算,再汇总结果。常见做法有两种:

一、最简单: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,多进程多卡)

生产环境、训练大模型首选,效率更高,支持单机多卡 & 多机多卡。

1. 启动方式(推荐 torchrun

torchrun \
--nproc_per_node=4 \
train.py

2. 示例代码(核心结构)

import 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)

DDP 关键点

  • 每个 GPU 一个进程
  • 使用 DistributedSampler 避免数据重复
  • 梯度在反向传播时自动 AllReduce
  • DataParallel 快很多

三、DataParallel vs DistributedDataParallel

对比项DataParallelDistributedDataParallel
多进程
效率较低
多机支持
推荐程度教学/调试生产/训练

四、常见坑

  • ❌ 在 DataParallel 里直接 model.module 忘了写
  • ❌ DDP 忘记用 DistributedSampler
  • batch_size 写成“总 batch”而不是“单卡 batch”
  • ✅ 大模型优先 DDP + 混合精度 (torch.cuda.amp)

如果你愿意,我可以:

  • 给你一个 可运行的 DDP 训练模板
  • 数据并行 vs 模型并行
  • 或针对 你现在的模型/硬件给具体方案
亿速云提供售前/售后服务

售前业务咨询

售后技术保障

400-100-2938

7*24小时售后电话

官方微信小程序