云镜收藏

稍后阅读

清单保存在当前浏览器,方便下次回来继续阅读。

清单还是空的

在笔记卡片或正文页点击“加入稍后阅读”即可收藏。

展开课程与本章目录
12 章学习路线01零基础准备:Python、NumPy 与正确安装02起步:环境、设备与第一枚 Tensor03张量基本功:索引、广播与线性代数04自动求导:从计算图到反向传播05神经网络工程:Module、损失与训练循环06数据管线:Dataset、DataLoader 与批处理07亲手搭 Transformer:从 Embedding 到注意力08稳定训练:初始化、AMP、裁剪与排错09性能与显存:Profiler、compile 与检查点10分布式训练:从 DDP 到 FSDP 与 TP11大模型生态:Transformers、PEFT、torchao 与 TorchTitan12方向实战:训练并优化一个迷你语言模型本章课程01先建立 rank、world size 与进程组概念02集合通信:all-reduce 与 reduce-scatter03DistributedDataParallel 的正确骨架04FSDP2 为什么能训练更大模型05Tensor Parallel 与二维并行06分布式挂起与排错清单

第 10 章 · 高级

分布式训练:从 DDP 到 FSDP 与 TP

理解进程、rank 和集合通信,掌握 DDP、FSDP2、张量并行的选择逻辑,以及分布式训练最常见的挂起问题。

6 节课程校订于 2026年8月31日
DDPFSDP2Tensor Paralleltorchrun

预计用时170 分钟

前置知识单卡训练闭环 · 基础网络与进程概念

完成标准选择 DDP/FSDP2/TP · 推导状态复制关系 · 定位 collective 挂起

本章进度0 / 6

先建立 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 等待的通信操作彼此不匹配,无法组成同一次集合通信。