S15:训练内存、Collective 与并行轴全景
分布式训练通过在设备间切分样本、参数、层、张量、专家或上下文来突破单设备容量与时间限制,同时引入通信、同步、负载不均和恢复成本。
内容类型:预习教材(不代表已完成)
日期:2026-12-07
阶段:P2 · AI Systems Engineering 90
总路线:Day 105 / 360
周次 / 节奏:W3 · 周一概念与阅读
状态:教材已备;学习未完成
主题:training memory、collectives、DP/FSDP/TP/PP/EP/CP
一句话定义
分布式训练通过在设备间切分样本、参数、层、张量、专家或上下文来突破单设备容量与时间限制,同时引入通信、同步、负载不均和恢复成本。
学习目标
- 能拆解参数、梯度、优化器状态、activation 与临时 buffer 的内存来源。
- 能说清 data parallel、FSDP/ZeRO、tensor、pipeline、expert 和 context parallel 分别切分什么。
- 能解释 all-reduce、all-gather、reduce-scatter 与 all-to-all 在训练中的基本角色。
- 能从模型大小、互联和工作负载出发讨论选择,而不是背“最佳方案”。
核心知识
训练显存不只等于参数大小。以混合精度 Adam 为例,可能同时存在低精度参数、梯度、FP32 master weights、两份 optimizer moments、activation、通信 bucket 和框架临时空间。Activation 随 micro-batch、序列长度、层数和隐藏维度增长,checkpointing 用额外重算换内存。
Data parallel 在每个 rank 放完整模型,切分 batch,并同步梯度;FSDP/ZeRO 进一步分片参数、梯度或优化器状态;tensor parallel 在一层内部切分矩阵计算;pipeline parallel 按层分 stage,以 micro-batch 填流水;expert parallel 把 MoE 专家分布到设备;context/sequence parallel 切分长序列相关计算。多种并行可组合,但协调复杂度和通信路径迅速增加。
Collective 是多 rank 协同的数据运动语义。All-reduce 常用于聚合梯度;reduce-scatter 聚合后每个 rank 留一片;all-gather 收集分片;all-to-all 常见于 token 路由到专家。通信开销取决于数据量、拓扑、带宽、延迟以及能否与计算重叠。
机制与推导
若参数量为 N,每参数字节分别为 b_p,b_g,b_o,基础模型状态近似:
[ M_{state}\approx N(b_p+b_g+b_o) ]
Adam 的 b_o 常包含两份 FP32 moment,实际还可能有 master weight。DP 每 rank 复制全部状态;ZeRO-1 分 optimizer,ZeRO-2 再分 gradient,ZeRO-3/FSDP 再分 parameter,理想状态项约除以 world size P,但 activation、buffer 和峰值 gather 不会同样消失。
训练步时间可粗分:T_step = T_compute + T_comm_unhidden + T_input + T_sync_wait。并行扩大后,局部计算减少但通信和 straggler 更显著。Pipeline 的 bubble 比例与 stage 数和 micro-batch 数相关;micro-batch 太少时设备空闲,太多则调度与 activation 状态增加。
最小练习或观察步骤
- 选一个假想 1B 参数模型,列出参数、梯度、Adam 状态与 activation 四栏。
- 分别画 DP、FSDP、TP、PP 的 4 卡示意图,只标“哪一维被切分”。
- 为每张图写主要 collective 或点对点通信,以及一次可能的等待。
- 假设设备互联很慢,判断哪些策略的代价会放大并说明原因。
- 记录一个无法从公式单独判断的现实因素,如 kernel、拓扑或数据加载。
常见误区与边界
- 用“参数×精度”估算训练内存,遗漏梯度、优化器和峰值 buffer。
- 认为 world size 翻倍,吞吐必定翻倍且每卡内存精确减半。
- 将 ZeRO 与 tensor parallel 混为一谈;前者主要分状态,后者分层内计算。
- 把 collective 名称当具体算法;ring/tree 等实现和硬件拓扑仍会不同。
- 把本地 Apple Silicon 当多 GPU 环境;本周可以公式、CPU/Gloo 或模拟学习。
系统场景连接
企业通常不会从零预训练超大模型,但理解训练系统有助于评估 fine-tuning 作业为何 OOM、checkpoint 为何缓慢、云 GPU 利用率为何不等于有效吞吐,也能判断厂商宣传中的“线性扩展”是否包含通信与失败恢复。模型 artifact 的分片格式还会直接影响后续加载和 serving。
自检问题
- 混合精度 Adam 训练为何远大于仅存模型权重的内存?
- FSDP 与 TP 分别切什么,通信时点有何不同?
- all-gather、reduce-scatter 和 all-to-all 各常出现在哪类并行?
- 为什么增加设备可能降低整体效率?
专业课程对齐
- 阅读 CMU Deep Learning Systems 的训练系统、自动微分与分布式执行相关课程,重点把计算图节点和张量放置连接到通信需求。
- 阅读 PyTorch Distributed 官方文档 的 process group 与 collective 概览,辨认 all-reduce、all-gather、reduce-scatter 的输入输出语义,不要求多机运行。
- 阅读 MIT 6.5840 的分布式系统课程主线,关注并发、故障与一致性为何同样约束训练协调,而非完成其全部实验。
深入学习提示
阅读顺序建议为内存账本→并行切分图→collective 数据流。每个并行策略都回答三问:每个 rank 持有什么、何时需要别人的数据、失败后从哪里恢复。用极端反例检查理解:模型放不下但网络很慢;模型能放下但 batch 太小;MoE 专家负载极不均。不要把策略按先进程度排序,而要把它们放进约束空间。
学后填写区
- 我的训练内存账本:
- 六种并行各切分什么:
- 最难理解的 collective:
- 一个通信/容量取舍:
- 尚未验证的假设: