W03 · S15—S21 · 总 Day 105—111
分布式训练:先算内存,再谈并行
拆解混合精度状态与 ZeRO 分片,估算通信下界。
作者准备的学习示例 · 不计真实学习进度 · 不代表生产 / GPU / 真机结果
核心问题
“7B 参数、半精度,约 14GB,为什么训练放不下?”因为参数只是账本的一行。梯度、优化器、主权重、激活、临时张量和通信缓冲区都占空间,而且峰值发生的时刻不同。
对应 S15~S21。这里用数量级估算建立系统直觉,不运行 GPU、不模拟真实计算性能。
1. 写清精度假设再计算
本例选择一种经典混合精度 Adam 布局:低精度权重 2 B、低精度梯度 2 B、FP32 主权重 4 B、两份 FP32 优化器状态 8 B,共 16 B/参数。并非所有框架都采用这套布局;实际梯度精度、主权重是否存在、优化器类型都会改变总量。
设参数量 P、设备数 N。这里只算每卡持久模型状态,单位是 bytes:
| 策略 | 简化状态量 / 每卡 | 分片的对象 |
|---|---|---|
| DDP | 16P | 状态均复制,数据分配到不同设备 |
| ZeRO-1 | 4P + 12P/N | 优化器与主权重状态 |
| ZeRO-2 | 2P + 14P/N | 再加梯度 |
| ZeRO-3 | 16P/N | 再加低精度参数 |
这些是特定假设下的稳态账本,不是训练峰值。ZeRO-3/FSDP 一类参数分片需要在计算时重建相关参数,临时聚集、prefetch、通信缓冲和激活可能显著抬高峰值。
2. 用输出校准数量级
npm run learning:p2 -- w03
7×10^9 参数时,DDP 每卡状态约 104.308 GiB,设备数量增加也不会把这份复制状态自动分掉。4 卡时 ZeRO-3 的持久状态约 26.077 GiB;8 卡约 13.039 GiB。这里 1 GiB = 2^30 bytes,不是十进制 GB。
“13.039 GiB 小于显存容量”不能直接推出能训练。先补上激活,再考虑最长序列、microbatch、重计算、通信与分配器,再用目标实现测峰值。显存账本用于排除明显不可行方案,不替代 profiling。
3. 分片省内存,通信仍有成本
对 N 个 rank 的简化 ring all-reduce,一份大小为 G 的梯度,每 rank 传输量近似 2(N-1)G/N。用它除以假设可用带宽,得到带宽项耗时。本例 4 卡、FP16 梯度和 25 GiB/s 假设下约 0.782 秒。
这个值不含消息启动延迟、链路竞争、拓扑、bucket 划分与计算通信重叠。它是指定模型中的通信耗时估算,不是步时,也不直接适用于 ZeRO 各阶段的全部通信。增加卡数可能减少每卡计算,却增加协调成本,不能只用“卡数翻倍”预测速度。
4. 与其他并行轴区别开
数据并行切样本;参数状态分片解决重复状态;张量并行切算子中的张量计算;流水并行切层并安排 microbatch;专家并行切专家。它们解决的瓶颈不同,可以组合,但组合后拓扑、通信路径与调度都更复杂。
轻量修改:把设备数固定为 4,只把带宽 25 改为 10。内存不应改变,通信带宽项耗时增大。然后把 P 从 7e9 改为 1e9,观察两个量如何缩放。不要据此推导训练质量或 scaling law。
5. 专业阅读与未来问题
- Stanford CS336 2025:选 systems、parallelism 相关课程,先问“哪一类状态放在哪里”。Assignment 2 是深入选项,不要求整套完成。
- PyTorch Distributed Overview:对照 DDP、FSDP、TP 和 pipeline 的概念边界;真正实现时再选择与本地版本匹配的教程。
AGI 研究不能只看算法,也受算力、带宽、实验周转时间和数据预算制约。但这份 TypeScript 计算器只能训练资源估算能力,不提供模型训练结果。可把“哪一个遗漏项最可能改变方案”作为本周唯一笔记问题。