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 为输出设备,一次迭代的数据流是:
前向传播:
- GPU0 把输入 batch 切分(Scatter)成多个 mini-batch 分发给各卡,同时把模型复制(Replicate)到各卡;
- 各卡并行前向传播,输出汇集(Gather)回 GPU0。
反向传播:
-
GPU0 计算损失和对各输出的梯度,再分发回各卡;
-
各卡各自反向传播,最终把参数梯度归约(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 行代码就是所有分布式训练框架的骨架:
- 初始化时广播参数------各副本起点一致;
- 每个进程读互斥的数据子集------效果上等价于 batch_size 扩大 world_size 倍;
- backward 后 allreduce 梯度求平均------各副本用同一份梯度更新,参数永远保持一致;
- 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 做了两件事:给每个需要求导的参数注册钩子,梯度一算完立刻发起非阻塞 allreduce ;optimizer.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)。
总结
回顾本文要点:
- 数据并行是主流:每卡一份模型副本、各读互斥数据、allreduce 平均梯度------一切框架都是这个骨架的实现;
- DataParallel:单进程多线程,一行代码可用,但负载不均衡且受 GIL 限制,官方已不推荐;
- 通信原语:broadcast 同步初始参数,allreduce 同步梯度,Ring-AllReduce 让每卡通信量与卡数解耦;
- DDP 四步 :
init_process_group→DistributedSampler(勿忘set_epoch)→ DDP 封装模型 →torch.distributed.launch/torchrun启动;梯度同步在 backward 中与计算重叠,多进程绕开 GIL; - Horovod :跨框架的分布式方案,
DistributedOptimizer靠 backward 钩子异步同步梯度,horovodrun启动; - 工程要点:小 batch 场景启用 syncBN;一切集合通信必须所有进程共同执行,否则程序无声挂起。
到这里,PyTorch 的核心机制(Tensor、autograd、nn)、工程工具(数据加载、可视化、GPU 与分布式)都已经打通。从下一篇开始进入实战环节:用一个规范的项目骨架组织完整的深度学习实验,告别散落一地的脚本。