首页 / PyTorch 入门教程 / 分布式训练入门:DDP

PyTorch 入门教程

分布式训练入门:DDP

本教程共 60 篇 · 第 56 篇 · 更新于 2026-08-17 · 约 4 分钟阅读

PyTorch分布式训练DDPtorchrun多卡训练NCCL

本节目标:搞懂分布式训练的基本概念,学会写一个能跑的 DDP 训练脚本,知道怎么用 torchrun 启动多卡训练。

为什么要分布式

模型越来越大,数据越来越多,单张 GPU 会撞上两堵墙:

  • 显存墙:模型参数、优化器状态、中间激活全要占显存,放不下。
  • 时间墙:单卡训练要几周甚至几个月,等不起。

分布式训练(Distributed Training)把计算摊到多张卡、多台机器上。并行方式分两大类:数据并行(Data Parallel)是每张卡放一份完整模型,各算各的数据,最后同步梯度;模型并行(Model Parallel)是把模型拆开分给多张卡。模型放得下单卡时,数据并行是绝对主流,本章就讲它。

先认识几个词

分布式训练是一堆进程在协作,术语绕不开:

  • rank:进程的全局编号,从 0 开始。rank 0 通常是主进程,负责汇总、保存。
  • world_size:参与训练的进程总数。
  • local_rank:进程在「本机」的编号。单机 4 卡时,local_rank 0-3 对应 GPU 0-3。
  • 进程组(process group):参与通信的进程集合,默认就是全部进程。
  • 后端(backend):通信实现。GPU 用 NCCL(最快),CPU 或调试用 Gloo。

进程之间靠环境变量互相找到对方:MASTER_ADDR 是 rank 0 所在机器的地址,MASTER_PORT 是协调端口。这些由启动器自动设置,代码里直接读。

Note

Windows 上 NCCL 不可用,torch.distributed 只支持 Gloo 后端,初始化要用文件或 TCP 方式(init_method="file:///...")。想正经跑多卡训练,还是 Linux 环境省心。

还有一个全局视角:DDP 把 batch 摊到 N 张卡上,等效单卡 batch 是 batch_size × world_size。想保持和单卡相同的收敛行为,总 batch 变大后通常要把学习率也相应调大,或把每卡 batch 调小。分布式不是免费的午餐,超参数要跟着卡数走。

DataParallel 不推荐了,用 DDP

老读者可能听过 torch.nn.DataParallel,包一层就能多卡,但它问题不少:单进程多线程、只能单机、主卡要汇总所有梯度导致显存和通信都成瓶颈。官方早已不推荐使用,多卡场景一律推荐 DDP(DistributedDataParallel,分布式数据并行)。

DDP 的做法是每个进程一份独立模型副本。反向传播时,每个进程算自己的梯度,然后通过 all-reduce 集合通信把梯度同步成平均值,再用一致梯度更新参数。这样所有进程的模型始终一致,效果等价于「把 batch 加大 N 倍」的单卡训练。

几个进程凑成一圈轮流传数据,就是集合通信(collective communication)。除了 DDP 依赖的 all_reduce(所有进程求和再平均),常用的还有:broadcast 把一份数据广播给所有进程、gather 把数据汇总到指定进程、barrier 让所有进程对齐到同一时刻。这些都是 torch.distributed 的底层原语,DDP 帮你调好了,你只需要认识名字。

DDP 的实现有两个细节值得知道。它给每个参数挂了一个 autograd 钩子,反向传播算出一个参数的梯度,钩子立刻把它丢进通信队列,通信和计算重叠进行;梯度按「桶」(bucket)打包一起传,减少小消息的启动开销,backward() 结束时所有梯度已经同步完成。传输算法用的是 Ring AllReduce:进程排成一圈,梯度分块轮流传,每块数据只走固定圈数就全同步完,带宽利用很均匀。

DDP 把同步细节全封装了,你的训练代码几乎不用变。梯度同步发生在 backward() 内部,和反向计算重叠,几乎不额外耗时。

完整的 DDP 脚本

把下面的代码存成 train_ddp.py。它假设你有两张 GPU 卡:

import os
import torch
import torch.nn as nn
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader, DistributedSampler, TensorDataset

def setup():
    # torchrun 会自动设置好 RANK / LOCAL_RANK / WORLD_SIZE
    local_rank = int(os.environ["LOCAL_RANK"])
    torch.cuda.set_device(local_rank)
    dist.init_process_group("nccl")   # 参数都从环境变量读
    return local_rank

class Net(nn.Module):
    def __init__(self):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(32, 64), nn.ReLU(), nn.Linear(64, 10))

    def forward(self, x):
        return self.net(x)

def main():
    local_rank = setup()
    world_size = dist.get_world_size()

    model = Net().cuda(local_rank)
    ddp_model = DDP(model, device_ids=[local_rank])   # 包装成 DDP

    # DistributedSampler 把数据切成 world_size 份,每进程拿一份
    dataset = TensorDataset(torch.randn(2000, 32), torch.randint(0, 10, (2000,)))
    sampler = DistributedSampler(dataset, shuffle=True)
    loader = DataLoader(dataset, batch_size=32, sampler=sampler)

    optimizer = torch.optim.Adam(ddp_model.parameters(), lr=1e-3)
    loss_fn = nn.CrossEntropyLoss()

    for epoch in range(3):
        sampler.set_epoch(epoch)        # 关键!每轮重新打乱数据划分
        for x, y in loader:
            x, y = x.cuda(local_rank), y.cuda(local_rank)
            optimizer.zero_grad()
            loss = loss_fn(ddp_model(x), y)
            loss.backward()             # DDP 在这里自动同步梯度
            optimizer.step()

        if dist.get_rank() == 0:
            print(f"epoch {epoch} loss {loss.item():.4f}")

    dist.destroy_process_group()

if __name__ == "__main__":
    main()

启动命令一行:

torchrun --nproc_per_node=2 train_ddp.py

--nproc_per_node 是本机进程数,一般等于 GPU 数。多机训练再加参数:--nnodes=2 --node_rank=0 --master_addr=主节点IP --master_port=29500,每台机器跑一次,node_rank 从 0 编号。

Note

脚本里用 torch.cuda.set_device 绑定 GPU,是沿用多年的经典写法。PyTorch 2.12 起新出的 torch.accelerator 是设备无关的新接口(torch.accelerator.set_device_indextorch.accelerator.current_accelerator()),一套代码通吃 CUDA、MPS、XPU。新项目可以往这个方向写,老代码不着急迁,功能等价。

torchrun 还内置容错:进程崩了自动重启(elastic 模式),配合检查点能从断点继续。这对多机训练很重要——机器越多,任一台掉线概率越大,没有容错的话一次故障毁掉整晚训练。

Warning

忘了 sampler.set_epoch(epoch) 是新手第一坑。不设置的话,每个 epoch 的数据划分都一样,等于一直用同一批数据训练,结果偏得离谱。另外 DataLoader 里不要再传 shuffle=True,会和 DistributedSampler 冲突。

DistributedSampler 还有两个参数值得认识。num_replicasrank 默认会从进程组自动读,一般不用传;drop_last=True 会在样本数不能被进程数整除时丢掉尾巴,避免各进程 batch 数量不一致导致同步等待。小数据集上建议加上。

进程组初始化时 init_process_group 还可以显式传 rankworld_size,配合 init_method="env://" 从环境变量读。用 torchrun 启动时这些环境变量已经备好,直接不传参最省事;自己写 mp.spawn 派进程时则必须手动传。

保存检查点:只在 rank 0 存

所有进程的模型参数始终一致,检查点只让 rank 0 保存一份就够了。其他进程等它存完再加载,中间用 barrier 对齐:

if dist.get_rank() == 0:
    torch.save(ddp_model.state_dict(), "model.pt")

dist.barrier()   # 等 rank 0 存完
map_location = {"cuda:0": f"cuda:{local_rank}"}
ddp_model.load_state_dict(
    torch.load("model.pt", map_location=map_location, weights_only=True))

map_location 必须写对。不然所有进程都会把权重加载到同一张卡上,直接崩。

常见问题

  • 卡死了:先查 MASTER_ADDRMASTER_PORT 是否一致,防火墙是否放行;同步点必须所有进程都到,一个进程报错全体挂起。
  • 速度不升反降:模型太小,通信开销盖过了并行收益;或 batch size 没跟着 world_size 放大。
  • 想打日志:只在 rank 0 打印,不然刷屏。
  • 调试:设 os.environ["NCCL_DEBUG"] = "INFO" 看通信细节;watch nvidia-smi 确认每张卡都在干活,如果只有一张卡在跑,多半是进程没绑对设备。
  • 只有一张卡:DDP 要求每进程独占一张 GPU。单卡环境就不要上 DDP 了,直接普通训练。
Tip

DDP 可以无缝叠加第 54 章的 AMP:把 autocastGradScaler 套进训练循环即可,梯度同步照常工作。模型大到单卡放不下时,DDP 就无能为力了,那是第 57 章 FSDP 的战场——把模型参数也分片到多张卡上。先把 DDP 跑熟,再往上走。

叠加 AMP 时只需要两个小改动:前向包进 autocast,反向换成 scaler.scale(loss).backward()。DDP 会在缩放后的梯度上做同步,等 scaler.step() 反缩放时每个进程的结果仍然一致,两者协作天衣无缝。