分布式训练与并行策略:DDP、张量并行、流水线并行与 ZeRO

从单卡到千卡集群的完整路线图:数据并行与 DDP 的梯度同步机制、all-reduce/all-gather 等通信原语、模型并行与张量并行切分方式、流水线并行的气泡与调度、ZeRO 三阶段与 FSDP 的显存账本、DeepSpeed 配置实战、梯度累积与混合精度配合,以及多机训练卡死、NCCL 超时、显存 OOM 的排错手册。

引言

当模型参数从百万级涨到十亿、百亿级,单张 GPU 立刻撞上两堵墙:显存装不下(参数 + 梯度 + 优化器状态轻松超过 80GB)和训练太慢(一个 epoch 要跑几周)。分布式训练就是拆掉这两堵墙的工程学——把计算、参数、数据切分到多张卡、多台机器上协同完成一次迭代。

本文按「先会用、再懂原理、最后能排错」的顺序展开:先用 DDP 把数据并行跑通,再理解 all-reduce 这类通信原语为什么是性能瓶颈,然后进入张量并行、流水线并行、ZeRO/FSDP 这些大模型必备的切分术,最后给出多机训练的排错清单。全文代码基于 PyTorch 官方 DDP + DeepSpeed。

前置:单卡训练循环与优化器基础见 https://plumephp.com/ml-deep-learning-advanced/;显存与量化压缩的配合见 https://plumephp.com/ml-model-compression-quantization/;优化器状态占用的原理见 https://plumephp.com/ml-gradient-descent-optimizers/;训练完成后上线见 https://plumephp.com/ml-model-deployment/。


目录


1. 为什么需要分布式训练

1.1 两道墙:显存与时间

训练一个大模型时,显存被四样东西瓜分:模型参数、梯度、优化器状态、激活值。以 Adam 为例,每个参数需要:2 字节权重(fp16)+ 2 字节梯度 + 4 字节一阶动量 + 4 字节二阶动量 + 4 字节 fp32 主权重 = 16 字节/参数。一个 7B 模型仅状态就要 112GB,远超单张 A100 的 80GB。

资源单卡瓶颈分布式解法
显存(参数/梯度/优化器)模型装不下ZeRO / FSDP / 模型并行
显存(激活值)batch 开不大梯度检查点 / 序列并行
计算吞吐迭代太慢数据并行
单机卡数上限8 卡封顶多机 + 集合通信

1.2 四种并行的坐标系

数据并行(Data Parallel)  :每卡一份完整模型,切分数据 → 最常用
张量并行(Tensor Parallel) :把单层矩阵切到多卡 → 层内切分
流水线并行(Pipeline Parallel):把不同层放到不同卡 → 层间切分
专家并行(Expert Parallel) :MoE 的不同专家分到不同卡 → 稀疏激活

前三种可组合成 3D 并行,是训练千亿模型的标配:数据并行 × 张量并行 × 流水线并行。

一句话:分布式训练的本质是「用通信换显存和时间」——切得越细,单卡负担越轻,但卡之间的通信开销越大,找到平衡点是全部工程的核心。

1.3 什么时候不需要分布式

不是所有任务都值得上多卡:模型能塞进单卡时单卡调好超参收益更高;卡间带宽低(无 NVLink、跨机房)时数据并行同步可能吃掉全部加速比;batch 太小时每卡样本不足,梯度噪声大、收敛变差。分布式是「不得已而为之」的手段,不是性能银弹。


2. 数据并行与 DDP

2.1 DP 与 DDP 的区别

PyTorch 有两代数据并行实现,理解差异很重要:

维度DataParallel (DP)DistributedParallel (DDP)
进程模型单进程多线程多进程(每卡一进程)
通信方式主卡 gather/scatterall-reduce 梯度
GIL 影响严重无
多机支持不支持原生支持
推荐度已废弃生产标准

DP 的所有前向都在主卡汇总,主卡成瓶颈;DDP 每个进程独立跑完整前向反向,只在反向结束时同步梯度,效率高得多,因此是生产标准。

2.2 DDP 的核心机制:梯度 all-reduce

DDP 的关键设计是反向传播时就通信,而不是等全部梯度算完再同步。它把梯度按 bucket 分桶,某个 bucket 的梯度算完立刻触发 all-reduce,与后续层的反向计算重叠(overlap),把通信藏进计算里。

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

def setup():
    dist.init_process_group(backend="nccl")
    local_rank = int(os.environ["LOCAL_RANK"])
    torch.cuda.set_device(local_rank)
    return local_rank

def main():
    local_rank = setup()
    model = MyModel().to(local_rank)
    model = DDP(model, device_ids=[local_rank])

    dataset = MyDataset()
    sampler = DistributedSampler(dataset, shuffle=True)
    loader = torch.utils.data.DataLoader(
        dataset, batch_size=32, sampler=sampler, num_workers=4
    )

    for epoch in range(num_epochs):
        sampler.set_epoch(epoch)  # 关键:每个 epoch 打乱不同
        for batch in loader:
            loss = model(**batch)
            loss.backward()
            optimizer.step()
            optimizer.zero_grad(set_to_none=True)

if __name__ == "__main__":
    main()

启动方式用 torchrun:

torchrun --nproc_per_node=8 --nnodes=2 --node_rank=0 \
  --master_addr=10.0.0.1 --master_port=29500 train.py

2.3 学习率与 batch size 的线性缩放

DDP 让全局 batch = 单卡 batch × 卡数。全局 batch 变大后梯度噪声降低,需要相应放大学习率:线性缩放 lr_new = lr_base × (global_batch / base_batch),或用平方根缩放。实践建议先做 warmup(前 5% 步数线性升温),全局 batch 超过 8K 后通常要换 LAMB 或带自适应裁剪的优化器。

一句话:DDP = 每卡完整模型 + 反向重叠的梯度 all-reduce;它解决的是「训练慢」,不解决「模型装不下」——显存问题要交给 ZeRO/FSDP。


3. 通信原语:all-reduce 与集合通信

3.1 集合通信五件套

分布式训练的通信全部由 NCCL 提供的集合操作构成:

原语语义典型用途
broadcast一对多初始化参数
all-reduce多对多归约(求和/平均)梯度同步
reduce-scatter归约后分片ZeRO 梯度分片
all-gather收集各卡分片ZeRO 参数重建
all-to-all全交换MoE / 序列并行

DDP 用 all-reduce;ZeRO 用 reduce-scatter + all-gather 组合替代 all-reduce,这是理解 ZeRO 通信量的关键。

3.2 Ring All-Reduce 的直觉

all-reduce 的高效实现是 Ring 算法:N 张卡组成环,分两阶段——reduce-scatter(每卡负责一段的部分和,沿环传递累加)+ all-gather(把各卡拥有的完整段沿环广播)。通信量是 2 × (N-1)/N × 数据量,与卡数近似无关,这是它能扩展的核心原因。判断通信是否成为瓶颈看计算通信比:卡内有 NVLink(A100 600GB/s)时数据并行扩展性良好;仅 PCIe(32GB/s)时通信可能吃掉 30%+ 时间;跨机器走以太网/RoCE 则必须靠 bucket 重叠和张量并行缓解。

# 测量实际可用带宽(NCCL 自带测试)
./build/all_reduce_perf -b 8 -e 4G -f 2 -g 8

一句话:all-reduce 是数据并行的通信核心,Ring 算法让它与卡数弱相关;但跨机带宽远低于 NVLink,所以大模型要在张量并行(卡内高带宽)和流水线并行(低通信)之间做取舍。


4. 模型并行与张量并行

4.1 模型并行:按层切

最朴素的模型并行是把不同层放到不同卡,问题是同一时刻只有一张卡在算,其余卡空转,利用率极低。它只在「模型单卡绝对装不下」时使用,实践中已被流水线并行取代。

4.2 张量并行:按矩阵切

张量并行(TP,Megatron-LM 提出)把单个矩阵乘法切到多卡。以 Transformer 的 MLP 为例,第一个线性层按列切(Column Parallel),第二个按行切(Row Parallel):

import torch
import torch.nn as nn
import torch.distributed as dist

class ColumnParallelLinear(nn.Module):
    """按输出维度切分:每卡持有一部分输出通道"""
    def __init__(self, in_features, out_features, world_size):
        super().__init__()
        self.out_per_rank = out_features // world_size
        self.weight = nn.Parameter(
            torch.empty(self.out_per_rank, in_features)
        )

    def forward(self, x):
        # x 在每卡相同,输出是完整输出的一段
        return torch.nn.functional.linear(x, self.weight)

class RowParallelLinear(nn.Module):
    """按输入维度切分:每卡持有一部分输入,输出 all-reduce 求和"""
    def __init__(self, in_features, out_features, world_size):
        super().__init__()
        self.in_per_rank = in_features // world_size
        self.weight = nn.Parameter(
            torch.empty(out_features, self.in_per_rank)
        )

    def forward(self, x):
        partial = torch.nn.functional.linear(x, self.weight)
        dist.all_reduce(partial)  # 各卡部分和相加
        return partial

列切 + 行切成对出现,中间无需通信(每卡算自己的分片),只在 Row Parallel 后做一次 all-reduce,这是 TP 通信量小的原因。代价是每层都通信,所以 TP 通常限制在同一节点内(走 NVLink)。Multi-Head Attention 天然适合 TP:按注意力头切分,8 个头分到 4 张卡每卡算 2 个,QKV 投影按列切、输出投影按行切,结构与 MLP 一致——这也是 TP 并行度常设为注意力头数约数的原因。

一句话:张量并行是「层内切矩阵」,通信频繁但量小,必须跑在 NVLink 上;它的价值是让单层就能跨卡,从而突破单卡装不下大层的限制。


5. 流水线并行与气泡

5.1 按阶段切分与微批

流水线并行(PP)把模型按层切成若干 stage,每卡一个 stage,数据切成多个**微批(micro-batch)**依次流过形成流水线:

时间 →
卡0: [F1][F2][F3][F4]            [B4][B3][B2][B1]
卡1:      [F1][F2][F3][F4]   [B4][B3][B2][B1]
卡2:           [F1][F2][F3][F4]  ...

5.2 气泡率与调度策略

流水线的致命伤是 气泡(bubble):填充和排空阶段有卡在空转。气泡占比近似 (P-1) / (M + P - 1)(P 为 stage 数、M 为微批数),微批越多气泡越小。

调度特点适用
GPipe全前向再全反向,显存占用高简单场景
1F1B前向反向交替,显存省主流默认
Interleaved 1F1B每卡持多段,气泡更小追求效率
# PyTorch 原生流水线(torch.distributed.pipeline.sync)
from torch.distributed.pipeline.sync import Pipe

model = nn.Sequential(layer1, layer2, layer3, layer4)
model = Pipe(model, chunks=8, checkpoint="except_last")
output = model(input)  # 自动切微批、调度 1F1B

5.3 3D 并行的组合逻辑

真实的大模型训练是三种并行的乘积 总卡数 = TP × PP × DP,例如 64 卡 = 8(TP,节点内 NVLink)× 2(PP,跨节点)× 4(DP)。分配原则:TP 放在节点内(通信最密),PP 跨节点(通信最稀),DP 用剩余卡。

一句话:流水线并行用「微批 + 1F1B」把层间切分做成了流水线,气泡是它的固有成本;3D 并行的排布口诀是 TP 在内、PP 在外、DP 补足。


6. ZeRO 与 FSDP:显存账本重构

6.1 显存账本的重新分配

ZeRO(Zero Redundancy Optimizer)的洞察是:数据并行里每张卡都存了一份完整的参数、梯度、优化器状态,这是巨大的冗余。ZeRO 把它们按卡分片,用时再 all-gather 回来。

阶段切分内容显存节省(N 卡)通信量
ZeRO-1优化器状态~4x与 DDP 相同
ZeRO-2+ 梯度~8x略增
ZeRO-3+ 参数~N倍显著增加

6.2 FSDP:PyTorch 原生实现

FSDP(Fully Sharded Data Parallel)是 PyTorch 对 ZeRO-3 的原生实现,核心是参数分片 + 用时 all-gather + 用完即弃:

import torch
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import MixedPrecision, ShardingStrategy

mp_policy = MixedPrecision(
    param_dtype=torch.bfloat16,
    reduce_dtype=torch.bfloat16,
    buffer_dtype=torch.bfloat16,
)

model = FSDP(
    model,
    sharding_strategy=ShardingStrategy.FULL_SHARD,  # 对应 ZeRO-3
    mixed_precision=mp_policy,
    device_id=torch.cuda.current_device(),
    use_orig_params=True,
)

ShardingStrategy 的三个选项对应不同 ZeRO 级别:FULL_SHARD(ZeRO-3)、SHARD_GRAD_OP(ZeRO-2)、NO_SHARD(DDP)。

6.3 分片带来的通信代价

ZeRO-3 每层前向要 all-gather 参数、反向要 all-gather + reduce-scatter,通信量约为 DDP 的 1.5 倍,带宽不足时 ZeRO-3 反而比 DDP 慢。经验法则:单节点内、模型装得下用 ZeRO-1/2 或纯 DDP;跨节点、模型装不下才用 ZeRO-3 + 大 bucket + 通信重叠。

一句话:ZeRO/FSDP 用「分片换显存」,把数据并行的冗余存储彻底消除;但分片越彻底、通信越多,ZeRO-3 的收益要建立在足够带宽之上。


7. DeepSpeed 实战配置

7.1 一个可用的 ZeRO-3 配置

DeepSpeed 把分布式策略抽象成一份 JSON,改配置就能切换并行方案:

{
  "train_batch_size": 256,
  "gradient_accumulation_steps": 4,
  "fp16": {
    "enabled": true,
    "loss_scale": 0,
    "initial_scale_power": 16
  },
  "zero_optimization": {
    "stage": 3,
    "offload_optimizer": { "device": "cpu", "pin_memory": true },
    "offload_param": { "device": "cpu", "pin_memory": true },
    "overlap_comm": true,
    "contiguous_gradients": true,
    "stage3_gather_16bit_weights_on_model_save": true,
    "sub_group_size": 1e9,
    "reduce_bucket_size": "auto"
  },
  "gradient_clipping": 1.0,
  "steps_per_print": 100
}

offload_optimizer 和 offload_param 把状态/参数卸到 CPU 内存,用带宽换显存,是「单卡跑大模型」的常用手段。

7.2 启动脚本与配置切换

deepspeed --num_gpus=8 --num_nodes=2 train.py --deepspeed ds_config_zero3.json \
  --per_device_train_batch_size 8 --gradient_accumulation_steps 4 --bf16

在 HuggingFace Trainer 中,只需把 JSON 路径传进 deepspeed 参数,其余不用改代码:

from transformers import TrainingArguments

args = TrainingArguments(
    output_dir="./out",
    per_device_train_batch_size=8,
    gradient_accumulation_steps=4,
    bf16=True,
    deepspeed="ds_config_zero3.json",
    gradient_checkpointing=True,
)

7.3 三种并行在 DeepSpeed 中的组合

DeepSpeed 通过 zero_optimization.stage 控制数据并行侧,通过外部参数控制 TP/PP:

deepspeed --num_gpus=64 train.py --deepspeed ds_config.json \
  --tensor_model_parallel_size 8 --pipeline_model_parallel_size 2 --zero_stage 1

注意:用了 TP/PP 后 ZeRO 一般只开到 stage 1,因为参数已经切分,再叠 stage 3 会重复分片、通信爆炸。

一句话:DeepSpeed 的价值是把「并行策略」变成配置文件,从 ZeRO-1 到 3D 并行靠改 JSON 就能切换;但配置之间不是叠加越多越好,TP/PP 与 ZeRO-3 通常互斥。


8. 梯度累积、混合精度与排错手册

8.1 梯度累积:小显存换大 batch

当显存开不大 batch 时,用梯度累积模拟大 batch:

accum_steps = 4
for i, batch in enumerate(loader):
    loss = model(**batch) / accum_steps
    loss.backward()
    if (i + 1) % accum_steps == 0:
        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
        optimizer.step()
        optimizer.zero_grad(set_to_none=True)

三个坑:除以 accum_steps 做平均(否则梯度放大)、只在真正 step 时裁剪梯度、与 DDP 的 no_sync 配合避免每个微批都同步(微批内不同步,最后一个微批才同步):

with model.no_sync():        # 微批内不同步梯度
    loss.backward()
loss.backward()              # 最后一个微批才同步

8.2 混合精度:fp16 vs bf16

精度动态范围是否需 loss scaling硬件要求
fp16窄,易溢出需要V100+
bf16宽,与 fp32 同不需要A100+

优先 bf16:省掉 loss scaling 的调试烦恼,数值更稳;fp16 则必须配 GradScaler。

8.3 多机排错清单

现象根因处理
卡在 init_process_group网络/端口不通检查 master_addr、防火墙、NCCL_IB
NCCL timeout某卡慢或掉队调大 NCCL_TIMEOUT,定位慢卡
显存 OOM 但单卡够ZeRO 分片不均检查 batch 能否整除、调 sub_group_size
loss 不降 / 变 NaN学习率未缩放、loss scale 溢出加 warmup、切 bf16、降 lr
多机比单机还慢通信未重叠开 overlap_comm、调大 bucket
结果不可复现随机种子未按 rank 设置seed = base_seed + rank

常用环境变量:

export NCCL_DEBUG=INFO           # 打印通信细节
export NCCL_IB_DISABLE=1         # 无 InfiniBand 时禁用
export NCCL_SOCKET_IFNAME=eth0   # 指定网卡

8.4 一个诊断流程

1. 单卡能否跑通?→ 不能则先修模型代码
2. 单机 8 卡能否跑通?→ 不能则查 NCCL/拓扑
3. 多机能否跑通?→ 查网络与时钟同步
4. 加速比是否线性?→ 不线性则 profile 通信占比;收敛是否与单卡一致?对比前 100 步 loss

一句话:梯度累积、混合精度是「单卡榨显存」的工具,多机排错的顺序永远是「单卡 → 单机 → 多机」,逐层排除。


9. 总结

9.1 策略选择路线

模型单卡装得下 + 想更快       → DDP 数据并行
模型单卡装不下 + 有 NVLink    → 张量并行(节点内)
模型很多层 + 想省通信         → 流水线并行(跨节点)
模型装不下 + 无高带宽         → ZeRO-1/2 + CPU offload
千亿参数 + 大规模集群         → 3D 并行(TP×PP×DP)+ ZeRO-1

9.2 关键决策点

问题选择
只是训练慢,模型装得下DDP,别引入复杂度
单层就装不下张量并行,限制在节点内
层数多、跨机通信贵流水线并行 + 多微批
显存紧张、带宽尚可ZeRO-2
显存极紧、带宽充足ZeRO-3 + offload
MoE 稀疏模型专家并行 + all-to-all

9.3 一句话心法

分布式训练没有银弹,只有「显存—通信—复杂度」的三角权衡——先问清瓶颈是显存还是时间,再选切分方式;切分越细、通信越贵,能单卡解决就别上多卡。


延伸阅读

  • https://plumephp.com/ml-deep-learning-advanced/ — 单卡训练循环、优化器与正则化基础
  • https://plumephp.com/ml-gradient-descent-optimizers/ — 优化器状态与显存占用的原理
  • https://plumephp.com/ml-model-compression-quantization/ — 量化与压缩如何减轻显存压力
  • https://plumephp.com/ml-model-deployment/ — 训练完成后的推理服务与版本管理
  • AI/ML 专题 — 大规模训练与系统工程深度文章
  • PyTorch 分布式文档
  • DeepSpeed 官方教程

继续阅读

探索更多技术文章

浏览归档,发现更多关于系统设计、工具链和工程实践的内容。

全部文章 返回首页

「ml」更多文章

  1. 联邦学习与隐私保护机器学习:FedAvg、非 IID、DP-SGD 与安全聚合
  2. 语音与音频机器学习:MFCC、CTC/RNN-T、Whisper、TTS 与声码器
  3. 扩散模型与生成式建模:DDPM、U-Net、潜在扩散与 LoRA 微调实战