第 06 章 · 进阶
大模型怎样在多张卡上训练
从显存组成出发,理解数据并行、ZeRO/FSDP、张量并行和流水线并行。
学习准备与本章目标
- 开始前
- 完整训练循环 · 模型参数与激活概念
- 学完后
- 估算主要显存占用 · 区分常见并行策略 · 理解通信为何成为瓶颈
本章路线图
分布式训练解决两个不同问题:单卡装不下,以及单卡算得太慢。学习任何并行策略时都回答三问:什么张量被切分、哪里发生通信、每张卡保存哪些状态。
一张卡为什么装不下训练
训练显存不只有模型权重,还包括梯度、优化器状态、激活、临时缓冲和通信空间。AdamW 常为每个参数维护一阶与二阶状态;激活又随 batch、序列长度、层数和隐藏维度增长。
低精度权重能减少一部分容量,但不会自动消除优化器和激活。判断瓶颈时要先分账,而不是只用“参数量 × 两字节”估算全部显存。
数据并行与参数分片
DDP 在每张卡复制完整模型,各自处理不同 batch,再同步梯度。它简单且吞吐高,但前提是单卡能放下模型、梯度和状态。
ZeRO 与 FSDP 会把优化器状态、梯度甚至参数分片到不同 rank,需要时再聚合。分片越彻底,单卡显存越低,但通信次数和实现复杂度越高。
张量并行与流水线并行
Tensor Parallel 把一个大矩阵运算沿维度切到多卡,每一层都可能发生集合通信,适合层本身就放不下单卡。Pipeline Parallel 则把不同层放到不同阶段,用 micro-batch 让流水线保持忙碌,但会出现气泡和调度问题。
实际大规模训练常组合数据、张量、流水线和序列并行。选择顺序从瓶颈出发:模型是否能装下、单层是否能装下、通信拓扑怎样、目标是省显存还是提吞吐。
分布式故障为何难排查
所有 rank 必须以一致顺序进入 collective;某个进程提前异常或执行不同分支,其余进程就可能一直等待。数据长度不一致、超时、网络错误和保存 checkpoint 也会放大问题。
排查时保留 rank、step、最后一次 collective 和数据批次信息,先在更少 GPU 上复现。吞吐要按有效 Token 计算,同时观察计算与通信是否重叠。
先做一张显存账单
以混合精度 AdamW 为例,模型状态可能包含低精度参数、梯度、FP32 master weight,以及两个 FP32 动量。粗略模型状态可达每参数 12–16 字节,具体取决于框架是否保留 master weight 与梯度 dtype。
训练峰值还包括激活、注意力临时量、通信 bucket、CUDA context、算子 workspace 和碎片。估算时要使用峰值而非稳定均值,并留出安全余量。
Activation checkpointing 通过前向少存、反向重算来降低激活显存;它不分片模型状态,也不会直接减少参数显存。
DDP 的同步过程
DistributedDataParallel 在每个 rank 保存完整模型,处理不同数据。反向计算某个梯度 bucket 完成后触发 all-reduce,各 rank 得到相同平均梯度,再独立执行同样的优化器更新。
DDP 几乎不省模型状态显存,优势是实现清晰且可让梯度通信与反向计算重叠。数据采样器必须让各 rank 获得不重复且步数一致的数据;日志、保存和验证也要明确由哪个 rank 执行。
梯度累积时,非最终 micro-step 可用 no_sync() 避免重复 all-reduce,最后一次再同步。
ZeRO 与 FSDP 分片了什么
ZeRO 可以逐级分片:
| 阶段 | 主要分片对象 | 直觉 |
|---|---|---|
| Stage 1 | 优化器状态 | 每个 rank 只保存部分 Adam 状态 |
| Stage 2 | 优化器状态 + 梯度 | 继续减少重复梯度 |
| Stage 3 | 优化器状态 + 梯度 + 参数 | 计算前按需聚合参数 |
FSDP 同样把参数、梯度和优化器状态分片,并围绕模块执行 all-gather 与 reduce-scatter。节省越多,通信与实现复杂度通常越高。包装粒度过大会造成聚合峰值,过小则产生许多细碎 collective。
分片后不能简单用“总状态除 GPU 数”预测峰值,因为计算当前层时仍需临时聚合完整参数,还存在预取和通信缓冲。
张量并行怎样切一层
当单层矩阵都装不下或算得太慢,可沿矩阵维度切分。Megatron 风格常把某些线性层列切分、下一层行切分,使中间结果局部保留,并在合适位置做 all-reduce。
注意力可按头切分,FFN 可按中间维切分。张量并行每层都需要高频通信,适合放在 NVLink 等高速互联域内;跨慢网络使用会严重损失吞吐。
Sequence Parallel 会把某些非张量并行计算沿序列维切分,降低重复激活;Context Parallel 则面向超长序列切分注意力上下文。名称相近,但切分对象和通信模式不同。
流水线并行与气泡
流水线并行把连续层分成多个 stage,不同 micro-batch 像流水线一样交错执行。最初填充和最后排空期间会有设备空闲,称为 pipeline bubble。
增加 micro-batch 数可以降低气泡比例,却会增加调度、激活保存和全局 batch。1F1B 等调度交替执行前向与反向,以控制激活峰值。层数和计算不均衡时还要手工平衡 stage。
流水线解决跨层切分,张量并行解决层内切分,数据并行复制一组模型切分处理更多数据;大规模训练常把三者组合成 3D 并行。
通信、拓扑与算术强度
常见 collective 包括 all-reduce、all-gather、reduce-scatter 和 all-to-all。选择并行方案时不仅看通信字节,还要看是否能与计算重叠、消息粒度和物理拓扑。
同一节点内 GPU 可能通过 NVLink/NVSwitch 连接,跨节点则经过网卡与交换网络。通常把高频张量并行放在高速域,把通信频率较低的数据并行扩到节点间。MoE 专家并行的 all-to-all 对拥塞尤其敏感。
分布式数据与随机性
每个 rank 应获得确定、互斥的数据分片,并在 epoch 或全局步上保持一致。变长序列可能造成某些 rank 计算更久,形成 straggler;按 Token 数平衡 batch 往往比按样本数更有效。
要复现训练,需记录全局 seed、数据采样器状态、每 rank 随机状态与全局 step。仅设置一个 seed 不保证不同并行规模下逐位一致,因为算子顺序和归约顺序会改变。
分片 Checkpoint 与弹性恢复
大模型 checkpoint 可能达到 TB 级,单 rank 汇总后再保存会造成内存和 IO 峰值。分布式 checkpoint 让各 rank 保存自己的分片,但还要记录分片布局、模型配置、优化器映射和版本。
理想恢复流程应支持:验证文件完整性、从中途失败继续、在允许范围内改变数据并行度,并定期实际演练。未验证可恢复的 checkpoint 不能算可靠备份。
挂起与性能问题的排查顺序
- 找到所有 rank 最后成功的 step 与 collective。
- 检查是否有 rank 更早出现 OOM、NaN 或数据异常。
- 用更少节点、更短序列和固定样本复现。
- 打开 collective 与网络日志,确认调用顺序一致。
- 用 profiler 区分计算、通信、数据加载和同步等待。
- 按有效 Token 计算吞吐,并观察最慢 rank,而非只看平均值。
章末检查
- DDP、FSDP、张量并行和流水线并行分别切分什么?
- ZeRO-3 为什么省显存却增加参数聚合通信?
- activation checkpointing 与参数分片解决的对象有何不同?
- 为什么张量并行更适合节点内高速互联?
- 一个 rank 挂起时,为什么其他 rank 可能没有错误只是在等待?
延伸资料
本章进阶内容
原理专题与代码实践
完成主教材后,按顺序阅读原理专题、完成代码实践,并用掌握标准复查本章内容。
- 01
从单卡内存账本理解为什么必须并行
- 02
比较 DP、TP、PP、ZeRO 与通信代价
- 03
实践 DDP、FSDP 与多维并行接口
- 04
为指定模型与集群画并行切分方案
本章术语
六个关键词
- Rank
- 分布式进程的唯一编号。
- All-Reduce
- 聚合各 rank 数据并把结果发送给全部 rank。
- DDP
- 复制模型、分配数据并同步梯度的数据并行。
- FSDP
- 分片参数、梯度和优化器状态的训练方式。
- Tensor Parallel
- 在多个设备之间切分单层矩阵计算。
- Pipeline Bubble
- 流水线填充和排空期间的设备空闲。
章节记录
