第 10 章 · 高级
分布式训练:从 DDP 到 FSDP 与 TP
理解进程、rank 和集合通信,掌握 DDP、FSDP2、张量并行的选择逻辑,以及分布式训练最常见的挂起问题。
先建立 rank、world size 与进程组概念
本节目标
本节围绕“先建立 rank、world size 与进程组概念”展开。运行示例时记录 rank、world size、数据分片与 collective 顺序,确认各进程步数一致。
清晰讲解
分布式训练通常一张 GPU 对应一个进程。rank 是进程全局编号,world size 是总进程数,进程组定义哪些进程参与通信。torchrun 负责注入环境变量并拉起进程。
核对“先建立 rank、world size 与进程组概念”中的 rank 映射、分片范围、collective 次序、通信量和同步边界。
代码示例
# 保存为 dist_hello.py,再运行:torchrun --standalone --nproc-per-node=2 dist_hello.py
import os, torch.distributed as dist
dist.init_process_group('nccl')
rank = dist.get_rank()
local_rank = int(os.environ['LOCAL_RANK'])
print('rank', rank, 'local_rank', local_rank, 'world', dist.get_world_size())
dist.destroy_process_group()
运行结果与观察
每个进程应打印唯一 rank、相同 world size 与正确 local device;程序结束前所有 rank 都应通过 barrier 并正常销毁进程组。
常见错误
- 多个进程绑定同一 GPU;每个 rank 都写同一文件;某个 rank 提前退出导致其余挂起。
- 排查“先建立 rank、world size 与进程组概念”时只查看单个 rank,没有对齐各进程的 step、collective、数据分片和错误日志。
与大模型方向的连接
理解 rank 后,数据切分、日志只写一次、checkpoint 谁保存等问题才有明确答案。
动手练习
在两进程脚本中只让 rank 0 打印最终汇总。
查看参考答案与验收点
先写出 world size、每 rank 数据和状态复制关系,再验证所有 rank 以相同顺序进入 collective;完整命令需在对应硬件执行。
官方资料
集合通信:all-reduce 与 reduce-scatter
本节目标
本节围绕“集合通信:all-reduce 与 reduce-scatter”展开。运行示例时记录 rank、world size、数据分片与 collective 顺序,确认各进程步数一致。
清晰讲解
all-reduce 先聚合再把结果发回每个 rank,DDP 用它同步梯度。reduce-scatter 则聚合后让每个 rank 只保留一片,FSDP 会结合 all-gather 实现分片参数训练。所有 rank 必须以相同顺序参与集合通信。
核对“集合通信:all-reduce 与 reduce-scatter”中的 rank 映射、分片范围、collective 次序、通信量和同步边界。
代码示例
import torch, torch.distributed as dist
x = torch.tensor([dist.get_rank() + 1.0], device='cuda')
dist.all_reduce(x, op=dist.ReduceOp.SUM)
print(dist.get_rank(), x.item()) # 两 rank 时均为 3
运行结果与观察
若每个 rank 初值为自身编号,sum all-reduce 后所有 rank 应得到 world_size*(world_size-1)/2;reduce-scatter 的各分片拼接后应等价于完整归约结果。
常见错误
- 某些 rank 跳过 collective;通信 Tensor 的 shape/dtype 不一致;在通信前频繁同步 CPU。
- 排查“集合通信:all-reduce 与 reduce-scatter”时只查看单个 rank,没有对齐各进程的 step、collective、数据分片和错误日志。
与大模型方向的连接
分布式并行策略的本质是决定哪些状态复制、哪些切片,以及何时通信重组。
动手练习
画出 4 rank all-reduce 前后每个 rank 上的数据变化。
查看参考答案与验收点
验收:各 rank 步数与 collective 一致,分片无重复遗漏,并能说明“集合通信:all-reduce 与 reduce-scatter”的通信和显存代价。
官方资料
DistributedDataParallel 的正确骨架
本节目标
本节围绕“DistributedDataParallel 的正确骨架”展开。运行示例时记录 rank、world size、数据分片与 collective 顺序,确认各进程步数一致。
清晰讲解
DDP 每个进程持有完整模型副本,处理不同数据,反向时同步梯度。适合模型能放进单卡、目标是提高数据吞吐的情况。要配合 DistributedSampler,每轮调用 set_epoch 改变洗牌。
核对“DistributedDataParallel 的正确骨架”中的 rank 映射、分片范围、collective 次序、通信量和同步边界。
代码示例
import os, torch
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader, DistributedSampler
local_rank = int(os.environ['LOCAL_RANK'])
torch.cuda.set_device(local_rank)
model = DDP(model.to(local_rank), device_ids=[local_rank])
sampler = DistributedSampler(dataset, shuffle=True)
loader = DataLoader(dataset, sampler=sampler, batch_size=8)
for epoch in range(epochs):
sampler.set_epoch(epoch)
train(model, loader)
运行结果与观察
每个 rank 处理不同数据分片,反向后对应参数梯度应一致;一个 epoch 后 sampler 的 set_epoch 会改变各卡顺序但仍不重不漏。
常见错误
- 使用 DDP 又设置 shuffle;忘记 set_epoch;只在 rank 0 反向。
- 排查“DistributedDataParallel 的正确骨架”时只查看单个 rank,没有对齐各进程的 step、collective、数据分片和错误日志。
与大模型方向的连接
DDP 不减少单卡模型状态显存;当模型本身放不下时,应考虑 FSDP 或模型并行。
动手练习
解释两卡 DDP 下 batch_size=8、累积 4 步时全局 batch 是多少。
查看参考答案与验收点
验收:各 rank 步数与 collective 一致,分片无重复遗漏,并能说明“DistributedDataParallel 的正确骨架”的通信和显存代价。
官方资料
FSDP2 为什么能训练更大模型
本节目标
本节围绕“FSDP2 为什么能训练更大模型”展开。运行示例时记录 rank、world size、数据分片与 collective 顺序,确认各进程步数一致。
清晰讲解
FSDP2 将参数、梯度和优化器状态跨 rank 分片。计算某个模块前 all-gather 所需参数,计算后释放,再以 reduce-scatter 聚合梯度。它节省模型状态显存,但增加通信并要求合理切分模块。
核对“FSDP2 为什么能训练更大模型”中的 rank 映射、分片范围、collective 次序、通信量和同步边界。
代码示例
# 概念示例:需要已初始化的分布式环境
from torch.distributed.fsdp import fully_shard
for block in model.blocks:
fully_shard(block)
fully_shard(model)
loss = model(input_ids).loss
loss.backward()
optimizer.step()
运行结果与观察
多卡运行时各 rank 只持有参数分片,峰值显存应低于复制完整模型的 DDP 对照;一步训练后 loss 有限,保存再加载的权重输出应一致。
常见错误
- 把 FSDP1 旧 API 与 FSDP2 混用;只包最外层造成重组峰值过大;每个 rank 保存完整 checkpoint。
- 排查“FSDP2 为什么能训练更大模型”时只查看单个 rank,没有对齐各进程的 step、collective、数据分片和错误日志。
与大模型方向的连接
当 Transformer 单卡放不下时,按 block 分片通常比 DDP 合适;FSDP2 是当前官方教程推荐入口。
动手练习
用文字比较 DDP 与 FSDP2 在参数、梯度、优化器状态上的复制情况。
查看参考答案与验收点
用一个极小 batch 过拟合,并记录 loss、梯度范数和参数更新;断点恢复后比较下一步结果,而不只检查文件存在。
官方资料
Tensor Parallel 与二维并行
本节目标
本节围绕“Tensor Parallel 与二维并行”展开。运行示例时记录 rank、world size、数据分片与 collective 顺序,确认各进程步数一致。
清晰讲解
张量并行把单个矩阵乘的权重按行或列切到多卡,每层都可能通信。它适合单个层太大或需要进一步扩展;实践中常与 FSDP 组成二维 mesh:一维做分片数据并行,一维做模型内张量并行。
核对“Tensor Parallel 与二维并行”中的 rank 映射、分片范围、collective 次序、通信量和同步边界。
代码示例
# 接口示意:具体计划取决于模型层命名
from torch.distributed.device_mesh import init_device_mesh
from torch.distributed.tensor.parallel import ColwiseParallel, RowwiseParallel, parallelize_module
mesh = init_device_mesh('cuda', (2,), mesh_dim_names=('tp',))
plan = {'wq': ColwiseParallel(), 'wo': RowwiseParallel()}
parallelize_module(attention, mesh['tp'], plan)
运行结果与观察
列并行与行并行组合后的输出应与未切分线性层在容差内一致;设备网格维度乘积必须等于 world size,并记录每个 rank 的局部 shape。
常见错误
- 模型能单卡放下却过早上 TP;行列切分计划不匹配前后层布局;忽视互联拓扑。
- 排查“Tensor Parallel 与二维并行”时只查看单个 rank,没有对齐各进程的 step、collective、数据分片和错误日志。
与大模型方向的连接
大模型训练不是并行方式越多越好;每增加一维,配置、通信和故障面都会扩大。
动手练习
对一个 [4096, 16384] Linear 描述 4 路列并行后每卡权重形状。
查看参考答案与验收点
先把每个轴写成语义:例如 [B,T,C] 分别是批次、序列和隐藏维;再用 assert tensor.shape == (...) 固化预期。
官方资料
分布式挂起与排错清单
本节目标
本节围绕“分布式挂起与排错清单”展开。运行示例时记录 rank、world size、数据分片与 collective 顺序,确认各进程步数一致。
清晰讲解
挂起通常意味着不同 rank 走了不同控制流、collective 顺序不一致、某个进程 OOM/异常或网络配置错误。先保存每个 rank 独立日志,打开分布式调试信息,再缩小到最少节点与批次。
核对“分布式挂起与排错清单”中的 rank 映射、分片范围、collective 次序、通信量和同步边界。
代码示例
# 启动前设置(Linux/macOS shell)
# TORCH_DISTRIBUTED_DEBUG=DETAIL NCCL_DEBUG=INFO torchrun ...
try:
train()
finally:
if dist.is_initialized():
dist.destroy_process_group()
运行结果与观察
用较短 timeout 故意让一个 rank 少调用一次 collective,应快速报出不匹配位置;修复后所有 rank 的通信序列、shape 与 dtype 必须一致并正常退出。
常见错误
- 只看 rank 0 日志;所有 rank 写一个日志文件互相覆盖;挂起后反复延长 timeout。
- 排查“分布式挂起与排错清单”时只查看单个 rank,没有对齐各进程的 step、collective、数据分片和错误日志。
与大模型方向的连接
大规模训练的首要能力不是“会启动”,而是能从 rank、step、collective 和硬件日志定位故障边界。
动手练习
设计一份每个 step 至少记录 rank、batch id 和显存的调试日志格式。
查看参考答案与验收点
先写出 world size、每 rank 数据和状态复制关系,再验证所有 rank 以相同顺序进入 collective;完整命令需在对应硬件执行。
官方资料
章末资料
小测、项目与相关面试题
1. 模型能放单卡、目标是提吞吐,优先选什么?
通常优先 DDP;它复制模型并用数据并行提高吞吐。
2. 为什么 collective 顺序不一致会挂起?
不同 rank 等待的通信操作彼此不匹配,无法组成同一次集合通信。
