二十一:分布式训练:DDP、DP 与 FSDP¶
来源:http://mp.weixin.qq.com/s?__biz=MzYyNTk3Njg1NA==&mid=2247484196&idx=1&sn=ef705c676e22283d49089e9f5424625f&chksm=f01eb05dc769394b2f6afa98eb688347979cf941bac796a7d2e15abec57682745b8f756124a2#rd
1. 学习定位¶
分布式训练用于解决单卡显存不足、训练速度慢、数据规模大和模型规模大的问题。第 21 天重点是 DDP。DDP 是 PyTorch 中最常用的数据并行训练方式,也是理解 DeepSpeed、FSDP、ZeRO、Megatron 等更复杂训练框架的基础。 本日知识链路:
单卡训练
-> 数据并行 Data Parallel
-> PyTorch DataParallel 的瓶颈
-> DistributedDataParallel
-> process group / rank / world_size / backend
-> 每卡一个进程
-> forward/backward 本地计算
-> gradient all-reduce 同步
-> optimizer step 各卡一致更新
-> DistributedSampler 保证数据切分
-> DDP 与 FSDP/ZeRO 的区别
2. 分布式训练的基本目标¶
分布式训练主要解决三类问题: - 速度:多 GPU 并行计算,提高吞吐。
-
显存:把模型、梯度、优化器状态或激活拆到多个设备。
-
规模:训练更大模型、更大 batch、更大数据集。
常见并行方式: 并行方式 拆分对象 典型方案 核心特点 数据并行 训练批次(batch)数据 DP、DDP 每张 GPU 保留完整模型副本,仅切分数据 参数 / 状态分片 模型参数、梯度、优化器状态 FSDP、ZeRO 跨卡拆分权重与优化器状态,极致省显存 张量并行 单算子 / 矩阵运算维度 Megatron-LM TP 拆分单层网络的张量,多用于大模型单层扩容 流水线并行 神经网络层结构 Pipeline Parallel 按网络层切分,多卡分段执行前向 / 反向 专家并行 MoE 模型的专家模块 Expert Parallel 专门拆分混合专家模型的各个 Expert
DDP 属于数据并行。它默认每个 GPU 拥有一份完整模型副本。
3. 数据并行的核心思想¶
数据并行把一个大 batch 切成多个 micro-batch,每个 GPU 处理其中一份。
每张卡: - 所有 GPU模型参数完全一致-
各卡加载独立数据分片
-
本地执行前向传播、反向传播
-
跨卡梯度同步聚合
-
统一执行优化器更新,最终所有卡参数再次保持一致
同步后每张卡的参数保持一致。数据并行的主要通信开销来自梯度同步。
4. DataParallel 的问题与分布式的通信原语¶
PyTorch nn.DataParallel 是单进程多线程模式,通常由主卡负责 scatter、gather 和参数更新。
问题:
- 主卡负载重,容易成为瓶颈。
-
Python GIL 和线程调度开销明显。受Python全局解释器锁影响,多线程无法真正并行,调度开销大。
-
多机扩展能力弱。仅适配单机多卡,不支持多机分布式。
-
通信和计算重叠能力较差。
因此生产训练通常使用 DistributedDataParallel,而不是 DataParallel。
在开始下一节前,需要补充几种分布式的通信原语:
这些通信原语包括:Broadcast、Scatter、Gather、All-Gather、Reduce、All-Reduce、Reduce-Scatter、All-to-All等。
▉ 1、Broadcast(1对多的广播)
这个最简单。当主节点执行Broadcast操作时,数据会从主节点发送至其他所有节点。
Broadcast是一个典型的分发、散播行为。在分布式机器学习中,Broadcast常用于网络参数的初始化。 ▉ 2、Scatter(1对多的发散) Scatter也是一种分发、散播行为。它也是将主节点的数据发送至其他所有节点。只不过,Broadcast发送的是完整数据,而Scatter是将数据进行切割后,再分发,就像分生日蛋糕。
▉ 3、Gather(多对1的收集) Gather,是将多个sender(发送节点)上的数据收集到单个节点上,可以理解为反向的Scatter。
▉ 4、All-Gather(多对多的收集) Gather是多个到一个,All-Gather是多个到多个。 All-Gather是将多个sender(发送节点)上的数据收集到多个节点上。它相当于多个Gather操作。或者说,是一个Gather操作之后,跟着一个Broadcast操作。
▉ 5、Reduce(多对1的规约) Reduce的英文意思是“减少、降低”。在集合通信里,它表示“规约”运算,是一系列简单运算操作(包括:SUM、MIN、MAX、PROD、LOR等)的统称。 例如SUM,就是求和。MIN,就是找出最小值。其实说白了,Reduce就是:输入多个数,执行操作后,得到更少的数(例如1个数)。 下面这个,就是以ReduceSum(求和规约)为例:
▉ 6、All-Reduce(多对多的规约) All-Reduce,这个是我们在文章开头提到的,AI领域非常常见的一个词组。 在大模型训练中,经常会用到数据并行(DP)这个并行方式。里面就有AIl Reduce这个关键操作。 我们以All Reduce Sum(求和)为例: 首先,对所有节点进行数据收集。然后,对数据进行求和。再然后,把结果重新发回给所有节点。
在大模型训练中,Server GPU节点收集的数据,就是各个Worker GPU节点计算得出的“梯度”。求和之后再发回的过程,是“更新梯度”。 ▉ 7、Reduce-Scatter(组合的规约与发散) Reduce-Scatter稍微有点复杂、烧脑。 它是先归约(Reduce),再分散(Scatter)。具体来说: 首先,在所有参与计算的GPU节点上,对位于相同位置或索引的数据块执行指定的规约运算(例如求和SUM)。 接着,将规约后的完整结果按维度切分,并将不同的数据块分发给各个节点。最终,每个节点只得到整个规约结果的一部分,而不是全部。
简单来说,它先对所有数据进行“汇总计算”,然后再将计算好的结果“分散下发”。 ▉ 8、All-to-All(多对多的全互连) AIl-to-AII也是AI领域出现频率很高的一个词组。它是全交换操作,可以让每个节点都获取其他节点的值。 在使用All-to-All时,每一个节点都会向任意一个节点发送消息,每一个节点也都会接收到任意一个节点的消息。每个节点的接收缓冲区和发送缓冲区都是一个分为若干个数据块的数组。
All-to-All的具体操作是:将节点i的发送缓冲区中的第j块数据发送给节点j。节点j将接收到的来自节点i的数据块,放在自身接收缓冲区的第i块位置。 All-to-All与All-Gather相比较,区别在于:All-Gather操作中,不同节点向某一节点收集到的数据是完全相同的。而在All-to-All中,不同的节点向某一节点收集到的数据是不同的。在每个节点的发送缓冲区中,为每个节点都单独准备了一块数据。 上面这个图,大家可以发现,它就是一个矩阵倒置。 All-to-All的核心目标是重分布。它不进行聚合运算,而是专注于在不同节点间重新分布数据块。 ▉ 9、Ring-base collective(基于环的集合) 最后还要提一个有趣的结构——环(Ring)。 Ring-base collective是将所有的通信节点通过首位相连形成一个单向环,数据在环上依次传输。 传输方式有两种,一种是一次性传输全部,还有一种,是对数据进行切割,然后分别发送。
All-Reduce里有一种Ring All-Reduce(环形全规约)算法。它是通过组合Reduce-Scatter和All-Gather两个操作来实现的。 Ring All-Reduce算法分为两个阶段: 第一阶段,将N个worker分布在一个环上,并且把每个worker的数据分成N份。
对于第k个worker,这个worker会把第k份数据发给下一个worker,同时从前一个worker收到第k-1份数据。
然后,第k个worker会把收到的第k-1份数据和自己的第k-1份数据整合,再将整合的数据发送给下一个worker。
以此循环N次之后,每一个worker都会包含最终整合结果的一份。
第二阶段,每个worker将整合好的部分发送给下一个worker。worker在收到数据之后,更新自身数据对应的部分即可。 很显然,这种环形算法可以解决传统All-Reduce中Server节点的能力瓶颈问题。 以上就是常见通信原语的具体工作原理。 AI大模型训练推理任务,是由海量的GPU共同完成的。而这些GPU之间的通信,就是基于上面这些通信原语模型。 接下来我们继续回到DDP的学习。
5. DDP (分布式数据并行)的核心结构¶
DDP 推荐一张 GPU 一个进程:
核心概念: -rank:全局进程编号。
-
local_rank:当前节点内的 GPU 编号。 -
world_size:总进程数。 -
process group:参与通信的一组进程。 -
backend:通信后端,GPU 通常用 NCCL,CPU 可用 Gloo。
DDP 启动后,每个进程加载模型副本,处理自己的数据分片,并在 backward 时同步梯度。
6. DDP 的初始化流程¶
典型 PyTorch DDP 流程:
import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
dist.init_process_group(backend="nccl") # 初始化进程组,指定通信后端
local_rank = int(os.environ["LOCAL_RANK"]) # 获取当前节点内GPU编号,并绑定设备
torch.cuda.set_device(local_rank)
model = MyModel().to(local_rank) # 模型移至对应GPU,再包裹DDP
model = DDP(model, device_ids=[local_rank])
torchrun 启动:
多机训练还需要设置节点数、节点 rank、master 地址和端口。
7. DistributedSampler¶
DDP 中每个进程必须读取不同数据分片。常用:
sampler = torch.utils.data.distributed.DistributedSampler(dataset)
loader = DataLoader(dataset, sampler=sampler, shuffle=False)
shuffle=True 和 sampler,或忘记 set_epoch。
8. 梯度同步与 All-Reduce¶
DDP 在 backward 过程中对梯度做 all-reduce:
同步后每个 rank 拿到相同平均梯度,因此 optimizer step 后参数保持一致。 DDP 会把参数分成 buckets。某个 bucket 的梯度一旦计算完成,就可以启动通信,即Bucket DDP,从而实现: - 单个 bucket 梯度计算完毕,立刻发起通信-
实现反向计算与梯度通信重叠执行
-
大幅降低通信等待开销,相比 DP 效率显著提升
具体来说其实就是对每个需要更新同步的参数设置对应的hook(DDP 为每个参数的梯度张量注册 autograd hook),每当完成了梯度计算就激活对应的hook,来开启通信同步: - 当某个参数反向传播算出梯度(梯度张量就绪)→ 触发 hook;
-
hook 内部检查:该参数所属的 bucket 是否所有参数梯度都已就绪;
-
bucket 集齐全部梯度后,立即发起 All-Reduce 通信。
分bucket的原因的,如果每个参数都执行一次通信请求也会产生不必要的开销,但是如果bucket分的太粗了又浪费了通信与计算并行进行的优势。所以bucket如何分也是一个值得思考的问题。 感兴趣可以尝试简单去对hook进行实现玩一玩。 这是 DDP 性能好的重要原因。
9. DDP 与 Loss Reduction¶
在训练的时候,每张卡上都有一个mini-batch,所有卡的模型需要同时进行更新。需要先聚合所有卡上的loss得到全局loss后再对这个整个的batch进行更新。
通常每张卡计算本地 mini-batch loss。若 loss 是本地 batch 的 mean,DDP all-reduce 后梯度等价于全局 batch mean 梯度。
需要注意:
- 不要手动再把 loss 除以 world_size,否则梯度会过小。
-
日志指标需要跨 rank 汇总,不能只看 rank 0 的局部 loss。
-
如果每卡 batch size 不同,要小心 mean/sum 的语义。
10. 全局 Batch 与学习率¶
DDP 训练时:
扩大 GPU 数量会扩大 global batch。如果其他超参数不变,优化行为会变化。 常见经验: - 线性学习率缩放:global batch 增大 k 倍,学习率尝试增大 k 倍。-
warmup:大 batch 更需要 warmup。大 batch 训练初期梯度波动更大,必须加长热身阶段,稳定训练。
-
gradient accumulation:显存不足时累积多步再同步或 step。显存受限场景下,用多步梯度累积等效扩大全局 batch,无需增加硬件。
这些只是经验,最终仍要验证。
11. Gradient Accumulation 与 no_sync¶
连续执行 K 轮微批次 只累积梯度、不更新参数,最后统一执行 optimizer.step(),等效放大全局 batch。
DDP 默认每次 backward 都同步梯度。为了减少通信,可在非最后一步使用:
等到最后一个 micro-step 正常 backward,同步累积后的梯度。12. find_unused_parameters 与静态图¶
如果模型某些参数在某次 forward 中未参与 loss,DDP 可能等待不存在的梯度,导致报错或卡住。 可设置:
但这会增加 autograd 图遍历开销。若模型图固定,优先保持所有参数参与训练,或使用 static graph 优化。 动态控制流、多任务模型、条件路由、MoE、部分 adapter 训练中常遇到 unused parameters。13. BatchNorm 与 SyncBatchNorm¶
普通 BatchNorm 在每个 rank 上只看本地 batch 统计。如果 per-GPU batch 很小,统计不稳定。 可使用 SyncBatchNorm:
它跨进程同步 BN 统计。但 LLM 中通常使用 LayerNorm/RMSNorm,不常用 BatchNorm。14. AMP(自动混合精度) 与 DDP¶
DDP 可与混合精度训练结合:
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
loss = model(input)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
15. Checkpoint¶
DDP 中每个 rank 都有完整参数副本。常见做法是只在 rank 0 保存:
注意: - DDP 包装后原模型在model.module。访问原生模型参数必须用 model.module,直接使用 model 拿到的是 DDP 包装对象。
-
断点恢复要恢复全套状态: model、optimizer、scheduler、scaler、epoch、step、sampler 状态。
-
多机保存要避免多个 rank 同时写同一路径。
16. DDP 常见性能瓶颈¶
瓶颈来源: - dataloader 太慢。
-
GPU 计算不饱和。
-
batch size 太小。
-
网络带宽不足。
-
gradient bucket 设置不合理。
-
all-reduce 占比过高。
-
rank 间负载不均。
-
频繁 CPU-GPU 同步。
排查工具: - PyTorch profiler。
-
NVIDIA Nsight。
-
nvidia-smi dmon。 -
NCCL debug log。
-
训练吞吐 tokens/sec 或 samples/sec。
17. DDP、FSDP 与 ZeRO¶
DDP:
FSDP: ZeRO: 小到中等模型优先 DDP;模型大到单卡完整副本放不下时考虑 FSDP/ZeRO。18. 多机训练注意事项¶
多机训练需要: - 网络连通。
-
MASTER_ADDR、MASTER_PORT。 -
nnodes、node_rank。 -
每节点 GPU 数一致或清楚配置。
-
NCCL 网络接口配置。
-
共享存储或分布式 checkpoint 策略。
常见问题: - 防火墙或端口不通。
-
rank 配置错误。
-
NCCL timeout。
-
某个 rank OOM 导致所有 rank 卡住。
-
数据集每机路径不一致。
19. 常见误区¶
误区一:DDP 会自动把模型切到多卡上。 DDP 是数据并行,每个进程仍有完整模型副本。 误区二:DDP 能解决单卡放不下模型的问题。 DDP 不能。要用 FSDP、ZeRO、tensor parallel、pipeline parallel。 误区三:多卡训练只需要 batch size 乘以卡数。 global batch 变化会影响学习率、warmup、泛化和收敛。 误区四:只看 rank 0 loss 就够。 rank 0 loss 只是本地数据,需要跨 rank 汇总指标。 误区五:DDP 卡住一定是通信库问题。 unused parameter、某个 rank 数据耗尽、异常未同步、dataloader 死锁都可能导致卡住。
20. 核心总结¶
第 21 天需要掌握的最小闭环:
DDP:
one process per GPU
each rank has full model replica
data is sharded by DistributedSampler
backward triggers gradient all-reduce
optimizer step keeps all replicas identical
Key concepts:
rank, local_rank, world_size, process group, backend
Global batch:
per_gpu_batch * world_size * grad_accum
DDP vs FSDP:
DDP replicates states
FSDP shards states
DDP pitfalls:
sampler.set_epoch
unused parameters
loss scaling
checkpoint model.module
NCCL and dataloader bottlenecks
21. 参考资料¶
-
PyTorch DistributedDataParallel 文档:https://pytorch.org/docs/stable/generated/torch.nn.parallel.DistributedDataParallel.html
-
PyTorch Distributed Overview:https://pytorch.org/tutorials/beginner/dist_overview.html
-
PyTorch FSDP 文档:https://pytorch.org/docs/stable/fsdp.html
-
DDP 实战视频:https://www.bilibili.com/video/BV1wS421w7ug/
-
对比 DP、DDP 和 FSDP:https://zhuanlan.zhihu.com/p/650002268
预览时标签不可点<div class="