分布式训练:沿数据、张量、流水线和状态四个轴拆分
从全局批量与 all-reduce,到张量并行、pipeline bubble、ZeRO/FSDP、3D 并行和故障检查点。
- 核算参数/状态/激活峰值
- 选择 DP/TP/PP/分片组合
- 按拓扑映射通信域
- 固定全局批量与数值语义
- trace 优化计算通信重叠
- 真实恢复并以质量/效率联合验收
1权重能放下,不代表训练能放下直觉
判断一块 GPU 能否训练某个模型,只看权重体积是不够的。以 7B 参数量的模型为例:7B 表示约 70 亿个参数,采用 FP16 存储时每个权重占 2 字节,推理时仅权重就需要 7×10⁹ × 2 bytes ≈ 14 GB。这看起来能放进 24GB 显存的 GPU,于是很容易误以为在同一块卡上做全量微调也没有问题。但推理只需要在前向过程中用一次这些权重;训练则要在反向传播中为每个参数保留一整套额外状态。
具体来说,混合精度训练通常维护一份更高精度的主权重用于参数更新,一份与权重同形状的梯度,以及 Adam 这类自适应优化器为每个参数维护的两组历史统计量(一阶矩与二阶矩)。此外,前向传播期间各层产生的中间结果需要暂存下来供反向传播使用,这部分称为激活。把这些加起来,混合精度下每个参数大致对应 12–20 字节的显存状态;对 7B 模型而言,总量轻松达到百 GB 量级,远超出单张 24GB 卡。如果序列很长、批量较大,激活部分甚至可以单独成为最大的内存开销。
内存账的输入是参数量、数值精度、优化器选择、批量大小、序列长度和临时缓冲;输出是每一类状态的峰值占用,以及据此判断每张卡能否容纳。它直接决定了该选择哪条并行轴:数据并行复制模型、ZeRO/FSDP 分片优化器状态、张量并行拆分单个层、流水线并行拆分层序列、激活检查点用重复计算换取更少的激活存储。这些策略解决的对象各不相同,选择它们的前提是先弄清内存到底被什么占据。
因此“权重能放下”只证明推理主体可以装入显存,不证明反向传播能运行。即便账目上各项之和小于显存容量,实现细节、显存碎片和通信缓冲也会吃掉额外空间,实际训练必须用实测的峰值占用并留出安全余量。正确的顺序是先逐项估算权重、梯度、优化器状态、激活、通信缓冲和碎片,再决定采用哪条并行轴。
2数据并行的语义是同步同一更新数据并行
数据并行的目标不是让各卡训练不同的模型,而是让所有卡以完全相同的方式更新同一个模型。每个 rank(参与训练的进程或设备编号)都持有一份完整的模型副本,各自读取不同的样本子集,在前向和反向之后计算出本地梯度。此时各副本的梯度彼此不同,因为它们看的是不同数据。要让参数保持同步,需要一次集合通信操作 all-reduce:所有 rank 把各自的本地梯度求和,再把同一个和发回每一个参与方。归约完成后,每个 rank 用同一个全局平均梯度执行参数更新,因此所有副本的参数保持一致。
平均的含义体现在公式里。设 b 为每张卡每个微步的局部 batch,n 为数据并行副本数,a 为梯度累积步数,则一次参数更新实际使用的全局批量 B全局 = b × n × a。每个 rank 把自己 a 个微步的梯度累加,得到第 r 个 rank 的累计梯度 gᵣ;all-reduce 对所有 gᵣ 求和并广播给每个副本,各副本再用平均梯度 ḡ = (Σᵣ gᵣ)/n 更新参数。三个因子相乘而不是相加,是因为每一个微步都同时在全部 n 张卡上发生,累积 a 次之后才触发一次更新:B全局 描述的是“一次更新实际消费多少样本”,b 描述的是“每张卡一次微步能放下多少样本”,a 描述的是“攒多少步才更新一次”。
这个关系直接约束了扩展方式。如果单纯增加设备数 n 而不调小 b 或 a,B全局 会随之增大;更大的全局批量意味着梯度噪声不同、所需的学习率调整不同,训练中实际消耗的 token 语义也改变了。此时若观察到收敛行为变化,原因首先是实验设置变了,而不是硬件本身。另一方面,梯度归约的数值精度与求和顺序也会影响结果:浮点加法不满足结合律,不同设备数或不同归约顺序可能得到略有差异的 ḡ。反过来,loss 数值看起来相同也不代表实验相同,数据顺序、随机数流、dropout 采样的差异都会把训练轨迹带到不同的位置。比较数据并行实验时,需要对齐的是 B全局、学习率与随机性配置,而不只是最终 loss。
3运行示例:8 卡全局批量与强扩展效率逐步演算
把数据并行的公式放进一个具体配置里。设每张卡每个微步处理 4 个样本(b = 4),累积 8 步才更新一次(a = 8),共 8 张 GPU(n = 8)。那么一次优化步实际看到的样本数是 B全局 = 4 × 8 × 8 = 256。八个数据并行 GPU 各自执行局部微批的前向与反向,并在本地累积梯度;累积期间不必每个 micro-step 都通信,只有第 8 个微步结束后才触发一次 all-reduce,同步之后执行一次全局参数更新。这正是梯度累积对通信成本的直接回报:通信频率从“每微步一次”降到“每 a 步一次”。
强扩展效率衡量的是固定总工作量下,增加设备能把耗时压到多短。假设单卡处理这批固定工作耗时 800ms,那么 8 张卡在完全理想的情况下应该只用 100ms,得到 8× 加速。实测耗时 140ms,加速比降为 800/140 ≈ 5.71×,效率为 800/(8×140) = 71.4%。理想与实测之间多出的 40ms 来自通信、同步等待和负载不均。结果可以用三行对照读出:1 设备用时 800ms,加速 1×,效率 100%;8 设备理想用时 100ms,加速 8×,效率 100%;8 设备实测用时 140ms,加速 5.71×,效率 71.4%。100% 的效率只有在通信零开销、同步零等待、负载完全均衡时才成立;只要这三项中任何一项存在,实测时间就会长于理想值,效率便相应下降。
| 设备 | 同一总工作时间 | 加速 | 效率 |
|---|---|---|---|
| 1 | 800ms | 1× | 100% |
| 8 理想 | 100ms | 8× | 100% |
| 8 实测 | 140ms | 5.71× | 71.4% |
4张量并行拆一层,通信发生在层内张量并行
当瓶颈不在整个模型而在单个层时,数据并行帮不上忙:它要求每张卡都装下完整的模型副本。以 131072×262144 的 MLP 权重矩阵为例,输入维度 131072、输出维度 262144(约 344 亿个参数),仅这一层的权重按 FP16 存储就约需 68.7GB,单卡放不下,更不用说同时保留它的梯度和激活。张量并行(TP)解决的就是这类问题:把同一层的矩阵或注意力头分给多张卡共同计算,而不是每张卡复制一份。
列并行是 TP 的基本切法:把输出列分给各设备,输入在设备间共享,每张卡只计算自己负责的那部分输出列。以 131072×262144 矩阵为例,若把 262144 个输出列分到 n 张卡,每张卡只需存放并计算 131072×(262144/n) 的局部矩阵。紧接其后的层可以采用行并行的形式:它消费上一层的局部特征,最后用 all-reduce 汇总各卡的部分和,拼回与单卡计算语义一致的层结果。这里的 all-reduce 属于 collective,即多设备共同参与的集合通信;同一服务器内,这类通信通常走 NVLink 这类高速 GPU 互连。通信之所以几乎每层都要发生,是因为层被切开了:层的输出分布在多张卡上,下一层要么需要完整输入、要么需要部分和,于是每经过一层通常就要一次集合通信来恢复语义。
TP 的输入是同一批激活和被分片的层权重,输出是经过集合通信拼回的、与单卡计算等价的层结果。它确实降低了每卡权重与部分激活的占用,代价是通信极其频繁,因此对小 batch、慢链路和跨节点延迟格外敏感:每层通信的开销必须摊到足够多的计算上才划算。并行度数也不能无限提高,度数太高时每张卡分到的局部矩阵太小,GPU 利用率反而下降。工程上靠算子融合和通信与计算的异步重叠来隐藏这部分开销。三条轴的取舍可以对照着看:
| 并行轴 | 拆什么 | 通信频率 | 适合 |
|---|---|---|---|
| 数据并行 | 样本 | 每次梯度同步 | 模型单卡可放或状态可分片 |
| 张量并行 | 单层矩阵/注意力头 | 几乎每层 | 单层太大、高速互连 |
| 序列并行 | 序列维激活 | 与张量并行配合 | 长序列激活 |
一个实用的判据是:如果加了设备后每卡内存下降了,但 tokens/s/GPU 也下降了,通常说明通信开销或过小的局部算子抵消了收益。张量并行更适合高速互连的域,不应仅仅为了“多用几张卡”就把它扩展到慢网络上。
| 轴 | 拆什么 | 通信频率 | 适合 |
|---|---|---|---|
| 数据并行 | 样本 | 每次梯度同步 | 模型单卡可放/状态可分片 |
| 张量并行 | 单层矩阵/头 | 几乎每层 | 单层太大、高速互连 |
| 序列并行 | 序列维激活 | 与 TP 配合 | 长序列激活 |
5流水线并行拆层,气泡由 micro-batch 填充流水线
流水线并行针对的是“模型层数太多、整份模型装不进单卡”的情形。它把连续的模型层分成 p 个阶段,每个阶段由若干设备负责;一个训练 batch 再拆成 m 个 micro-batch,依次流过各阶段完成前向与反向。一个 micro-batch 在阶段间传递时,下游阶段在等数据、上游阶段在等梯度,等待而没有有效计算的这段时间称为气泡。气泡是流水线并行的核心代价,问题在于它占多大比例。
在 GPipe 这类最简单的同步调度下,可以给出理想化的估计:效率近似 m/(m+p−1),气泡比例近似 (p−1)/(m+p−1)。以 p = 4 个阶段、m = 8 个 micro-batch 为例,效率约为 8/(8+4−1) = 8/11 ≈ 72.7%,对应气泡约 27.3%。公式的含义很直接:f气泡 是理想化的空闲时间比例,p 是流水线阶段数,m 是每个批次拆出的 micro-batch 数;分母 m+p−1 对应整条流水线从空载到完成全部 micro-batch 的等效时间单位数,分子 p−1 对应流水线填充与排空阶段的空闲单位数。这个式子只适用于各阶段成本相近的简化调度,不包含通信开销、前向与反向的计算差异以及交错调度带来的变化,只能作为估算起点。
减少气泡的直接手段是增大 m,但这并非免费:micro-batch 变多意味着更多中间激活驻留显存,调度开销增大,而每个 micro-batch 过小时内核利用效率又会下降。1F1B 调度让各阶段在稳定状态下交替执行一次前向和一次反向,从而降低峰值激活占用,是控制这一代价的常用选择。此外,阶段之间算力不均时,最慢的阶段决定整条流水线的节拍,切分阶段应按层的计算成本而非层数平均分配。流水线还存在状态版本问题:异步或交错调度下,不同 micro-batch 可能使用不同版本的权重,此时必须明确一致性语义,否则同一批数据会与不同时刻的模型混合计算。
6ZeRO / FSDP 分片数据并行冗余状态状态分片
数据并行的朴素实现有一个明显的冗余:每个 rank 都保存完整的优化器状态、梯度和参数副本,而它们本该是同一份。ZeRO 与 FSDP 的思路是让数据并行 rank 不再各自长期持有完整的训练状态,而是按 rank 分摊。分片是渐进的:ZeRO-1 分片优化器状态,ZeRO-2 再分片梯度,ZeRO-3 与 FSDP 进一步分片参数本身。
各阶段每卡保存的内容可以对照着看:普通数据并行对优化器状态、梯度、参数全部复制;ZeRO-1 优化器状态分片、梯度与参数仍复制;ZeRO-2 优化器状态与梯度都分片、参数仍复制;ZeRO-3/FSDP 三者都分片,参数按需聚合。分片不是静态切分之后就各算各的:计算某一层之前,各卡用 all-gather 从其他卡收集本层所需的参数,凑齐后完成该层计算;反向之后,用 reduce-scatter 对各卡梯度求和,并只把不同分片留在各卡,使每卡长期持有的仍只是自己负责的那部分状态。
| 阶段 | 优化器 | 梯度 | 参数 |
|---|---|---|---|
| 普通 DP | 复制 | 复制 | 复制 |
| ZeRO-1 | 分片 | 复制 | 复制 |
| ZeRO-2 | 分片 | 分片 | 复制 |
| ZeRO-3/FSDP | 分片 | 分片 | 分片/按需聚合 |
从外部看,输入是按 rank 分片的参数与状态,输出仍应等价于一次同步数据并行更新:分片改变的是存储布局和通信模式,而不是更新语义。参数 all-gather 的粒度、预取与重计算需要与网络带宽和层大小匹配,粒度太大则单次通信重,粒度太小则通信次数多。offload 把部分状态移到 CPU 内存或 NVMe 存储,能容纳更大的模型,却可能被搬运带宽限制。每卡峰值内存下降而集合通信占比上升,正是用通信换内存的预期结果;真正的检验在部署弹性上——若从检查点恢复或改变 world size 时无法重新分片,这套方案在弹性伸缩上仍不合格。
| 阶段 | 优化器 | 梯度 | 参数 |
|---|---|---|---|
| 普通 DP | 复制 | 复制 | 复制 |
| ZeRO-1 | 分片 | 复制 | 复制 |
| ZeRO-2 | 分片 | 分片 | 复制 |
| ZeRO-3/FSDP | 分片 | 分片 | 分片/按需聚合 |
73D 并行是按拓扑组合,而非随意相乘组合
单轴并行解决不了大模型的全部约束,实际部署常把数据并行、张量并行和流水线并行组合起来,称为 3D 并行。以 64 张 GPU 为例,配置 DP=8、TP=4、PP=2,即 8 个数据并行副本、每层张量并行度 4、流水线 2 个阶段,总设备数 8×4×2=64。三个轴的通信特征完全不同:TP 组内几乎每层都要通信,应优先放在同一节点内的高速互连上;PP 只在相邻阶段之间传递激活;DP 组同步梯度或状态,频率相对低但单次数据量大。若模型还使用专家并行,会加入 all-to-all 通信,每个设备同时向多方发送 token 并从多方接收。因此 rank 到设备的映射必须感知拓扑:把 TP 组跨在慢网络上,等于让最频繁的通信走最差的链路。
这些组合并不改变数据并行的计数规则。全局批量只由数据并行副本数、局部 micro-batch 和梯度累积步数决定:B全局 = B微批 × a × N数据并行,其中 B微批 是每个数据并行副本一次处理的样本数,a 是梯度累积步数,N数据并行 是处理不同样本的数据副本数。公式里不乘张量并行度或流水线并行度,因为 TP 和 PP 只是共同完成同一副本的模型计算,并不增加这次更新看过的样本。把 64 张卡全部乘进 batch,等于把 TP 的 4 和 PP 的 2 误当成独立的样本维度,会把全局批量算大 8 倍或更多,随后学习率和训练轨迹都会按错误的批量被设定,实验结论自然失真。
配置的输入是设备拓扑、三种并行度和 batch 配置,输出是 rank 分组、通信域以及不变的优化器语义:无论 TP、PP 怎么切,一次优化步消费的样本数只由数据并行维度决定。
8检查点与故障恢复是规模化的一部分可靠性
训练规模扩大后,故障从偶发事件变成必须按常规处理的情形:训练两周后某个节点掉线,要保证的不是“祈祷它别发生”,而是“发生后不从头再来、也不恢复出错”。检查点的内容因此远比“存一份权重”宽泛:模型、优化器、学习率调度器、随机数状态、数据游标、损失缩放和并行元数据都要落盘,缺了任何一项,恢复后的训练都会从一个微妙错位的地方继续。
分片训练对检查点有额外要求。ZeRO/FSDP 这类方案下,状态按 rank 分片存放,检查点要么支持在不同 world size 上重新分片,要么明确声明做不到;写入过程采用临时产物加原子清单的方式,避免一半 rank 写成功、一半失败后留下损坏的快照。检查点间隔是写入开销与丢失计算量之间的取舍:存得越密,故障后要重算的越少,但训练被频繁打断;而设备数越多,整体故障率越高,两者共同决定间隔该取多密。验证检查点也不能只看文件是否存在,而要定期做真实的恢复演练,比较恢复后的 loss 曲线与样本序列是否与原轨迹衔接。
弹性训练还引入成员变化问题:节点加入或退出后,batch 划分和随机流都必须重新调整。数据“恰好一次”的语义很难实现——恢复时可能重复或跳过样本,因此 sampler 状态和全局 step 必须持久化,使重复与跳过可控、可解释,而不是随机发生。
9扩展评测要区分强扩展和弱扩展评测
评估分布式训练的效果,先要确定自己在测什么。GPU 翻倍时,如果保持总问题大小不变,测的是强扩展:固定总模型与总工作,看增加设备后同一任务的加速有多少。如果保持每卡工作量不变,测的是弱扩展:总问题随设备数同步增大,看单位效率能否维持。两者回答的问题不同——强扩展关心“同样一件事能不能更快做完”,弱扩展关心“规模做大后每张卡是否还同样划算”。
评测的输入是不同设备数下保持一致的实验协议,输出不应只是吞吐,而是一份联合报告:速度、效率、内存、功耗、恢复与最终质量。吞吐指标要分清口径:tokens/s 是集群每秒处理的 token 数,tokens/s/GPU 是平摊到每卡的效率;MFU 是实际有效模型计算相对硬件理论峰值的比例;collective 占比是集合通信占一个 step 时间的比例。数据等待、气泡、重计算、显存峰值、功耗和恢复时间同样要报告,任何一项都可能让“吞吐好看”的训练在实际中不可持续。
与此同时必须验证训练语义:global batch、学习率、token 数、数据顺序、dropout 随机流、loss scaling 与归约精度都要对齐。吞吐提高但最终质量下降,不是成功的扩展。反过来,MFU 低只说明硬件没有充分用于模型计算,不能单独判断瓶颈来自数据、通信还是过小的算子。定位瓶颈要用 trace 把一个 step 拆成数据读取、前向、反向、集合通信、优化器和检查点,把各部分的时间占比与等待依赖对应起来;GPU 利用率低只是症状,trace 里的等待关系才指向病因。
11把因果链连起来综合
把整条路径串起来,分布式训练从“能否放下”的问题出发,最终要落到可验证的实践上。第一步是核算参数、优化器状态与激活的峰值占用,这一步决定哪些轴真正必要——账没算清就选并行策略,等于在不知道瓶颈的情况下开药方。第二步据此选择 DP、TP、PP 与分片的组合:状态冗余大用 ZeRO/FSDP,单层太大用张量并行,层序列太长用流水线并行,只有样本维度才交给数据并行。第三步按设备拓扑映射通信域,把最频繁的 TP 通信放在高速互连内,避免把几乎每层都发生的集合通信甩给慢网络。第四步固定全局批量与数值语义:B全局 = B微批 × a × N数据并行,只由数据并行维度决定;学习率、token 数、数据顺序、dropout 随机流、loss scaling 与归约精度都要与基线对齐,否则性能结论不可比。第五步用 trace 把 step 拆开,找到数据读取、前向、反向、集合通信、优化器和检查点之间的等待依赖,通过计算与通信的异步重叠把空闲填满。最后一步是真实恢复与联合验收:演练故障恢复、比较恢复前后的 loss 与样本序列,并以质量与效率的联合结果——而非单独的吞吐——判定这次扩展是否成立。任何一步被跳过,前面各步留下的因果都会在下游被放大。
- Megatron-LM:张量模型并行
- ZeRO:训练状态分片
- GPipe:流水线并行与 micro-batch
- PyTorch FSDP:全分片数据并行实现语义