跳转至

第十二章 分布式训练:并行策略、通信与故障恢复

分布式训练的本质是把参数、数据、激活、梯度和优化器状态分配到多个设备,并让通信与计算协调。目标不是“用更多卡”,而是在内存上可容纳、数值上等价、吞吐上有效、故障后可恢复。

分布式并行策略

图 12-1 数据并行复制模型;张量/流水线/专家并行切分模型计算;ZeRO/FSDP 切分训练状态。

12.1 DP 与 DDP

先建立直觉。 DataParallel 在单进程聚合,DDP 为每个进程维护模型副本并用 All-Reduce 同步梯度。DDP 通常性能和隔离更好,但每张卡仍保存完整权重与优化器。 读这一节时,先不要急着记缩写,而要不断追问:输入是什么,经过了哪些可观察变换,输出又怎样被验证。

把过程拆开看:

  1. 每个 rank 初始化进程组。
  2. DistributedSampler 切分数据。
  3. 前向反向触发梯度 bucket 通信。
  4. 所有 rank 一致更新参数。

最小例子。 4 卡 DDP 若每卡 micro batch=8、累积 2 步,有效 batch 通常为 64,但最后不齐批次和 sampler 设置会影响样本数。 例子越小,越容易手算、打印中间量并判断实现是否偏离定义。

容易踩坑。 每个 rank 读取相同样本、只在 rank0 backward、随机种子完全相同导致增强重复、保存冲突,都会出错。 排错时应保存失败输入和版本,先定位最早出现异常的步骤,再决定是否调整模型或超参数。

动手任务:运行两卡最小 DDP,打印每个 rank 样本 id 并验证无重复覆盖。 完成后不要只保存最终输出,还要保存代码、配置、中间张量或检索结果,以及你对结果的解释。

数据并行让每张卡持有模型副本,处理不同微批次,反向后同步梯度。单进程 DataParallel 有主卡瓶颈,实践优先一进程一卡的 DistributedDataParallel(DDP)。DDP 通过 All-Reduce 聚合梯度,所有 rank 以相同梯度和优化器状态更新,因此参数保持一致。

# torchrun --standalone --nproc_per_node=4 train_ddp.py
import os
import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader, DistributedSampler, TensorDataset

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

    x = torch.randn(4096, 32)
    y = torch.randn(4096, 1)
    dataset = TensorDataset(x, y)
    sampler = DistributedSampler(dataset, shuffle=True)
    loader = DataLoader(dataset, batch_size=64, sampler=sampler, pin_memory=True)

    model = DDP(torch.nn.Linear(32, 1).cuda(local_rank), device_ids=[local_rank])
    opt = torch.optim.AdamW(model.parameters(), lr=1e-3)
    for epoch in range(3):
        sampler.set_epoch(epoch)
        for xb, yb in loader:
            xb, yb = xb.cuda(local_rank, non_blocking=True), yb.cuda(local_rank, non_blocking=True)
            opt.zero_grad(set_to_none=True)
            loss = torch.nn.functional.mse_loss(model(xb), yb)
            loss.backward()
            opt.step()
    if dist.get_rank() == 0:
        torch.save(model.module.state_dict(), "model.pt")
    dist.destroy_process_group()

if __name__ == "__main__":
    main()

DistributedSampler 防止所有 rank 读取相同样本;每轮 set_epoch 改变一致的 shuffle;只有 rank 0 写普通 checkpoint。真实代码还需全局指标 All-Reduce、异常协调和恢复逻辑。

12.2 ZeRO 与 FSDP

从问题出发。 ZeRO/FSDP 将优化器状态、梯度乃至参数分片到各 rank,降低单卡冗余。节省显存的代价是更多集合通信、参数聚合和更复杂 checkpoint。 先把名词放到一边,沿着输入、状态变化和输出走一遍;只有能指出证据落在哪一步,概念才算真正掌握。

沿数据流逐步检查:

  1. Stage1 分优化器状态。
  2. Stage2 再分梯度。
  3. Stage3/FSDP 再分参数。
  4. 计算前按需聚合、计算后释放或重分片。

用小数据走一遍。 参数分片并不意味着前向永远只看本地参数;计算某层前往往要 All-Gather 得到完整层。 先在纸面预测结果,再让程序打印中间状态;预测与运行不一致的地方,正是需要继续追查的知识缺口。

这里最容易出现的误解。 把 ZeRO stage 当单纯开关、wrap 粒度太细、CPU offload 受 PCIe 限制、保存 full state 时 OOM,都会影响训练。 出现异常时先冻结样本、配置和环境,从最靠近输入的环节开始验证;不要同时更换模型、数据和超参数。

小实验:为一个模型估算 DDP 与 ZeRO-1/2/3 的理论状态显存,并说明遗漏项。 提交物应包含预期、运行证据和失败记录;只有别人能按同样步骤复现,结果才具有学习价值。

Adam 混合精度训练中,参数、梯度和优化器状态占大量显存。ZeRO 分阶段切分:Stage 1 切优化器状态;Stage 2 再切梯度;Stage 3 再切参数,并在前向/反向需要时 All-Gather。FSDP 同样围绕参数分片、按需聚合和梯度 Reduce-Scatter。

Stage 越高显存越省,但通信和实现复杂度增加。Offload 把状态移到 CPU/NVMe,可进一步省显存,却可能受 PCIe、CPU 和磁盘瓶颈限制。

{
  "bf16": {"enabled": true},
  "gradient_accumulation_steps": 8,
  "train_micro_batch_size_per_gpu": 1,
  "zero_optimization": {
    "stage": 3,
    "overlap_comm": true,
    "contiguous_gradients": true,
    "reduce_bucket_size": 50000000,
    "stage3_prefetch_bucket_size": 50000000
  }
}

配置中的 batch 三元关系必须一致。桶大小影响通信聚合与峰值内存,不应盲目复制;用 profiler 测量。

12.3 模型并行

并行策略的切分轴

图 12-2 数据、张量、层和专家是四种不同切分维度,通信模式也随之不同。 先看它解决什么。 Tensor Parallel 切单层矩阵,Pipeline Parallel 切层,Sequence/Context Parallel 切序列,Expert Parallel 切专家。选择取决于模型哪一维无法放入设备和通信拓扑。 理解的标准不是会复述定义,而是能画出数据流、预测中间结果,并设计一个让错误暴露出来的检查。

可以把实现分成以下环节:

  1. 识别最大单层与总状态。
  2. 选择切分轴。
  3. 安排集合通信与流水线。
  4. 再叠加数据并行形成多维网格。

一个可以手算的例子。 列并行线性层把输出列分到多卡,后续操作若需要完整输出就 All-Gather;行并行常需要 Reduce-Scatter 或 All-Reduce。 这里故意不用大模型或大数据,因为小输入能把每个轴、分数和状态变化完整暴露出来。

排错时先看这些地方。 只乘并行度不考虑整除、跨节点放高频通信、流水线 microbatch 太少导致 bubble,都会降低效率。 修复不能止于‘这次跑通’,还要把失败样本变成自动测试,防止同类问题在下一版本重新出现。

现在动手:为 8 卡单机和 2×8 卡集群分别设计并行网格并解释通信。 把关键断言写进测试,并在 README 说明怎样运行、怎样判断正确、当前实现还不支持什么。

张量并行(TP)把单个矩阵乘按列或行切到多卡,需要在层内频繁集合通信;适合高速互连。流水线并行(PP)把层分阶段,微批次流水执行;会有 pipeline bubble,需调度与微批数量平衡。序列并行(SP)在序列维切分部分计算,长上下文下有价值。专家并行(EP)把 MoE 专家分布到设备,token 通过 All-to-All 路由。

混合并行通常组合 DP×TP×PP×EP。映射要尊重硬件拓扑:把通信最频繁的 TP 放在节点内高速连接,把数据并行扩到节点间。并行维度乘积应等于总设备数。

12.4 集合通信

抓住这一节的主线。 集合通信是并行算法的语言:All-Reduce 汇总并复制结果,Reduce-Scatter 汇总后分片,All-Gather 收集分片,All-to-All 重新分发不同数据。 遇到新术语,先找它在系统中的位置:它读取什么、保存什么、改变什么,以及失败时会留下什么信号。

真正动手时按这个顺序走:

  1. 明确每个 rank 初始张量。
  2. 写出通信后每个 rank 所有内容。
  3. 计算传输量与次数。
  4. 匹配 DDP、FSDP、TP 或 EP。

先做最小实验。 DDP 梯度 bucket 常用 All-Reduce;FSDP 梯度可用 Reduce-Scatter;MoE token 路由常用 All-to-All。 把随机性固定并保留中间量,这个例子就能成为后续优化时的回归基线。

别被表面现象带偏。 把带宽和延迟混为一谈、通信与计算无法重叠、张量大小不均、某 rank 先退出造成死锁,都会损害稳定性。 先区分定义错误、实现错误、数据问题和性能瓶颈。四类问题需要的证据不同,不能靠盲目调参混在一起处理。

验证任务:用四个 rank 的小向量手工演示四种集合通信结果。 实验结束后写三句话:观察到了什么、这些证据支持什么结论、还有哪种解释尚未排除。

常见原语:All-Reduce 汇总并复制结果;Reduce-Scatter 汇总后分片;All-Gather 收集所有分片;All-to-All 交换不同目标的数据。通信时间近似由延迟和带宽共同决定,小消息受延迟支配,大消息受带宽支配。

DDP 用梯度 bucket 尝试让反向计算与 All-Reduce 重叠。若存在未使用参数、控制流不一致或不同 rank 执行不同步,可能死锁。所有 rank 必须以一致顺序进入集合通信。

12.5 混合精度

先把概念落到可观察对象上。 FP16 范围小,常需 loss scaling;BF16 指数范围接近 FP32,通常更稳但尾数精度低。混合精度把敏感运算保留更高精度,并使用适配 kernel。 这一节的重点是建立因果链。先说明问题,再看计算或流程,最后用可重复实验验证结论。

把抽象概念还原成操作:

  1. 选择 autocast dtype。
  2. loss scaling 后反传。
  3. 更新前反缩放与裁剪。
  4. 检查非有限梯度并调整 scale。

把它缩小到能逐项检查。 Softmax、归一化统计和优化器状态常需更高精度;具体由框架与 kernel 实现决定。 如果这个小例子还不能解释清楚,扩大数据只会让错误更难发现。

需要特别守住的边界。 仅把模型 .half、裁剪缩放前梯度、硬件不支持 BF16、不同 rank 出现 Inf 却继续同步,都会导致错误。 保留完整 trace 比猜原因更重要。找到第一次偏离预期的位置,通常比分析最终错误输出更高效。

本节练习:比较 FP32、FP16+scaler、BF16 的损失曲线、显存和吞吐。 除代码外,请保留一份结果说明,标明环境、随机种子、输入规模和你主动检查过的边界。

FP16 指数范围小,通常需 loss scaling;BF16 指数范围接近 FP32,训练更稳但尾数精度较低。部分归一化、Softmax、损失和统计仍常在 FP32 计算。混合精度是否更快取决于硬件与 kernel。

梯度累积时,若用 DDP,可在非最后一个微步使用 no_sync() 避免每次同步。梯度裁剪应在反缩放之后、优化器更新之前执行。

12.6 Checkpoint 与断点续训

先建立直觉。 真正可续训的 checkpoint 包含模型、优化器、scheduler、scaler、随机状态、数据位置和并行拓扑信息。分片 checkpoint 还要能适配 world size 或有转换流程。 读这一节时,先不要急着记缩写,而要不断追问:输入是什么,经过了哪些可观察变换,输出又怎样被验证。

把过程拆开看:

  1. 在一致 step 建立保存屏障。
  2. 写临时目录并校验完整。
  3. 原子发布成功标记。
  4. 恢复后做短程连续性验证。

最小例子。 只恢复权重会丢失 Adam 动量与学习率进度,loss 可能突变;数据迭代位置丢失则重复或跳过样本。 例子越小,越容易手算、打印中间量并判断实现是否偏离定义。

容易踩坑。 所有 rank 写同一文件、成功标记先于分片完成、恢复后重置 seed、从不演练损坏分片,都会让备份失效。 排错时应保存失败输入和版本,先定位最早出现异常的步骤,再决定是否调整模型或超参数。

动手任务:训练 20 步在第 10 步保存,比较不中断与恢复路径第 11—20 步的指标。 完成后不要只保存最终输出,还要保存代码、配置、中间张量或检索结果,以及你对结果的解释。

分片训练的 checkpoint 可能包含每 rank 的分片、元数据和优化器状态。恢复时世界大小变化并非所有格式都支持。保存内容包括模型、优化器、调度器、随机数、数据位置、全局步和 scaler。

采用临时目录写完后原子发布,保留完成标记和校验和;定期在独立作业中恢复并跑若干步。只测试“能加载模型权重”不足以证明能继续训练。

12.7 Accelerate 与 DeepSpeed

从问题出发。 Accelerate 提供较薄的设备与分布式抽象,DeepSpeed 提供 ZeRO、offload 与训练引擎。工具减少配置工作,但不能替代对 batch、状态分片和通信的理解。 先把名词放到一边,沿着输入、状态变化和输出走一遍;只有能指出证据落在哪一步,概念才算真正掌握。

沿数据流逐步检查:

  1. 先在单卡验证数据和 loss。
  2. 用配置文件声明精度与并行。
  3. 检查启动后的实际 world size。
  4. 保存并恢复完整状态。

用小数据走一遍。 同一 YAML 在库版本变化后默认行为可能改变,运行日志应打印最终解析配置,而不只保存输入文件。 先在纸面预测结果,再让程序打印中间状态;预测与运行不一致的地方,正是需要继续追查的知识缺口。

这里最容易出现的误解。 同时让多个框架接管梯度累积、配置键拼错被忽略、只在 rank0 初始化不一致对象,都会产生隐蔽 bug。 出现异常时先冻结样本、配置和环境,从最靠近输入的环节开始验证;不要同时更换模型、数据和超参数。

小实验:用 Accelerate 将单卡脚本改为多卡,并逐项说明代码变化与未变化部分。 提交物应包含预期、运行证据和失败记录;只有别人能按同样步骤复现,结果才具有学习价值。

Accelerate 统一设备放置、混合精度和多进程启动,适合把单卡脚本平滑扩展。DeepSpeed 提供 ZeRO、流水线、优化器与推理能力。框架简化入口,但不会替你决定正确并行策略,也不会自动修复数据或通信瓶颈。

调试顺序:单卡小数据过拟合;单机两卡确认数值;扩大卡数并比较全局 batch;再加 ZeRO/混合精度;最后做多机与故障恢复。一次引入所有优化会让问题难以定位。

12.8 性能诊断

先看它解决什么。 性能诊断从时间线入手:数据等待、前向、反向、通信、优化器和 checkpoint 各占多少。GPU 利用率低只是现象,不直接告诉根因。 理解的标准不是会复述定义,而是能画出数据流、预测中间结果,并设计一个让错误暴露出来的检查。

可以把实现分成以下环节:

  1. 用 profiler 捕获稳定窗口。
  2. 查看 kernel 空洞与通信重叠。
  3. 统计数据加载和 CPU。
  4. 一次只改变一个瓶颈。

一个可以手算的例子。 所有 rank 在 All-Reduce 前等待同一个慢 rank,可能来自样本长度不均或硬件降频,而不是 NCCL 本身。 这里故意不用大模型或大数据,因为小输入能把每个轴、分数和状态变化完整暴露出来。

排错时先看这些地方。 只看 nvidia-smi 瞬时利用率、profile 包含预热、用更大 batch 掩盖数据错误、跨节点未绑定网卡,都会误诊。 修复不能止于‘这次跑通’,还要把失败样本变成自动测试,防止同类问题在下一版本重新出现。

现在动手:对一次训练 step 做时间分解,提出证据支持的三项优化并复测。 把关键断言写进测试,并在 README 说明怎样运行、怎样判断正确、当前实现还不支持什么。

记录每步时间、数据时间、前向/反向、优化器、通信、MFU、网络带宽和显存。卡利用率低可能是数据慢、通信慢、微批过小、CPU 同步或频繁 checkpoint。某 rank 变慢会拖住所有同步 rank,需排查慢卡和数据倾斜。

12.9 练习与资料

抓住这一节的主线。 练习分布式不能只追求跑通,还要验证数值等价、样本覆盖、故障恢复和性能缩放。规模扩大前先在小集群注入错误。 遇到新术语,先找它在系统中的位置:它读取什么、保存什么、改变什么,以及失败时会留下什么信号。

真正动手时按这个顺序走:

  1. 单卡与多卡 loss 对齐。
  2. 检查全局 batch 样本。
  3. kill 一个进程观察行为。
  4. 恢复后比较状态。

先做最小实验。 固定有效 batch 与随机性后,DDP 结果应与单卡在容差内接近;完全逐位相同通常不现实。 把随机性固定并保留中间量,这个例子就能成为后续优化时的回归基线。

别被表面现象带偏。 多卡更快就认为正确、没有 barrier 超时、异常进程未清理、恢复只看能启动,都会遗漏。 先区分定义错误、实现错误、数据问题和性能瓶颈。四类问题需要的证据不同,不能靠盲目调参混在一起处理。

验证任务:完成一份 DDP/FSDP 验收清单,包含正确性、吞吐、显存和恢复。 实验结束后写三句话:观察到了什么、这些证据支持什么结论、还有哪种解释尚未排除。

练习:用 1/2/4 卡 DDP 保持全局 batch 不变并比较 loss;估算 ZeRO-1/2/3 各自切分哪些状态;画出 8 卡上 DP=2、TP=2、PP=2 的 rank 分组。

延伸阅读:PyTorch DDPPyTorch FSDPDeepSpeed ZeRO 教程

本章配套代码

下面的脚本与正文使用相同符号。建议先在代码中打印形状和中间量,再运行断言;若依赖尚未安装,至少先阅读入口函数、输入输出和测试部分。

本章端到端实验:把知识变成可复现证据

本实验不是把本章代码重新抄一遍,而是把概念、实现、测试和解释串成一个小型工程。请新建独立目录,保存 README.md、环境文件、源代码、测试、运行日志和结果图。README 至少说明任务、输入输出、运行命令、预期现象、已知限制和复现条件。

实验步骤

  1. 运行两卡最小 DDP,打印每个 rank 样本 id 并验证无重复覆盖。
  2. 为一个模型估算 DDP 与 ZeRO-1/2/3 的理论状态显存,并说明遗漏项。
  3. 为 8 卡单机和 2×8 卡集群分别设计并行网格并解释通信。
  4. 用四个 rank 的小向量手工演示四种集合通信结果。
  5. 比较 FP32、FP16+scaler、BF16 的损失曲线、显存和吞吐。
  6. 训练 20 步在第 10 步保存,比较不中断与恢复路径第 11—20 步的指标。

每完成一步,先写下预期,再运行代码。若结果与预期不一致,不要覆盖旧日志;建立 failures.md,记录现象、假设、证据、修复和回归测试。这样得到的不是一次性 Demo,而是一份能证明你真正理解本章内容的实验档案。

验收标准

  • 全新环境能够按照 README 从头运行,依赖和随机种子已记录。
  • 关键函数至少有正常、边界和错误输入三类测试;涉及数值计算时检查有限值与合理误差。
  • 结果包含一个基线和至少一个受控改动,能够说明变化来自哪里。
  • 日志保留输入规模、耗时、内存或显存、软件版本和失败样本。
  • 结论区分“实验直接证明的事实”“根据事实做出的推断”和“仍未验证的猜想”。

本章自测

  1. 不看正文,用自己的话解释“DataParallel 在单进程聚合,DDP 为每个进程维护模型副本并用 All-Reduce 同步梯度”,并给出一个可以证伪的测试。
  2. 不看正文,用自己的话解释“ZeRO/FSDP 将优化器状态、梯度乃至参数分片到各 rank,降低单卡冗余”,并给出一个可以证伪的测试。
  3. 不看正文,用自己的话解释“Tensor Parallel 切单层矩阵,Pipeline Parallel 切层,Sequence/Context Parallel 切序列,Expert Parallel 切专家”,并给出一个可以证伪的测试。
  4. 不看正文,用自己的话解释“集合通信是并行算法的语言:All-Reduce 汇总并复制结果,Reduce-Scatter 汇总后分片,All-Gather 收集分片,All-to-All 重新分发不同数据”,并给出一个可以证伪的测试。
  5. 不看正文,用自己的话解释“FP16 范围小,常需 loss scaling”,并给出一个可以证伪的测试。
  6. 不看正文,用自己的话解释“真正可续训的 checkpoint 包含模型、优化器、scheduler、scaler、随机状态、数据位置和并行拓扑信息”,并给出一个可以证伪的测试。

回答时先画图或写形状,再给结论。若只能说出术语而不能给出最小例子、边界条件和验证方法,说明这一节仍需要回到代码中重做。