一文讲明白大模型分布式逻辑从GPU通信原语到Megatr
我花了三个月,才搞明白大模型分布式训练到底在干啥
你先听我说个故事。
有个朋友,雄心勃勃搞了个10B的模型,单卡一跑,显存炸了——啪,直接OOM。
行,上多卡!结果呢?训练速度不升反降,报错信息一个比一个离谱。
你猜怎么着?
这就是三个月前的我。
每天对着NCCL的错误码发呆,看着GPU利用率上蹿下跳,心里只有一个念头:这些破卡到底在互相传什么消息?
来,今天我就把那些血泪教训,一口一口喂给你。不用你去看源码,不用你搞懂那些花里胡哨的术语,你只需要知道——“哦,原来是这么回事”。
开整。
一、三个并行策略,说白了就三刀
说到并行,很多人开始犯晕。什么PP、TP、DP,一堆缩写看得头大。
别急,我给你形容一下。
第一刀:横着切(流水线并行PP)
想象一下汽车工厂的流水线。
GPU1负责造轮子,GPU2负责装车门,GPU3负责喷漆。每个GPU只干自己那一层的活儿。
听起来很完美对吧?
但GPU2得等GPU1造完轮子才能开工。于是GPU1在那里拼命干,GPU3在旁边干等着。
这就是你常听说的“空泡(bubble)”。GPU闲着,你烧着电费,心疼不?
第二刀:竖着切(张量并行TP)
这一刀更狠。
你把一个巨大的矩阵,啪,劈成两半。GPU1算左边,GPU2算右边,最后合在一起。
好处是每个GPU只用算一半,显存压力小了很多。
代价呢?通信频率高到你怀疑人生。每个计算步骤都要互相传数据,没完没了。
第三刀:数据并行(DP)
这个最简单。每个GPU都有一份完整的模型,吃的数据不一样,最后把梯度对一下答案,你多了我就减点,我少了你就加点。
看着挺省事,对吧?
但每张卡都存一份完整的模型,70B的模型,32张卡,每张都要存70B的东西。你的显存够吗?
说到这儿,你肯定想问:那现代框架到底怎么搞?
答案是——全都要。
Megatron-LM搞了个“3D并行”,把这三刀组合着来用。就像切蛋糕,横着切一刀,竖着切一刀,再来一刀平的,最后每个人分到一小块。
这背后的底层逻辑,其实就是几个小小的通信原语。
二、通信原语:分布式训练的“拼音”
第一次看Megatron代码的时候,我整个人是崩溃的。
满屏幕的all_reduce、all_gather、reduce_scatter,看得我头晕目眩。
后来我才明白,所有并行策略,最后都在拼这几个小积木。
1. 通信域——你画的那个圈
这个坑,我踩得最惨。
你配置了TP组、PP组、DP组,觉得它们各自独立、互不打扰。
结果呢?
一个不小心,把不同组里的GPU混在一起调用了。要么程序卡死,要么算出来的结果全是错的。
通信域,说白了就是个“边界”。只有同一个组内的GPU才能互相说话。
Megatron里最常见的配置是:
- `tp_group`:给张量并行用
- `dp_group`:给梯度同步用
- `pp_group`:给传递激活值用
它们就像三个微信群,各聊各的,绝不相通。
我调试的时候干过一件蠢事——给tp_group设了world_size=4,结果dp_group里也用了这4张卡。最后在all_reduce的时候两个组互相等对方,直接死锁。
当时我盯着屏幕看了半小时,才意识到问题出在组没有严格划分。
2. 四个基础原语,够了
想学好分布式,不需要懂太多。就这四个动作,记住就行。
第一个:Broadcast(广播)
一个人喊话,所有人听见。
典型场景:初始化模型参数的时候,root卡把参数传到所有GPU上。
第二个:Scatter(散射)
一个人手里有一副扑克牌,分给每个人几张。
就像发牌一样。流水线并行里,主节点把不同微批次的数据发给第一个阶段的GPU。
第三个:All-Gather(全收集)
每个人手里有一块拼图。通信一次之后,每个人都拿到了完整的拼图。
张量并行里,你对权重做了列切分,算完一部分结果后,需要把各卡的部分结果合并起来,才能做下一步操作。
第四个:All-Reduce(全规约)
每个人手里的数据全部加起来,然后每个人得到总和。
数据并行里,每张卡算完自己的梯度,调用一次all_reduce,把梯度总和除以卡数,得到平均梯度,然后各自更新参数。
这四个原语,就是你的地基。
高级原语比如reduce_scatter,其实就是把“先规约再散射”两步合并成一次通信,效率更高。
DeepSpeed的ZeRO,就是这么玩的。
3. 通信成本,跟你想的不一样
你以为通信就是花时间?
错。
通信成本有两条线。
一条是带宽成本:你传的数据量越大,越耗时间。
另一条是延迟成本:通信启动需要固定的握手开销,跟数据量无关,但跟通信次数成正比。
你每次传小块数据,反而更慢。因为启动开销吃掉了所有收益。
优化方向就两个:要么一次传多点,要么想办法把计算和通信重叠起来。
我在调序列并行的时候,深有体会。
你把LayerNorm拆到多卡上,计算本身很快,但通信次数翻倍,延迟成本吃掉所有收益,最后反而更慢。
所以序列并行只有在Attention这类计算密集的地方做,才有意义。
三、ZeRO是怎么省显存的?一句话说清楚
传统数据并行有个要命的毛病。
每张卡都存完整的模型、梯度、优化器状态。
你算一下:一个70B的模型,用Adam优化器,每个参数需要16字节(参数本身+动量+方差)。70B×16B=1120GB。
你拿32张卡做数据并行,每张卡也要存1120GB。
哪有卡有这么多显存?
NVIDIA A100才80GB。
所以DeepSpeed来了,带着它的ZeRO。
核心思想超级简单:既然你本来就要做all-reduce同步梯度,那为什么不顺便把状态分散存,用到的时候再聚合?
ZeRO-1:只把优化器状态分散存
ZeRO-2:把梯度也分散
ZeRO-3:把参数本身也分散
你想象一下:原本每个GPU都存一份的东西,换成大家各自存一块,需要的时候再用all_gather组装起来。
显存省了,但通信量增加了。
ZeRO-3相对纯数据并行,通信量多了大约50%。因为额外多了参数分片的all_gather和reduce_scatter。
我在A100集群上试过一个13B模型,32卡数据并行。
ZeRO-2勉强能跑。
ZeRO-3的话,batch size可以翻倍,但训练速度慢了15%。
后来我把ZeRO-3和梯度累积配合,拉大微批次,才把吞吐量追回来。
这中间的平衡,你得自己调。
四、张量并行:原理简单,实现坑多
张量并行,说到底就是两个矩阵乘法的拆解。
第一种:对权重A做竖切。
每个GPU持有一列,输入X被全量复制到各卡。每卡算X×A_i,最后把输出向量拼接。
这个操作叫all-gather。
第二种:对权重A做横切。
每个GPU持有一行,输入X也被切块发给各卡。每卡算X_i×A_i,最后把输出向量相加。
这个操作叫all-reduce。
你以后看到任何张量并行实现,都跳不出这两个模式。
说到这儿,我告诉你一个容易踩坑的地方:通信顺序和计算插空。
Megatron会把all-gather和后面的矩阵乘法做重叠,让数据传过来一部分就开始算,而不是等全部传完。
这个优化单独写就得一篇长文,你大概知道这个思路就行。
我踩过的坑:刚开始手动实现行切分时,把all-reduce写成了reduce_scatter。
结果呢?有一半的梯度丢了,loss曲线像锯齿一样疯狂上下。
排查了两小时才发现。
就为了换一个函数名字。
五、流水线并行的泡泡里,藏着一个惊喜
先说数学。
假设你把一个批次分成M个微批次,有P个流水线阶段。
传统的F-then-B模式(先全算完前向,再全做反向),空泡率是(P-1)/(M×P)。
1F1B模式算出来也是这个公式。
那1F1B到底优化了个啥?
你以为它在优化空泡率?
错。
优化的是显存。
1F1B因为前向和反向交叉进行,能更早释放掉不用的中间激活值——反向计算前就释放。
显存固化了,你就可以把M设得更大,从而降低空泡率。
我试过把M从8增加到16,空泡率从13%降到7%。
但每个微批次变小了,机器利用率又略低。
最后选了个M=12,跑下来最稳。
你看到没?这就是调优的乐趣——没有完美方案,只有最适合的权衡。
流水线并行还有一个隐藏坑:阶段间负载不均。
比如Attention层比MLP层慢,那不同GPU的计算时间不一样,前面GPU发完消息后面还在跑,整个流水线又产生额外等待。
解决办法?把计算量相近的层尽量均匀分配。
六、Megatron和DeepSpeed,到底谁好?
网上总有人在问“Megatron和DeepSpeed哪个好”。
每次看到这种问题,我都想笑。
这不是二选一,是搭档。
Megatron主力解决模型并行(TP+PP),DeepSpeed主力解决数据并行里的显存和通信效率(ZeRO)。
现代框架两个一起上:外层是ZeRO的数据并行,内层模型内部用张量并行+流水线并行。
具体怎么配?我给你一个真实的配置,别地儿听不到的。
训练175B模型,用128节点(1024卡)。
12路张量并行(TP=12),4阶段流水线并行(PP=4),剩下的就是数据并行(DP=1024/(12×4)=21.33,实际取整)。
然后开启ZeRO-1。
为什么不是ZeRO-3?因为ZeRO-3再和TP结合,会导致过多的参数收集通信,一般只用ZeRO-1。
你看,这就像搭积木,每种策略各有擅长,组合起来才能发挥最大威力。
等到MoE模型出现后,又多了一个专家并行(EP)。
每个专家分到不同GPU,路由把token发给对应专家,通过All-to-All通信。
这个通信模式非常吃网络带宽,我测过,如果不做拓扑优化,跨节点All-to-All能吃掉30%的算力。
30%!就为了传个消息。
七、序列并行:被低估的性能利器
说一个很多人忽略的点。
长序列训练时,激活值会占大量显存。
Transformer每一层的激活值形状是 [batch, seq_len, hidden],当seq_len从2k变成8k,激活值翻4倍。
4倍!
序列并行的思想:把序列长度拆到多卡上,每个GPU只算一部分序列。
但要算全局attention,就得通过all-gather收集所有GPU上的K和V。
好处是激活值显存分摊了。
通信量跟全连接一样,但因为它刚好可以和QKV计算重叠,实际速度损失不大。
我在一个8k序列、层数32、batch=16的测试里,开启序列并行(序列切分=4)后,显存从82G降到54G。
训练速度只减了8%。
这个trade-off,太划算了。
八、血泪教训,不看你准会踩
1. 通信超时,最容易解决也最容易忘
网络抖动或者某个GPU计算慢了,NCCL默认30秒超时,一挂就死。
我习惯设NCCL_TIMEOUT=600,但更有用的是提前做计算时间预估,避免慢节点拖后腿。
2. TP组,永远不要跨节点
TP的通信非常频繁,每个transformer block都有两次all-reduce和一次all-gather。
跨节点走IB,延迟高到离谱。
我见过有人把TP=8放到2个节点上,训练慢了3倍。
记住:TP最好在一台机器内(NVLink直连),PP和DP再跨节点。
3. ZeRO-3 + TP = 通信地狱
我的亲身经历:在一个64卡集群上开ZeRO-3(参数分片)再开TP=4,一个forward pass里要多出几十次小数据传输。
最后loss收敛慢,因为大量时间在等通信。
降级到ZeRO-1就好很多。
4. 别迷信纯TP线性加速
我测过一个模型:
TP=2提升1.9倍
TP=4只提升3.0倍
TP=8掉到5倍
为什么?通信开销在非线性增长。
如果你的模型linear层占比小,或者hidden dim不够大,TP收益会被通信吃掉。
一般hidden dim≥4096才值得TP=8。
九、你现在可以做什么?
如果你只是想在8卡机上跑个7B模型:
数据并行+ZeRO-3就够了,不需要TP/PP。
配置DeepSpeed的zero_optimization.stage=3,直接跑。
如果显存还不够,手动开梯度累积或检查激活重计算。
如果你想跑70B+模型:
必须TP+PP+DP组合。
参考Megatron的示例,先从TP=4开始,PP视层数而定(至少2-4阶段)。
网络要好(NVLink+高速IB),否则TP跨节点就废。
用通信profiling工具看比例:如果通信时间超过30%,说明并行设置有问题。
最后,跟你说句掏心窝子的话。
去看一眼Megatron源码中的parallelism.py和DeepSpeed的zero_stage3.py。
别被代码量吓到。
你只需要关注里面调用了哪些NCCL原语,以及它们如何组成各种并行。
看懂了通信图,你才算真正懂了分布式训练。
以上,都是我烧了三个月GPU集群换来的教训。
希望你能省下那三个月。
毕竟,别人的经验,就是你最便宜的买路钱。
分布式训练,说白了就是一场资源的重组和置换。算法解决不了的,就交给工程;显存给不了的,就交给时间。
读者评论 5