最早接触分布式训练的时候,我的理解特别朴素:把模型均匀拆到几张显卡上,算完梯度再同步一下,不就完事了。直到自己动手跑一个7B级别模型的训练,看着显存被瞬间吃光、日志里频繁出现卡死和OOM,才发现这个“朴素理解”几乎没有一步是对的。
这篇笔记是我在系统梳理分布式训练基础理论时整理的,面向已经跑过单卡训练、正要转向多卡训练的开发者。我不会把篇幅花在安装工具这类入门操作上,而是把数据并行、张量并行、流水线并行、ZeRO分片、混合精度、显存规划这几件事的原理和取舍讲透。你会发现分布式训练的核心不是“怎么切卡”,而是“通信和显存怎么取舍”这一件事。
1. 为什么单卡跑不动大模型:先算一笔显存账
1.1 7B模型训练需要多少显存
很多人刚接触大模型时,第一个直觉是:7B参数听起来不多,模型文件也就14GB左右,一张80GB的显卡总该够了吧?这个直觉是错的,因为训练状态下模型在显存里放的东西远不止参数本身。
业内估算训练显存有一个经典公式:按每个参数16字节估。以7B模型为例,7 × 10^9 × 16 = 112GB。这意味着训练一个7B模型,理想情况下也需要超过112GB的显存,一张卡根本放不下。那16字节是怎么来的?
- 模型参数,FP16存储,2字节/参数
- 梯度,FP16存储,2字节/参数
- 优化器状态,FP32的动量、方差以及一份FP32主权重,12字节/参数
四项加起来正好是16字节。注意优化器状态才是大头,占了3/4。卡上最贵的不是模型本身,而是Adam这家“账房”在记账时产生的各种中间状态。
1.2 除了模型状态,还有两座隐形大山
模型中有一个很容易被忽略的显存消耗点:激活值(activation)。前向传播时,每一层的中间输出都要保留,反向传播时要用它们算梯度。一个batch里若包含上万token、几十层Transformer、隐藏层维度4000以上,激活值轻松吃掉几十GB。激活值的大小不仅和模型规模有关,更和序列长度、batch大小直接相关,这解释了为什么大batch训练时即使模型状态能塞下,激活值仍然可能直接撑爆显存。
还有一类开销来自框架运行时:CUDA context、通信缓冲区、内存碎片、NCCL的临时空间。经验上是数百MB到数GB不等,单卡跑小模型时没感觉,一旦卡上显存被模型填到90%以上,这些开销就会成为压死骆驼的最后一根稻草。
算完这笔账之后,核心结论就出来了:单卡跑不动大模型不是因为算力不够,而是因为显存不够。分布式训练首先解决的是显存容纳问题,其次才是算力扩展问题。理解了这一点,后面几种并行策略的动机就很清晰了。
2. 数据并行:把同一份模型复制到每张卡上
2.1 梯度AllReduce的直觉理解
数据并行是最早出现、也最容易理解的一种并行方式。它的做法是:每张卡上都放一份完整的模型副本,把训练数据切成N份分给N张卡,每张卡独立做前向和反向,算出一份梯度,然后把所有卡的梯度做一次全局聚合,每张卡拿到聚合后的梯度,再各自更新参数。
为什么可以这样?因为梯度本质上是对一批样本的误差方向的统计。如果只有一张卡,它算的是一个batch中全部样本的平均梯度;现在有N张卡,每张卡算的是1/N的样本的平均梯度,只要把N份梯度求平均,数学上就等价于整批样本的平均梯度。这就是数据并行能保证训练结果和单卡一致的根本原因。
这里的全局聚合操作就是AllReduce。用生活场景类比一下:如果几个人分别记了一天的账,对不上账的时候,每个人都要把自己的账本交给别人汇总一遍,最终每人都拿到一份完全相同的汇总结果。AllReduce做的就是这件事。
2.2 通信量其实比想象中大得多
数据并行最大的隐藏成本是通信。每次训练迭代,每张卡都要把完整梯度发给别人,同时收别人的完整梯度。以7B模型为例,FP16梯度是14GB,在Ring AllReduce这种高效的算法下,每个GPU每步大约要发送和接收各14GB左右的数据,合计接近28GB的通信量。
这意味着数据并行的扩展性取决于“计算时间”和“通信时间”的比值。如果单卡算一个batch要10秒,通信只占1秒,那并行效率还算理想;但如果单卡算一个batch只要0.5秒,通信要1秒,那么多卡并行反而可能比单卡还慢。这也是为什么数据并行在早期“显卡少、模型小”的时代很好用,到了大模型时代却必须配合下面几种并行一起用。
2.3 从DDP到FSDP:数据并行也在进化
传统的数据并行实现(每卡完整副本+全梯度AllReduce)有两大痛点:一是每张卡都要塞下完整模型和优化器状态,模型一大就放不下;二是每一步的梯度AllReduce通信量固定为模型规模的倍数,模型越大越吃力。
后来的FSDP(全分片数据并行)做了关键改进:既然每张卡都存完整参数很浪费,不如把参数、梯度、优化器状态都按卡切分,每张卡只存一部分,等需要计算某层的参数时再临时从其他卡取回来。这等于用“通信换显存”,后面会详细展开。
我在实践中的体会是:如果一个模型在单卡上刚好能塞进显存,只是想加快训练速度,优先考虑传统数据并行就好,实现简单、数学等价、调试成本低。只有当模型在单卡上塞不下了,才需要FSDP这类分片方案。
3. 张量并行:把一张矩阵按行和列切开
3.1 按列切分与按行切分的经典组合
数据并行解决不了“单卡放不下模型”的问题,那就必须把模型本身切开,让模型的不同部分放在不同卡上。按层切分是后面要讲的流水线并行,而在层内部把矩阵运算切开,就是张量并行。
Transformer里最常见的大矩阵运算是Y = XW,X是激活矩阵,W是权重矩阵。张量并行最直觉的做法是把W按列切成W1和W2,分别放在两张卡上,两张卡同时算XW1和XW2,然后拼接结果。关键问题是:拼接完输出之后,下一步是什么?如果下一步是残差连接和LayerNorm,那需要把完整输出拼回来才能算,这里就产生了一次跨卡通信。
更精细的做法是让矩阵乘法本身按行和列同时切分,比如把一个输出维度为4096的线性层切成2×2的网格,用4张卡配合,每张卡算其中一块子矩阵,最后一轮用AllReduce把子矩阵的部分和汇总成完整结果。这种二维切分在工程上广泛使用,因为它在保证计算分摊的同时,把通信量压到了只有列切分方案的1/2左右。
3.2 张量并行消耗的是通信带宽,省的是显存
张量并行看似只是一种“拆分计算”的技巧,对通信却极其敏感。每个Transformer块都有若干处需要跨卡同步点在attention输出、MLP输出、LayerNorm之前。也就是说,每过一个block,就要做几次全卡级别的AllReduce。模型层数越多,这种同步的次数就越多。
这直接决定了张量并行的适用边界:它内部节点之间的通信频率太高,只能依赖NVLink这种几百GB/s的卡间高速互联,跨节点做张量并行是极不划算的。实践中几乎都把张量并行的规模限制在单节点内的GPU数量,一般就是8卡或16卡,很少跨节点扩。
3.3 为什么说TP是“高带宽亲儿子”
如果读者和我一样买过一些机子做过小规模测试,建议务必先搞清楚自己集群的拓扑。同样是8张卡,有的机器全部卡都在一个NVSwitch上,任意两卡通信带宽都很高;有的是几组卡各连各的,跨组通信要走PCIe甚至网卡,带宽差一个数量级。
张量并行选多大维度,主要不是看有多少卡,而是看卡之间的通信带宽撑不撑得住高频AllReduce。哪怕有32张卡,如果它们分布在多台机器上、靠普通以太网连接,也不适合全用TP;相反,如果只有4张卡但都在同一个NVLink域里,那先上TP就是很自然的选择。
4. 流水线并行:把Transformer一层层分到不同机器
4.1 朴素的按层切割为什么低效
流水线并行的想法更直接:把模型的若干层切成几段,每段放在一张卡上,数据按顺序从第一段流到最后一段。在数学上,这种方式也能放下大模型,而且每张卡只需要负责一小部分层的存储和计算。
但它有一个“天然低效”的问题:气泡。假设模型切成4段,第一张卡算完前几层把结果传给第二张卡,第二张卡开始算,此时第一张卡已经闲着没活干了。依次类推,数据在4张卡上“接力跑”,大部分时间只有一张卡在工作,其他卡在空等。如果不做任何优化,4卡流水线并行的加速比甚至可能低于1,比单卡还慢。
4.2 MicroBatch与1F1B调度为什么能救场
解决气泡的办法是把一个大的训练batch再切成更小的microBatch,然后让这些microBatch像流水线的工件一样,一个接一个地流经各段。前一个microBatch还在后段计算时,前段已经开始处理下一个microBatch,空转时间被大幅压缩。
气泡率有一个近似公式:(p-1)/(m+p-1),p是流水线段数,m是microBatch数量。直观感受一下:如果p=4、m=4,气泡率约43%;如果m=16,气泡率降到了约16%。所以实践中microBatch数量通常要数倍于流水线段数。
在调度方式上,最常用的是1F1B调度:每处理完一个microBatch的前向,立刻处理它的反向,而不是像早期方案那样等所有microBatch前向都结束再统一反向。这样做最大的好处是显存占用大幅下降,因为反向计算结束后,对应microBatch的激活值就可以释放,不需要把所有microBatch的激活一起留在显存里。
4.3 气泡、显存与带宽的三方博弈
流水线并行不是没有代价。第一,切分点附近需要传输大量的中间激活,切分越碎、传输次数越多,通信开销越高。第二,microBatch太小时每个microBatch的计算量偏少,GPU利用率下降;microBatch太大时气泡率又降不下去。第三,每个stage上模型的层数不同会造成负载不均,切分时要考虑每层计算量的差异,不能简单地“层数除以卡数”。
实践中我的经验是:先把流水线段数设定为机器数的倍数,保证每一台机器内部有连续的若干层;再用小规模实验去测不同microBatch数量下的吞吐,找一个平台期。流水线并行更像“容量解决方案”,它主要解决模型装不下的问题,而不是追求极致加速,所以对吞吐的微调优先级应该放在参数切分正确性之后。
5. 从ZeRO到FSDP:用通信换显存的极限手段
5.1 三种冗余:参数、梯度、优化器状态
回到数据并行那张图:每张卡都有完整参数、完整梯度、完整优化器状态。当模型大到单卡放不下时,这种冗余直接导致无法训练。但仔细想一下,这些冗余真的有必要吗?
- 参数:反向传播时需要用到整层参数,但同一时间点只用到少数层的参数,没有必要让所有层常驻显存
- 梯度:最终梯度是要在所有卡之间求和的,不是每张卡都需要保存完整梯度
- 优化器状态:Adam更新时每个参数的动量和方差只由对应参数决定,可以按参数分片存储
ZeRO的核心思想就是把这三类状态分别做分片。它有三个阶段,逐级递进:第一阶段只分片优化器状态,第二阶段连梯度也分片,第三阶段连参数也分片。每前进一步,省下的显存都更可观,但通信开销也随之增加。
5.2 每个阶段分别省掉了什么
ZeRO-1对应的是最克制的方案:把优化器状态按卡分片。训练时仍需要所有卡的完整梯度做AllReduce,但每张卡只负责更新自己那部分参数对应的动量状态,更新完再广播给所有卡。这个阶段就能把Adam那12字节/参数的显存压力降下来,是最划算的一步。
ZeRO-2更进一步,把“算完整梯度”也切碎了。原本的AllReduce被替换成Reduce-Scatter和All-Gather的组合,每张卡最终只拿到完整梯度的1/N。这时每张卡在训练中需要保存的梯度只有原来的1/N,显存进一步节省。
ZeRO-3则是把参数也分片存。前向和反向时用到哪一层,就把哪一层的参数临时All-Gather回来,用完就丢。从显存角度看,这是最极致的分片,单张卡几乎只需要存模型参数的1/N,配合优化器和梯度分片后,理论显存需求可以降到原来的1/N甚至更低。
5.3 为什么ZeRO-3不总是最优解
很多新人看到ZeRO-3的显存收益后,觉得所有场景都应该用它。但实际训练大模型时的工程判断恰恰相反:ZeRO-3引入了更频繁的通信。本来每层参数只需要存在本地,现在每次前向和反向都要先All-Gather,这意味着每个batch都会发生几十上百次全局通信,对网络延迟和带宽的考验远超ZeRO-1、ZeRO-2。
正因为如此,业界在超大模型训练上几乎都不是单用ZeRO-3,而是把数据并行、张量并行、流水线并行进行组合,让通信尽量发生在节点内的高带宽互联上。ZeRO-3和FSDP更适合这样的场景:模型规模还没大到需要复杂并行策略,或者是主流的训练框架里实现的好、开箱即用。
我个人的使用建议是:如果只是想用现有框架快速把模型跑起来、卡的数量不多,优先选FSDP这类现成方案;如果做到数百卡以上、追求极致效率,那手动规划TP、PP、DP的组合才是真正要下的功夫。
6. 混合精度与激活值:显存管理里最不起眼的大头
6.1 为什么优化器状态比模型本身还占显存
前面算过,模型参数只有2字节/参数,Adam优化器状态却要12字节/参数。这个差距来自混合精度训练的经典设计:模型参数在每次前向和反向时用FP16计算,但更新时必须在FP32的“主权重”上进行,更新完再转回FP16。为什么呢?因为FP16只有大约3位有效十进制数字,多次累加更新后误差会累积;FP32则稳妥得多。Adam本身又额外保存每个参数的一阶动量和二阶动量各一份,都是FP32,所以优化器状态膨胀得非常快。
这就造成一个看似矛盾的现象:一个14GB的FP16模型,训练时配齐优化器状态后需要112GB,其中模型本身还不到1/8。理解了这一点,在显存不够时就知道该往哪优化——优先处理优化器状态,比如换用Adafactor这类无动量优化器或做优化器状态分片,而不是去压缩模型参数量。
6.2 激活值Checkpointing的取舍逻辑
激活值之所以常被忽略,是因为它们和“模型大小”没有直接关系,只和batch大小、序列长度、层数、隐藏维度有关。但它经常是压垮显存的最后一根稻草。
解决激活值显存最通用的技术是Activation Checkpointing,也叫激活重计算。它的思路很反直觉:前向传播时,不保存每一层的激活值,只保存各层的输入;反向传播要用某层激活时,临时把该层的前向重新算一遍。这样显存占用从“所有层激活都保留”降到“只保留少量输入”,大幅下降,代价是大约多算30%-40%的前向FLOPs。
工程上几乎不需要自己实现这个机制,主流框架都有开关,但在开启前要想清楚:如果模型是计算密集型且GPU算力有富余,重计算的开销几乎无感;如果模型本身就是通信瓶颈,那就需要衡量了。
6.3 混合精度训练的两个容易忽略的细节
一个是Loss Scaling。FP16能表示的数值范围很窄,如果梯度太小,可能在反向传播时直接变成0。处理办法是给损失函数乘一个大数,比如1024或动态调整的Scale值,让梯度在FP16范围内保持可表示,完成更新后再把scale加回去。很多分布式训练里的“loss突然变0”问题,都跟Loss Scaling处理不当有关。
另一个是BF16和FP16选谁。BF16保留了更多指数位,数值范围大,不会因为梯度过小而直接下溢,但尾数位少,精度却明显低于FP16。某些模型对精确度敏感时,BF16训练效果不如FP16混合精度;而另一些大规模场景下BF16又表现更稳定。碰到训练loss震荡,可以先试试切回FP16混合精度或降低batch对比。这些都不是教科书里会强调的,但排查起来很常见。
7. 多卡训练的工程选型与3D并行排布
7.1 3D并行怎么排卡
把前面几种技术组合起来,就构成了业界常说的高维并行:数据并行(DP)、张量并行(TP)、流水线并行(PP)。举一个8台机器共64卡的排布实例:每台机器内部8张卡先做TP,因为这一步对节点内带宽要求最高;然后按机器划分PP,比如把模型切给4组机器,每组机器负责连续若干层;最后在多组机器之间叠加DP,通过梯度同步把数据吞吐顶上去。
“先TP后PP再DP”的顺序几乎是通用经验。原因是通信敏感度从高到低递减:TP需要最频繁的同步,必须在最快互连范围内;PP的同步只发生在切分点,频率低很多;DP的同步每步只要一次AllReduce,虽然总量大但对延迟宽容度最高。排卡顺序错了,性能可能差好几倍。
7.2 框架选型与集群上的现实约束
不同框架对不同并行策略的支持成熟度差异很大。如果只谈通用方案,主流的分布式数据并行接口、全分片数据并行、分布式的ZeRO系列,以及一些专门为超大模型预训练设计的并行框架,各有各的侧重点。
工程选型的现实约束往往不是算法能力强弱,而是硬件条件。跨节点的网络如果是高带宽低延时的InfiniBand或RoCE等支持RDMA的网络,ZeRO-3跑起来很顺畅;如果只是千兆以太网,那每步通信都可能变成瓶颈,只能优先考虑尽量扩大TP和PP规模、减少跨节点DP通信。所以选型前第一步应该是摸清自己的网络拓扑,而不是先看哪个框架参数更炫。
7.3 通信计算重叠,比减少通信更重要
减少通信量只是手段,更高级的手段是“隐藏通信”。分布式训练中,反向传播算梯度是逐层进行的,某些层的梯度算完后,通信操作可以立刻在后台启动,同时其他层的前向计算继续跑。只要通信时间不大于计算时间,通信就可以被完全“藏”在计算里,训练吞吐几乎不受影响。
一个典型事件是NCCL的异步传播和内核融合:把多个小的通信合并成一个大通信,减少启动延迟。新手在多卡调优时,经常盯着通信量看,但实际吞吐瓶颈常常是“通信等待计算”或“计算等待通信”的串行化问题。用profiler观察每个step的时间线,如果GPU空闲占比高,就要考虑调整并行策略或者开启关键通信的异步化。
8. 从理论到实操:我在分布式训练里踩过的坑
8.1 先搞清OOM是静态还是动态
分布式训练里的OOM常常让人摸不着头脑,因为同样的配置在8卡上跑没问题,换成16卡就崩了。排查的第一步是区分:是“模型静态加载时就超出显存”,还是“前向过程中激活值把显存吃满了”。做法很简单——用一个很小的batch先跑几层,如果小batch能过、大batch不能过,那多半是激活值超限;如果小batch也不能启动,那多半是模型和优化器状态本身超出了容量。
不同问题处理方式完全不一样:前者可以用激活重计算、减少batch、降精度;后者只能上并行策略或分片方案。如果一上来就调整batch大小,方向错了,浪费时间。
8.2 卡死问题大概率是集合通信挂起
分布式训练中常见的另一种“症状”是:日志停在一个地方不动,GPU利用率降到0。这种时候基本可以判定是集合通信在等待某个rank返回,常见原因有:参数没对齐导致某张卡在跑while循环等数据、通信超时阈值太小、某个节点掉线导致NCCL不断重试。
我的排查顺序是:先看日志里最后一条打印的是哪个阶段,确认是不是卡在AllReduce;再用一个小脚本只做梯度同步,排除业务逻辑干扰;最后检查各个节点之间的网络是否有丢包。多数卡死问题不是模型代码的错,而是集群通信环境的问题。
8.3 随机种子与可复现性,比想象中重要
数据并行虽然数学上等价于单卡大batch,但分布式场景下数据切分顺序、随机种子、硬件执行顺序都会影响结果。同一份代码,两次训练出来的loss曲线可能有细微差异,这未必是bug。
为了排查的方便,我习惯把下面三样东西固定下来:训练数据的shuffle种子、模型初始化种子、DDP的随机种子。同时每次实验把完整配置存档,包括并行策略、batch大小、学习率、精度设置、网络拓扑。大模型训练一轮很贵,配置丢失才是最大的浪费。
另外有一点很实用:换并行策略之后,不要直接跑全量训练,先用较小的模型、较少的数据做一次“对照训练”,把loss曲线和单卡小batch的曲线放在一起看。如果趋势不一致,大概率是并行实现或通信逻辑有bug,而不是模型的锅。
8.4 我个人的调试验证顺序
最后分享一个这几年摸索出来的跑大规模训练之前的必做流程:
- 单卡小模型小batch跑通,确认模型代码本身没问题
- 节点内TP=8,小模型,验证张量并行切分和拼接逻辑
- 两台机器之间P2P测试,用NCCL自带的带宽测试工具量一下真实通信带宽
- 小规模完整并行配置,跑50步,观察GPU利用率和loss走势
- 开启激活重计算和混合精度,再跑50步,对比显存和吞吐
- 最后才上全量配置,同时全程打开profiler采样
这套流程看起来繁琐,但从“卡死半小时才发现问题”和“提前10分钟发现断连”之间,差的不是半小时,而是一整天的心态。
分布式训练的基础理论,说穿了就是“显存不够怎么办”和“通信太贵怎么办”两件事。你每次在框架里调整一个并行参数,其实都是在权衡这两件事。把这套权衡逻辑想明白,再去看框架源码、看技术社区的优化方案,会顺畅很多。