《深度学习框架PyTorch入门与实践》系列:09-分布式与并行训练之DataParallel、DDP与Horovod

09-分布式与并行训练:DataParallel、DDP与Horovod

一块 GPU 训不动大模型、跑不完大数据集时,多卡训练就是必修课。本文按"原理 → API → 工程实践"的顺序讲透 PyTorch 多卡训练:并行的两种形态 (数据并行/模型并行)与 DataParallel 的工作机制及两大缺陷、集合通信原语 (broadcast/scatter/gather/allreduce 以及 Ring-AllReduce 的带宽优势)、用 40 行 MPI 代码手写一个"极简 DDP" 来理解梯度同步的本质、DistributedDataParallel 的完整四步流程与启动命令Horovod 的用法,最后是 syncBN 与分布式训练最容易踩的同步陷阱。核心代码来自陈云《深度学习框架PyTorch:入门与实践(第2版)》配套仓库。

文章目录

  • 09-分布式与并行训练:DataParallel、DDP与Horovod
    • 一、并行的基本概念:数据并行与模型并行
    • 二、nn.DataParallel:一行代码的单机多卡
      • [2.1 工作原理](#2.1 工作原理)
      • [2.2 使用示例](#2.2 使用示例)
      • [2.3 两大缺点](#2.3 两大缺点)
    • 三、分布式系统与集合通信原语
      • [3.1 rank、world_size、local_rank](#3.1 rank、world_size、local_rank)
      • [3.2 点对点通信与集群通信](#3.2 点对点通信与集群通信)
      • [3.3 Ring-AllReduce:为什么它更快](#3.3 Ring-AllReduce:为什么它更快)
      • [3.4 用 mpi4py 做一次分布式计算](#3.4 用 mpi4py 做一次分布式计算)
    • [四、用 MPI 手写一个"极简 DDP"](#四、用 MPI 手写一个"极简 DDP")
    • [五、DistributedDataParallel 完整流程](#五、DistributedDataParallel 完整流程)
      • [5.1 四步构建](#5.1 四步构建)
      • [5.2 启动方式](#5.2 启动方式)
      • [5.3 DDP 的内部机制](#5.3 DDP 的内部机制)
    • [六、Horovod 简介](#六、Horovod 简介)
    • [七、工程实践:syncBN 与同步陷阱](#七、工程实践:syncBN 与同步陷阱)
    • 总结

一、并行的基本概念:数据并行与模型并行

先厘清两个词:并行 (Parallel)通常指一台服务器上的多块 GPU 协同;分布式(Distributed)指多台服务器上的多块 GPU 协同,涉及跨机通信,更复杂。

无论并行还是分布式,切分工作的方式只有两种:

  • 模型并行:把模型拆开放到不同 GPU 上。比如网络前几层放 GPU0,后几层放 GPU1。适合单卡装不下的超大模型,但层与层之间存在依赖,GPU 利用率不高,实现也麻烦。
  • 数据并行 :每块 GPU 持有完整的模型副本,各自处理一部分数据,反向传播后同步梯度再更新参数。实现简单、扩展性好,是绝对的主流。

本文聚焦数据并行。数据并行的全部难点浓缩成一句话:如何让各卡上的模型参数始终保持一致------初始化时靠广播对齐参数,每一步靠梯度求平均对齐更新。后面所有工具(DataParallel、DDP、Horovod)都是对这句话的不同工程实现。

二、nn.DataParallel:一行代码的单机多卡

2.1 工作原理

python 复制代码
model = nn.DataParallel(model.cuda(), device_ids=gpus, output_device=gpus[0])

DataParallel单进程多线程控制多卡。设 GPU0 为输出设备,一次迭代的数据流是:

前向传播:

  1. GPU0 把输入 batch 切分(Scatter)成多个 mini-batch 分发给各卡,同时把模型复制(Replicate)到各卡;
  2. 各卡并行前向传播,输出汇集(Gather)回 GPU0。

反向传播:

  1. GPU0 计算损失和对各输出的梯度,再分发回各卡;

  2. 各卡各自反向传播,最终把参数梯度归约(Reduce)到 GPU0,只在 GPU0 上更新参数;下一次迭代前 GPU0 再把新参数广播给各卡。
    #mermaid-svg-xZzONh0CA0LTn2FL{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-xZzONh0CA0LTn2FL .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-xZzONh0CA0LTn2FL .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-xZzONh0CA0LTn2FL .error-icon{fill:#552222;}#mermaid-svg-xZzONh0CA0LTn2FL .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-xZzONh0CA0LTn2FL .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-xZzONh0CA0LTn2FL .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-xZzONh0CA0LTn2FL .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-xZzONh0CA0LTn2FL .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-xZzONh0CA0LTn2FL .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-xZzONh0CA0LTn2FL .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-xZzONh0CA0LTn2FL .marker{fill:#333333;stroke:#333333;}#mermaid-svg-xZzONh0CA0LTn2FL .marker.cross{stroke:#333333;}#mermaid-svg-xZzONh0CA0LTn2FL svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-xZzONh0CA0LTn2FL p{margin:0;}#mermaid-svg-xZzONh0CA0LTn2FL .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-xZzONh0CA0LTn2FL .cluster-label text{fill:#333;}#mermaid-svg-xZzONh0CA0LTn2FL .cluster-label span{color:#333;}#mermaid-svg-xZzONh0CA0LTn2FL .cluster-label span p{background-color:transparent;}#mermaid-svg-xZzONh0CA0LTn2FL .label text,#mermaid-svg-xZzONh0CA0LTn2FL span{fill:#333;color:#333;}#mermaid-svg-xZzONh0CA0LTn2FL .node rect,#mermaid-svg-xZzONh0CA0LTn2FL .node circle,#mermaid-svg-xZzONh0CA0LTn2FL .node ellipse,#mermaid-svg-xZzONh0CA0LTn2FL .node polygon,#mermaid-svg-xZzONh0CA0LTn2FL .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-xZzONh0CA0LTn2FL .rough-node .label text,#mermaid-svg-xZzONh0CA0LTn2FL .node .label text,#mermaid-svg-xZzONh0CA0LTn2FL .image-shape .label,#mermaid-svg-xZzONh0CA0LTn2FL .icon-shape .label{text-anchor:middle;}#mermaid-svg-xZzONh0CA0LTn2FL .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-xZzONh0CA0LTn2FL .rough-node .label,#mermaid-svg-xZzONh0CA0LTn2FL .node .label,#mermaid-svg-xZzONh0CA0LTn2FL .image-shape .label,#mermaid-svg-xZzONh0CA0LTn2FL .icon-shape .label{text-align:center;}#mermaid-svg-xZzONh0CA0LTn2FL .node.clickable{cursor:pointer;}#mermaid-svg-xZzONh0CA0LTn2FL .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-xZzONh0CA0LTn2FL .arrowheadPath{fill:#333333;}#mermaid-svg-xZzONh0CA0LTn2FL .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-xZzONh0CA0LTn2FL .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-xZzONh0CA0LTn2FL .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-xZzONh0CA0LTn2FL .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-xZzONh0CA0LTn2FL .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-xZzONh0CA0LTn2FL .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-xZzONh0CA0LTn2FL .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-xZzONh0CA0LTn2FL .cluster text{fill:#333;}#mermaid-svg-xZzONh0CA0LTn2FL .cluster span{color:#333;}#mermaid-svg-xZzONh0CA0LTn2FL div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-xZzONh0CA0LTn2FL .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-xZzONh0CA0LTn2FL rect.text{fill:none;stroke-width:0;}#mermaid-svg-xZzONh0CA0LTn2FL .icon-shape,#mermaid-svg-xZzONh0CA0LTn2FL .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-xZzONh0CA0LTn2FL .icon-shape p,#mermaid-svg-xZzONh0CA0LTn2FL .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-xZzONh0CA0LTn2FL .icon-shape .label rect,#mermaid-svg-xZzONh0CA0LTn2FL .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-xZzONh0CA0LTn2FL .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-xZzONh0CA0LTn2FL .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-xZzONh0CA0LTn2FL :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} Scatter数据 + Replicate模型
    Scatter数据 + Replicate模型
    并行forward
    并行forward
    分发梯度
    Reduce梯度到GPU0
    GPU0: 完整batch
    GPU0 子batch
    GPU1 子batch
    GPU0: Gather输出/算loss
    各卡backward
    GPU0 更新参数, 下轮广播

2.2 使用示例

python 复制代码
import torch
import torch.nn as nn

device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")

# toy model:只打印输入形状,观察数据是怎么被切分的
class Model(nn.Module):
    def forward(self, input):
        print("In Model:", input[0].shape, input[1].shape)
        return input

model = Model()
if torch.cuda.device_count() > 1:
    print("You have", torch.cuda.device_count(), "GPUs!")
    model = nn.DataParallel(model)
model.to(device)

tensorA = torch.Tensor(17, 2).to(device)
tensorB = torch.Tensor(15, 3).to(device)
print("Outside:", tensorA.shape, tensorB.shape)
output = model([tensorA, tensorB])

两卡机器上的输出:

复制代码
You have 2 GPUs!
Outside: torch.Size([17, 2]) torch.Size([15, 3])
In Model: torch.Size([9, 2]) torch.Size([8, 3])
In Model: torch.Size([8, 2]) torch.Size([7, 3])

可以看到 17 个样本被切成 9+8 两份,模型内部拿到的是子 batch。除了包一层 DataParallel,训练代码一行都不用改,这是它至今仍受欢迎的原因。

2.3 两大缺点

  • 负载不均衡:所有梯度都要归约到 GPU0,损失计算也在 GPU0,导致 GPU0 显存和计算负担明显高于其他卡。模型一大,经常是 GPU0 先 OOM 而其他卡还很空。
  • 速度受限:单进程多线程的设计绕不开 Python GIL;每轮参数更新后 GPU0 还要向所有卡广播新参数,通信开销随卡数线性增长。而且单进程注定只能单机,多台机器无法使用。

正因为这些问题,PyTorch 官方现在明确推荐:即使单机多卡,也用 DistributedDataParallel 代替 DataParallel

三、分布式系统与集合通信原语

3.1 rank、world_size、local_rank

分布式训练会同时启动多个进程(通常一个进程管一块 GPU),每个进程都完整执行一遍 python main.py。几个贯穿始终的概念:

  • group:进程组,初始化时默认只有一个组,一般无需手动配置;
  • rank:进程的全局唯一编号,rank=0 通常是主进程;
  • world_size:进程总数;
  • local_rank :进程在本台机器内的编号。例如 2 台机器各 4 卡共 8 个进程:rank 取 0~7,而每台机器内部的 local_rank 都是 0~3。绑定 GPU 用 local_rank(torch.cuda.set_device(local_rank)),跨机通信、控制"只做一次"的逻辑用 rank。

3.2 点对点通信与集群通信

进程间通信分两类。点对点通信 是一个进程发、一个进程收(又分阻塞的 Send/Recv 和非阻塞的 Isend/Irecv,后者能把计算与通信重叠起来提升效率)。深度学习里更常用的是集群通信(Collective Communication),核心原语有:

  • broadcast(广播):一个进程把数据复制给组内所有进程------DDP 初始化时同步模型参数用的就是它;
  • scatter(分发):一个进程把数据切片后分给不同进程;
  • gather(汇集):所有进程的数据收到某一个进程;变体有 all-gather(收完后每个进程都有全量数据);
  • reduce(归约):对所有进程的数据做求和/取最大等运算,结果放到一个进程上;
  • allreduce :reduce 的变体,归约结果同步到每一个进程------数据并行中同步梯度用的正是它;
  • barrier(屏障):所有进程在此对齐,先到的等后到的。

PyTorch 支持三种通信后端:Gloo (CPU 通用)、NCCL (NVIDIA GPU 首选,走 NVLink/PCIe 优化路径)、MPI。GPU 训练无脑选 NCCL 即可。

3.3 Ring-AllReduce:为什么它更快

朴素的 allreduce 实现是"所有卡把梯度发给主卡,主卡求平均后再广播回去",主卡的通信量随卡数线性增长,卡越多越堵。NCCL 和 Horovod 采用的 Ring-AllReduce 把 N 块卡连成一个环,每块卡只和左右邻居通信:
#mermaid-svg-UjrKZoKroLuczca5{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-UjrKZoKroLuczca5 .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-UjrKZoKroLuczca5 .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-UjrKZoKroLuczca5 .error-icon{fill:#552222;}#mermaid-svg-UjrKZoKroLuczca5 .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-UjrKZoKroLuczca5 .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-UjrKZoKroLuczca5 .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-UjrKZoKroLuczca5 .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-UjrKZoKroLuczca5 .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-UjrKZoKroLuczca5 .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-UjrKZoKroLuczca5 .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-UjrKZoKroLuczca5 .marker{fill:#333333;stroke:#333333;}#mermaid-svg-UjrKZoKroLuczca5 .marker.cross{stroke:#333333;}#mermaid-svg-UjrKZoKroLuczca5 svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-UjrKZoKroLuczca5 p{margin:0;}#mermaid-svg-UjrKZoKroLuczca5 .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-UjrKZoKroLuczca5 .cluster-label text{fill:#333;}#mermaid-svg-UjrKZoKroLuczca5 .cluster-label span{color:#333;}#mermaid-svg-UjrKZoKroLuczca5 .cluster-label span p{background-color:transparent;}#mermaid-svg-UjrKZoKroLuczca5 .label text,#mermaid-svg-UjrKZoKroLuczca5 span{fill:#333;color:#333;}#mermaid-svg-UjrKZoKroLuczca5 .node rect,#mermaid-svg-UjrKZoKroLuczca5 .node circle,#mermaid-svg-UjrKZoKroLuczca5 .node ellipse,#mermaid-svg-UjrKZoKroLuczca5 .node polygon,#mermaid-svg-UjrKZoKroLuczca5 .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-UjrKZoKroLuczca5 .rough-node .label text,#mermaid-svg-UjrKZoKroLuczca5 .node .label text,#mermaid-svg-UjrKZoKroLuczca5 .image-shape .label,#mermaid-svg-UjrKZoKroLuczca5 .icon-shape .label{text-anchor:middle;}#mermaid-svg-UjrKZoKroLuczca5 .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-UjrKZoKroLuczca5 .rough-node .label,#mermaid-svg-UjrKZoKroLuczca5 .node .label,#mermaid-svg-UjrKZoKroLuczca5 .image-shape .label,#mermaid-svg-UjrKZoKroLuczca5 .icon-shape .label{text-align:center;}#mermaid-svg-UjrKZoKroLuczca5 .node.clickable{cursor:pointer;}#mermaid-svg-UjrKZoKroLuczca5 .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-UjrKZoKroLuczca5 .arrowheadPath{fill:#333333;}#mermaid-svg-UjrKZoKroLuczca5 .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-UjrKZoKroLuczca5 .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-UjrKZoKroLuczca5 .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-UjrKZoKroLuczca5 .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-UjrKZoKroLuczca5 .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-UjrKZoKroLuczca5 .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-UjrKZoKroLuczca5 .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-UjrKZoKroLuczca5 .cluster text{fill:#333;}#mermaid-svg-UjrKZoKroLuczca5 .cluster span{color:#333;}#mermaid-svg-UjrKZoKroLuczca5 div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-UjrKZoKroLuczca5 .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-UjrKZoKroLuczca5 rect.text{fill:none;stroke-width:0;}#mermaid-svg-UjrKZoKroLuczca5 .icon-shape,#mermaid-svg-UjrKZoKroLuczca5 .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-UjrKZoKroLuczca5 .icon-shape p,#mermaid-svg-UjrKZoKroLuczca5 .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-UjrKZoKroLuczca5 .icon-shape .label rect,#mermaid-svg-UjrKZoKroLuczca5 .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-UjrKZoKroLuczca5 .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-UjrKZoKroLuczca5 .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-UjrKZoKroLuczca5 :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 梯度分块传递
GPU0
GPU1
GPU2
GPU3

算法分两个阶段:先做 N-1 步 scatter-reduce (每步每卡把梯度的一个分块发给下家并累加收到的分块,结束时每卡各持有一个"已聚合完毕"的分块),再做 N-1 步 all-gather (把聚合好的分块沿环传一圈)。每块卡的总通信量约为 2(N-1)/N × 梯度大小与卡数几乎无关------这就是 Ring-AllReduce 能把数据并行扩展到大规模集群的原因。使用 NCCL 后端时这一切自动发生,不需要写任何代码,但理解它有助于解释"为什么 DDP 比 DataParallel 快"。

3.4 用 mpi4py 做一次分布式计算

用 MPI 的 Python 封装 mpi4py 直观感受一下这些原语(pip install mpi4py,Linux 还需 apt install libopenmpi-dev openmpi-bin):

python 复制代码
# Matrix.py
import mpi4py.MPI as MPI
import numpy as np

comm = MPI.COMM_WORLD
rank = comm.Get_rank()
size = comm.Get_size()          # world_size

# 只在 rank-0 初始化数据
if rank == 0:
    array = np.arange(8)
    splits = np.split(array, size)   # 切成 size 份
else:
    splits = None

# scatter:把切片分发给各进程
local_data = comm.scatter(splits, root=0)

# 各进程本地求和,再 allreduce 得到全局和(每个进程都拿到结果)
local_sum = local_data.sum()
all_sum = comm.allreduce(local_sum, op=MPI.SUM)

# 各进程本地平方,再 allgather 汇集
result = np.vstack(comm.allgather(local_data ** 2))

if rank == 1:
    print("元素和为:", all_sum)          # 28
    print("平方结果:\n", result)

mpiexec -n 2 python Matrix.py 启动------mpiexec 会启动 2 个 Python 进程执行同一份代码,只是各自拿到不同的 rank。注意直接 python Matrix.py 跑则 world_size=1,这是分布式脚本和普通脚本最根本的差别。

四、用 MPI 手写一个"极简 DDP"

在用封装好的框架之前,先用 MPI 手写一遍数据并行训练,把"分解、同步、聚合、控制"四个要点看得清清楚楚:

python 复制代码
# MPI_PyTorch.py
import torch
import torchvision as tv
import mpi4py.MPI as MPI

## 第一步:环境初始化
comm = MPI.COMM_WORLD
rank = comm.Get_rank()
size = comm.Get_size()
torch.cuda.set_device(rank)     # 每个进程绑定一块 GPU

## 第二步:数据【分解】------每个进程只拿数据集的 1/size
dataset = tv.datasets.CIFAR10(root="./", download=True,
                              transform=tv.transforms.ToTensor())
# X[rank::size]:从第 rank 个元素起每隔 size 取一个,天然互斥
dataset.data = dataset.data[rank::size]
dataset.targets = dataset.targets[rank::size]
dataloader = torch.utils.data.DataLoader(dataset, batch_size=512)

## 第三步:模型【同步】------广播 rank0 的初始参数,保证各副本起点一致
model = tv.models.resnet18(pretrained=False).cuda()
for name, param in model.named_parameters():
    param_from_rank_0 = comm.bcast(param.detach(), root=0)
    param.data.copy_(param_from_rank_0)

lr = 0.001
loss_fn = torch.nn.CrossEntropyLoss().cuda()

## 第四步:训练------各进程独立 forward/backward,梯度【聚合】后更新
for ii, (data, target) in enumerate(dataloader):
    output = model(data.cuda())
    loss = loss_fn(output, target.cuda())
    loss.backward()
    # 核心:allreduce 求所有进程梯度的平均值,再做梯度下降
    for name, param in model.named_parameters():
        grad_sum = comm.allreduce(param.grad.detach().cpu(), op=MPI.SUM)
        grad_mean = grad_sum / size
        param.data -= lr * grad_mean.cuda()

## 【控制】:只在 rank-0 保存模型
if rank == 0:
    torch.save(model.state_dict(), "./000.ckpt")

mpiexec -n 4 python MPI_PyTorch.py 即可 4 卡训练。这 40 行代码就是所有分布式训练框架的骨架:

  1. 初始化时广播参数------各副本起点一致;
  2. 每个进程读互斥的数据子集------效果上等价于 batch_size 扩大 world_size 倍;
  3. backward 后 allreduce 梯度求平均------各副本用同一份梯度更新,参数永远保持一致;
  4. rank 判断控制副作用------保存、打印只做一次。

理解了这个骨架,DDP 和 Horovod 就只是"把第 3 步做得更快、更自动"而已。

五、DistributedDataParallel 完整流程

5.1 四步构建

torch.distributed 是 PyTorch 官方的分布式接口,DistributedDataParallel(DDP)是其高层封装。对照上一节的手写版本,改动小得惊人:

python 复制代码
# distributed_PyTorch.py
import torch
import torch.distributed as dist
import torchvision as tv

## 第一步:初始化进程组,NCCL 后端
dist.init_process_group(backend='nccl')
local_rank = dist.get_rank()
torch.cuda.set_device(local_rank)

## 第二步:数据------DistributedSampler 自动为每个进程划分互斥子集
dataset = tv.datasets.CIFAR10(root="./", download=True,
                              transform=tv.transforms.ToTensor())
sampler = torch.utils.data.DistributedSampler(dataset)
dataloader = torch.utils.data.DataLoader(dataset, batch_size=512, sampler=sampler)

## 第三步:模型------DDP 封装,自动完成参数广播与梯度 allreduce
model = tv.models.resnet18(pretrained=False).cuda()
model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[local_rank])
loss_fn = torch.nn.CrossEntropyLoss().cuda()

## 第四步:训练------写法与单卡完全相同
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
for epoch in range(num_epochs):
    # 每个 epoch 设置一次,保证各 epoch 的 shuffle 不同且各进程一致
    sampler.set_epoch(epoch)
    for data, target in dataloader:
        optimizer.zero_grad()
        output = model(data.cuda())
        loss = loss_fn(output, target.cuda())
        loss.backward()     # 梯度 allreduce 在这一步自动发生
        optimizer.step()

if local_rank == 0:
    torch.save(model.state_dict(), "./000.ckpt")

三个关键点:

  • init_process_group 必须最先调用,它读取启动器注入的环境变量(RANK、WORLD_SIZE、MASTER_ADDR 等)完成组网;
  • DistributedSampler 不能省 :DataLoader 默认每个进程都会遍历全量数据,等于白白重复训练;用了 sampler 后 shuffle 参数交给 sampler.set_epoch(epoch) 控制;
  • 手写版里那段逐参数 allreduce 的循环,现在浓缩在 loss.backward() 里自动完成。

5.2 启动方式

DDP 需要专门的启动器为每个进程注入 rank 等环境变量。单机多卡:

bash 复制代码
# nproc_per_node 一般等于本机 GPU 数
python -m torch.distributed.launch --nproc_per_node=2 distributed_PyTorch.py

多机多卡(2 台机器各 4 卡为例,两台机器都要执行,node_rank 分别为 0 和 1):

bash 复制代码
python -m torch.distributed.launch \
    --nproc_per_node=4 --nnodes=2 --node_rank=0 \
    --master_addr="serverA" distributed_PyTorch.py

另一种方式是在代码里用 torch.multiprocessing.spawn 自行拉起多个进程。新版本 PyTorch 中 torch.distributed.launch 已由 torchrun 接替,参数基本一致,迁移成本很低。

5.3 DDP 的内部机制

DDP 的行为可以用一段伪代码概括:

python 复制代码
class DistributedDataParallel(nn.Module):
    def forward(self, input):
        # 前向传播不做任何同步
        return self.model(input)

    def backward(self, loss):
        loss.backward()
        # 反向传播时多了一步:把本进程梯度与所有进程 allreduce
        for name, param in self.model.named_parameters():
            allreduce(name, param)

实际实现更聪明:DDP 把参数分桶(bucket),某个桶内所有梯度一算完就立刻发起异步 allreduce,让通信和剩余层的反向计算重叠进行。配合 Ring-AllReduce,这就是 DDP 对 DataParallel 的双重优势来源:多进程绕开 GIL、通信隐藏在计算背后。

六、Horovod 简介

Horovod 是 Uber 开源的第三方分布式训练框架,支持 TensorFlow/PyTorch/MXNet。一个好记的类比:NCCL 相当于 GPU 版的加强 MPI,Horovod 相当于 PyTorch 版的 mpi4py 。安装需要指定 CUDA 与 NCCL 路径(这也是它最大的使用门槛,注意 CUDA 版本必须与 torch.version.cuda 一致):

bash 复制代码
HOROVOD_NCCL_HOME=/usr/local/nccl-2 HOROVOD_CUDA_HOME=/usr/local/cuda \
HOROVOD_GPU_OPERATIONS=NCCL pip install --no-cache-dir horovod

训练代码与 DDP 高度同构:

python 复制代码
# Horovod_PyTorch.py
import horovod.torch as hvd
import torch
import torchvision as tv

## 初始化
hvd.init()
torch.cuda.set_device(hvd.local_rank())

## 数据:同样使用 DistributedSampler
dataset = tv.datasets.CIFAR10(root="./", download=True,
                              transform=tv.transforms.ToTensor())
sampler = torch.utils.data.DistributedSampler(
    dataset, num_replicas=hvd.size(), rank=hvd.rank())
dataloader = torch.utils.data.DataLoader(dataset, batch_size=512, sampler=sampler)

## 模型:手动广播初始参数(对应 DDP 构造时自动做的事)
model = tv.models.resnet18(pretrained=False).cuda()
hvd.broadcast_parameters(model.state_dict(), root_rank=0)
loss_fn = torch.nn.CrossEntropyLoss().cuda()

## 优化器:用 DistributedOptimizer 包装,梯度同步由它接管
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
optimizer = hvd.DistributedOptimizer(optimizer,
                                     named_parameters=model.named_parameters())
# 训练循环与单卡完全一致(略)

if hvd.rank() == 0:
    torch.save(model.state_dict(), "./000.ckpt")

DistributedOptimizer 做了两件事:给每个需要求导的参数注册钩子,梯度一算完立刻发起非阻塞 allreduceoptimizer.step() 前等待所有同步完成。伪代码:

python 复制代码
class DistributedOptimizer(Optimizer):
    def __init__(self, optimizer, params):
        self.optimizer = optimizer
        for param in params:
            # 反向传播中逐参数异步同步梯度
            param.register_backward_hook(lambda tensor: hvd.allreduce(tensor))

    def step(self):
        wait_for_all_reduce_done()   # 确保所有 allreduce 完成
        self.optimizer.step()

启动用 horovodrun(对 mpirun 的封装):

bash 复制代码
horovodrun -np 4 python main.py                          # 单机 4 卡
horovodrun -np 8 -H serverA:4,serverB:4 python main.py   # 两机各 4 卡

DDP 还是 Horovod?两者性能相当、原理相同。DDP 零安装成本、与 PyTorch 生态贴合最紧;Horovod 的优势在跨框架统一 API 和一些附加能力(如缓解大 batch 收敛问题的 Adasum 算法、支持进程数动态伸缩的 Elastic Horovod)。纯 PyTorch 项目建议首选 DDP。

七、工程实践:syncBN 与同步陷阱

跨卡同步 BN(syncBN) 。BatchNorm 统计的是 batch 内的均值方差,而数据并行下每个进程只见到 batch_size/world_size 个样本,BN 统计量的噪声会明显变大,卡数多、单卡 batch 小时掉点严重。解决办法是让 BN 层跨进程通信、用全局 batch 统计量,torch.distributed 与 Horovod 都已封装。DDP 下一行转换即可:

python 复制代码
# 把模型中所有 BatchNorm 层替换为跨卡同步版本,须在 DDP 封装之前调用
model = torch.nn.SyncBatchNorm.convert_sync_batchnorm(model)
model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[local_rank])

同步陷阱。分布式 bug 的一大特征是"不报错,只挂起"------某个集合通信操作只有部分进程执行到了,其余进程会永远等待。几个真实易犯的例子:

python 复制代码
# 陷阱一:条件跳过 backward。某进程跳过后,其他进程的 allreduce 永远等不到它
for data in dataloader:
    if data is None:
        continue          # 危险!各进程的迭代步数可能不一致
    loss = forward(); loss.backward(); optimizer.step()

# 陷阱二:变量只在 rank-0 定义,广播时其他进程直接 NameError
if rank == 0:
    b = load_checkpoint()
b = broadcast(b, root=0)  # 其他进程里 b 未定义

# 陷阱三:rank-0 建目录,其他进程立刻写文件------目录可能还没建好
if rank == 0:
    os.makedirs('/path/to/logdir')
_ = hvd.allreduce(torch.Tensor(), name='barrier')  # 用 allreduce 充当屏障
logger = Logger(f'/path/to/logdir/{rank}.log')

配套的几条工程经验:

  • 可视化、保存 checkpoint 等副作用统一包在 if rank == 0: 里,避免多进程写同一文件;
  • 手动修改、新增了参数后记得重新广播,否则各副本悄悄不一致,训练效果变差却不报错;
  • 调试遵循"先小后大":先 --nproc_per_node=1 单进程跑通(可用 ipdb),再两进程验证同步逻辑,最后上大集群,能极大降低复现成本;
  • 卡数增多等价于增大全局 batch_size,每个 epoch 的更新次数变少,学习率策略需要相应调整(简单线性放大学习率可能让训练发散,可了解 warmup 或 Horovod 的 Adasum)。

总结

回顾本文要点:

  1. 数据并行是主流:每卡一份模型副本、各读互斥数据、allreduce 平均梯度------一切框架都是这个骨架的实现;
  2. DataParallel:单进程多线程,一行代码可用,但负载不均衡且受 GIL 限制,官方已不推荐;
  3. 通信原语:broadcast 同步初始参数,allreduce 同步梯度,Ring-AllReduce 让每卡通信量与卡数解耦;
  4. DDP 四步init_process_groupDistributedSampler(勿忘 set_epoch)→ DDP 封装模型 → torch.distributed.launch/torchrun 启动;梯度同步在 backward 中与计算重叠,多进程绕开 GIL;
  5. Horovod :跨框架的分布式方案,DistributedOptimizer 靠 backward 钩子异步同步梯度,horovodrun 启动;
  6. 工程要点:小 batch 场景启用 syncBN;一切集合通信必须所有进程共同执行,否则程序无声挂起。

到这里,PyTorch 的核心机制(Tensor、autograd、nn)、工程工具(数据加载、可视化、GPU 与分布式)都已经打通。从下一篇开始进入实战环节:用一个规范的项目骨架组织完整的深度学习实验,告别散落一地的脚本。

相关推荐
心运软件2 小时前
基于深度学习的宝石图像分类系统
人工智能·python·深度学习·机器学习·分类·数据挖掘
雨晨源码(同名B站)2 小时前
【2027届人工智能专业选题】基于yolov8的农业病虫害图像识别与分类系统 |深度学习 计算机视觉
人工智能·深度学习·yolo·计算机视觉·分类
OpenApi.cc3 小时前
MoCode — AI Agent Server (Docker) + macOS Desktop Client(开源项目)
数据结构·人工智能·深度学习·神经网络
一个王同学3 小时前
从零到一 | CV转多模态大模型 | week22 | 实战项目-DocuMind-VL:基于 OCR 与 Qwen-VL 的文档多模态问答系统(二)
人工智能·深度学习·机器学习·计算机视觉·ocr
美狐美颜SDK开放平台3 小时前
第三方美颜SDK接入教程:直播APP实现实时美颜、滤镜与美型效果的方法
人工智能·深度学习·音视频·sdk·美颜sdk·视频美颜sdk·直播app开发
狂奔蜗牛(bradley)15 小时前
深度学习三大基础激活函数详解:Sigmoid、Tanh、ReLU 公式、导数、图像与优缺点对比
人工智能·深度学习
船厂电气自动化ai大模型16 小时前
AI大模型与数学 第32课 函数凹凸性与二阶导数:拐点求解、凹凸区间计算(10道二阶导数计算题)
数据结构·人工智能·python·深度学习·算法
小O的算法实验室17 小时前
IEEE TII,学习为多目标深度学习生成偏好
人工智能·深度学习·学习