分布式训练实战:DDP、FSDP、ZeRO 与大规模训练工程

大模型训练离不开多卡分布式,本文系统讲解 DDP 数据并行原理与 AllReduce、混合精度训练(FP16/BF16)、梯度累积与梯度裁剪、ZeRO 分片(ZeRO-1/2/3)、FSDP 全分片数据并行、张量并行与流水线并行、checkpoint 与断点续训、以及分布式训练排错(NCCL 超时/OOM/收敛异常)的完整工程方法论。

单卡放不下的模型,只能用分布式训练。但分布式不是「多卡各跑一份再求和」那么简单——通信开销、显存分片与收敛一致性,每一步都是工程陷阱。本文从 DDP 原理讲起,给出可落地的多卡训练方案。

为什么需要分布式训练

三个朴素理由推动分布式训练:模型放不下、数据跑不动、时间等不起。

  • 模型放不下:70B 参数在 BF16 下约 140GB,远超单卡 80GB,必须把权重切分到多卡。
  • 数据跑不动:单卡一次只能处理一个小 batch,训练大数据集太慢,需要并行数据流。
  • 时间等不起:同样的迭代轮数,10 卡理论上能比单卡快接近 10 倍(理想线性加速)。

但分布式是有成本的:通信开销会吃掉一部分加速比。卡越多,AllReduce 的通信量越大,数据并行在 8~16 卡后收益递减,需要切换到分片与并行混合策略。工程上的核心是「在显存、通信与算力之间找平衡」。

DDP:数据并行的黄金标准

**DDP(DistributedDataParallel)**是最经典、最稳妥的数据并行实现,PyTorch 官方推荐。核心思想:每张卡持有一份完整模型副本,各自处理不同的数据分片,前向/反向各自独立,最后通过 AllReduce 同步梯度。

# DDP 最小启动示例
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel

dist.init_process_group("nccl")          # 初始化进程组
model = DistributedDataParallel(model)   # 包装模型

# 训练循环与单卡几乎一致,loss.backward() 后 DDP 自动同步梯度

DDP 的两个关键工程点:

  • AllReduce 的通信量:每步训练都要把全部梯度做一次全局归约,梯度总量 = 参数量 × 2 字节。7B 模型每步约同步 14GB 数据,通信成为瓶颈。
  • 梯度按桶(Bucket)通信:DDP 把参数按大小分桶,桶内梯度凑齐后异步发起通信,降低通信延迟。桶越小通信越频繁,桶越大内存峰值越高,默认桶大小(25MB)通常已够用。

DDP 适合模型单卡放得下、只是数据太多的场景;模型超过单卡显存时,就要靠下面的分片方案。

混合精度训练:FP16 与 BF16

混合精度是每个训练任务的标配:主权重用 FP32 保精度,前向反向用 FP16/BF16 提速省显存。PyTorch 的 torch.autocast 与 GradScaler 组合即可。

# 混合精度训练(AMP)标准写法
scaler = torch.cuda.amp.GradScaler()   # FP16 需要梯度缩放

for data, label in loader:
    with torch.autocast(device_type="cuda", dtype=torch.float16):
        loss = model(data)                # 前向在 FP16 下计算
    scaler.scale(loss).backward()         # 缩放梯度防下溢
    scaler.step(optimizer)
    scaler.update()
  • FP16:速度与显存最优,但动态范围小,小梯度易下溢,需要 GradScaler 动态缩放。
  • BF16:动态范围与 FP32 相同,无需缩放,训练更稳,但占用与 FP16 相同、速度略慢。大模型预训练几乎都用 BF16。

一句话选型:推理用 FP16/INT8,训练用 BF16。若显存仍不够,再叠加下面的分片方案。

梯度累积:小显存跑大 batch

显存不够放大 batch 时,**梯度累积(Gradient Accumulation)**用时间换空间:多个小 batch 反向后不立即更新,累加梯度,累积到目标步数再 optimizer.step()。

# 梯度累积:4 个小步合成一个大步
accum_steps = 4
optimizer.zero_grad()
for i, (data, label) in enumerate(loader):
    loss = model(data)
    (loss / accum_steps).backward()       # 除累积步数防梯度爆炸
    if (i + 1) % accum_steps == 0:
        optimizer.step()
        optimizer.zero_grad()

注意三个坑:

  • BN 失效:累积后 batch 统计量变小,BatchNorm 的均值和方差漂移——大模型多用 LayerNorm/GroupNorm 规避此问题。
  • 梯度爆炸:累积梯度会越积越大,必须配梯度裁剪 torch.nn.utils.clip_grad_norm_。
  • 学习率调度:LR scheduler 应在大步(step)上更新,而非每个小步。

ZeRO:把冗余全部切掉

ZeRO(Zero Redundancy Optimizer)的洞察是:DDP 每张卡都存了一份完整权重、梯度、优化器状态,这三份是冗余的。ZeRO 把这三份状态分片到各卡,分为三档:

  • ZeRO-1:分片优化器状态(如 Adam 的 momentum/variance)。显存省约一半,通信几乎不变,是性价比最高的一档。
  • ZeRO-2:再分片梯度。省得更多,但要额外一次通信归约。
  • ZeRO-3:连参数也分片,前向反向时需要实时收集参数(All-Gather)。可训练超出单卡内存数倍的模型,但通信开销显著上升。
# DeepSpeed 启用 ZeRO-2 的最小配置(deepspeed_config.json 片段)
{
  "zero_optimization": {
    "stage": 2,
    "allgather_partitions": true,
    "overlap_comm": true,
    "reduce_scatter": true
  },
  "train_batch_size": 32,
  "fp16": { "enabled": true }
}

选档建议:先 ZeRO-1 拿最大性价比;单卡仍放不下模型再上 ZeRO-2;只有权重大到必须切参数量时(如 30B+ 预训练)才用 ZeRO-3。

FSDP:PyTorch 原生的全分片

FSDP(Fully Sharded Data Parallel)是 PyTorch 官方的 ZeRO-3 等价实现,与 DeepSpeed 相比 API 更亲民、与 HuggingFace 生态集成更顺。核心差异在于 FSDP 以模型层为分片粒度:每层参数单独分片,前向时按层 All-Gather,算完即释放,显存峰值大幅下降。

# FSDP 包装模型(HuggingFace Trainer 也内置 support)
from torch.distributed.fsdp import (
    FullyShardedDataParallel as FSDP,
    CPUOffload,
    MixedPrecision,
)

model = FSDP(
    model,
    mixed_precision=MixedPrecision(
        param_dtype=torch.bfloat16, reduce_dtype=torch.bfloat16
    ),
    cpu_offload=CPUOffload(offload_params=True),  # 可选的 CPU offload
)

FSDP 的关键调参是 auto_wrap_policy:按 Transformer 层(transformer_auto_wrap_policy)分片粒度适中、通信最少。开启 CPU offload 能显著压显存,但 CPU-GPU 拷贝会成为新瓶颈,吞吐会掉。

张量并行与流水线并行

数据并行与分片解决「显存」,但单卡放不下单层(如超大 Embedding、MoE 专家)时,需要把单个层切开:

  • 张量并行(Tensor Parallel):把矩阵按行/列切到多卡,前向时用 AllReduce 聚合部分和。Megatron 的核心技术,适合超大单层,但通信密集、需要同机 NVLink。
  • 流水线并行(Pipeline Parallel):把模型按层切成几段,每卡负责一段,数据像流水线一样流过。减少峰值显存,但会产生气泡(bubble)——微批次(micro-batch)切得越小气泡越少。
# 三种并行叠加的经典分工(Megatron-LM 风格)
# 节点间: 流水线并行(每节点一段)
# 节点内: 张量并行(同机 8 卡切开单层)
# 数据并行: 在 (TP×PP) 组之间复制副本

实践中 90% 的团队用不到 TP/PP——它们是为千亿参数预训练准备的。7B~70B 微调用 FSDP/ZeRO 就够了,只有基座预训练或超大模型才引入 TP/PP。

Checkpoint 与断点续训

分布式训练动辄数天,任何一次 OOM 或断网都可能重来。断点续训是生产级训练的必备能力,核心是三样东西一起存:

# HuggingFace Trainer 自动保存(save_steps / save_total_limit 控制轮数与数量)
# 保存内容:model + optimizer state + scheduler state + random RNG state
# 续训:Trainer.train(resume_from_checkpoint=args.resume_from_checkpoint)
  • 模型权重:当前参数。
  • 优化器状态:Adam 的 momentum/variance——丢了就等于从头开始,这是续训最关键的一环。
  • 随机状态(RNG):数据顺序与 dropout 的随机种子,保证续训后数据流与原始计划一致。

工程实践:每 N 步自动保存、保留最近 K 份滚动覆盖;长任务先验证「保存→恢复→梯度对齐」再正式开跑;OOM 时优先尝试减少 batch 而非重训。

分布式训练排错手册

多卡训练的错误千奇百怪,最常见的三类:

# 1) NCCL 超时:卡间通信卡死
# 现象: "NCCL error: timeout in collective operation"
# 排查: 检查同机 NVLink/IB 是否连通、防火墙是否放行 29500-29509 端口、
#       os.environ 是否设置了 MASTER_ADDR/MASTER_PORT、卡间是否同构

# 2) CUDA OOM:显存超限
# 排查: 减 batch → 关梯度累积 → 开 ZeRO/FSDP → 开 CPU offload,逐级降
# 3) loss = NaN 或发散
# 排查: 降低 LR → 开梯度裁剪 → 检查输入是否含 NaN/Inf → 确认 loss 是否正确归一

一条黄金建议:先用 2 卡把训练跑到「能稳定收敛」,再水平扩到全部卡。分布式 bug 与训练 bug 混在一起时极难定位,先把非分布式问题清零。

总结

分布式训练的工程主线是「先数据并行拿加速,再分片省显存,最后才上并行切割」:DDP 解决数据太多,混合精度与梯度累积降低单卡压力,ZeRO/FSDP 解决模型放不下,TP/PP 服务千亿预训练。配合断点续训与系统化排错,多卡训练才能从「能跑」走到「稳跑」。记住:通信是分布式唯一的税,一切方案都是在给这笔税找最优缴法。

继续阅读

探索更多技术文章

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

全部文章 返回首页

「ai」更多文章

  1. 类别不平衡与异常检测:从重采样到半监督方法
  2. Embedding 深入:对比学习、双塔架构与向量检索工程
  3. MLOps 治理与可复现:模型注册、漂移监控与合规