图解大模型训练之:Megatron源码解读2,模型并行
先讲个让我“头皮发麻”的故事
你第一次看到 Megatron 的模型并行代码是什么感觉?
我是这样的:当时已经写了大半年分布式训练,自认为对数据并行、AllReduce 门儿清。结果那天有人问我:“Megatron 的模型并行到底怎么回事啊?”我第一反应:完了,又要翻车了。
不是因为它难。说真的,三句话就能讲清楚“张量并行切权重,流水线并行切层”。但真要我去翻代码,看它怎么把模型切开、塞进几十张 GPU,还要保证没错——我愣是摔了好几跤,才从坑里爬出来。
所以今天这篇,我不搞从头铺到尾的教程了。直接回答你心里最痒的几个问题。你带着问题来,看完就能回去翻代码,少走弯路。相信我,看完你会拍大腿:原来这么回事!
第一关:模型并行到底并行了个啥?你绝对猜错了一半!
先问你个问题:你怎么理解“模型并行”?
很多人脱口而出:“就是把不同层放到不同 GPU 上嘛。”对了一半,但另一半才要命。Megatron 的模型并行其实是 两个层次的拆分绑在一起干,代码里根本分不开。
张量并行(TP):把一个 Transformer 层里的权重矩阵切成两块,分到两张卡。每张卡只存一半权重,前向、反向各算一半,最后通过集合通信拼起来。就像两个人一起搬一张大桌子,每人抬一半,到地方再合上。
流水线并行(PP):把不同 Transformer 层放到不同设备,输入数据切成微批量,像工厂流水线一样往前推。
问题是,这两个在 Megatron 里是 绑死在同一套进程组体系里 的。不是 AA、BB 分开的模块,而是 AA、BB 共享一个进程分配系统。
你去看 initialize_megatron 那一段,什么 tensor_model_parallel_group、pipeline_model_parallel_group、data_parallel_group……一下子冒出来七八个组,你知道吗,我第一次跑 pretrain_bert_distributed.sh 时,看到 --tensor-model-parallel-size 8,心想:简单,把模型按层切 8 份呗。结果初始化完打印 parallel_state,发现自己的进程同时参加了三个不同的 group!那种感觉就像进了迷宫,每个岔路口都有三个方向,走哪个都错。
数据并行只需要一个 AllReduce group 就够了。模型并行?你需要搞清楚每个算子在哪个 group 里通信。
反直觉吧?大多数人以为模型并行是切模型,其实首先要切的是 进程组。
第二关:代码里怎么切模型的?一个 Linear 层为什么要拆成两个?
这是当年卡了我三天的坑。你敢信?
在 megatron/core/tensor_parallel/layers.py 里,你会发现所有线性层都被换成了 ColumnParallelLinear 或 RowParallelLinear。这两个名字已经说出了切分的秘密:
- **ColumnParallelLinear**:把权重矩阵按列切。每张卡只拿到一部分列。输入数据完整进入,每张卡算一部分输出,然后通过 AllGather 拼起来——或者干脆不拼,等下一层再处理。好比一家公司分两个部门接客户,每个部门只处理自己那一半的业务,最后再汇总。
- **RowParallelLinear**:把权重矩阵按行切。每张卡只拿到一部分行。输入数据先通过 ReduceScatter 分发到各卡,每张卡算一部分输出,最后 AllReduce 求和。就像两个厨师各拿一半食材一起做一道菜,最后把各自的成果混在一起。
为什么搞这么复杂?因为 Transformer 里有两种主要计算:MLP 扩展(线性层把 hidden dim 放大到 4 倍) 和 Attention 投影。按照论文的设计:
- 第一个 MLP 线性层:用 ColumnParallelLinear,输出分到各卡,但不立即拼接。因为下一层是 GELU,element-wise 操作,不需要通信。完美省掉一次 AllReduce!
- 第二个 MLP 线性层:用 RowParallelLinear,输入已经分散在各卡(上一层的输出分片),乘完自己的部分后做 AllReduce,得到完整输出。刚好一次通信搞定。
你看,这里面藏着设计哲学:能省通信就省通信,但必须保证数学等价。
有个坑我摔得特别惨:手写模型时,我们习惯把权重定义成完整的矩阵。但在这套代码里,权重初始化时就已经切了一半了!你打印 model.layers[0].mlp.fc.weight.shape,得到的是 (hidden_size // tp_size, 4 * hidden_size),而不是完整的形状。我第一次调试时愣了半天,以为模型加载错了。后来才明白:代码里从头到尾就没有完整的权重矩阵,它就是让你存切片的。
还有,Embedding 也是通过 ParallelEmbedding 实现的,同样切了词表维度。如果词表大小不是 TP size 的整数倍,程序直接报错。这事儿我踩过,建议你提前把词表补成 TP size 的整数倍,省得被官方那个 assert 卡掉。
第三关:TP 和 PP 怎么结合?进程组怎么分配到硬件上?
这个问题被问得最多,也是我最想吐槽的——因为绝大多数文章只讲概念,不讲进程组怎么映射到硬件。就跟只告诉你菜谱,却不给你厨房工具一样。
举个例子。CodeGeeX 的配置(我反复拿它说事):8 头 TP(同机 8 卡),192 头 DP(跨机)。注意它 PP size=1,所以没有流水线。但如果你看 Megatron 自己的配置例子,很多是 TP=2, PP=4, DP=64 这种组合。好,一旦组合起来,进程分配的逻辑就变成这样了:
1. 总进程数 = TP size × PP size × DP size。
2. 先按 PP size 拆成多个 pipeline stage group,每个 stage 里再按 TP size 拆成 tensor parallel group。
3. 剩下的维度就是 data parallel group。注意,data parallel group 是 跨 pipeline 和 tensor 并行 的,每个组里包含那些互为数据并行的副本。
用人话说就是:你有一堆 GPU,先分成几组流水线,每个流水线组里再根据 TP size 分成若干 TP 组,每个 TP 组里做张量并行。所有 TP 组里同一位置的 GPU,加上所有流水线组里同一位置的 GPU(如果有数据并行的话),组成一个 DP 组。
Megatron 的 initialize_model_parallel 函数干的就是这件事。它会创建 _MODEL_PARALLEL_GROUP、_TENSOR_MODEL_PARALLEL_GROUP、_PIPELINE_MODEL_PARALLEL_GROUP 等一堆全局变量。魔鬼就在细节里:如果一个进程既在 TP group 又在 PP group,它在通信的时候必须知道用哪个 group 的句柄。比如做张量并行内的 AllReduce,要用 mpu.get_tensor_model_parallel_group(),不能错用成 mpu.get_data_parallel_group()。
我犯过一个低级错误:手写自定义算子需要做 AllReduce,结果在 TP context 里用了默认的 group(全局所有进程),直接炸了。后来查了半天,发现是 group 没设对。从此我养成了习惯:在任何涉及集合通信的地方,显式传入 group 参数。这习惯救了我很多次。
第四关:负载不均衡怎么办?有什么坑?
模型并行最恶心的问题就是负载不均衡。流水线并行天生有气泡(bubble),虽然通过交错调度和 microbatch 可以压到很小,但依然存在。张量并行相对均匀,但如果某个线性层的 hidden size 不是 TP size 的整数倍,切分就会不均匀,导致有的卡忙死、有的卡闲死。
另一个坑是内存。TP 会把每张卡上的模型大小压到 1/TP_size,但激活值(activation)的大小是翻倍的。因为每张卡仍需存完整的输入 batch 的激活(比如 layer norm 前的输入),这对显存是个挑战。Megatron 于是引入了序列并行(Sequence Parallelism),把 LayerNorm 和 Dropout 也切到序列维度,从而进一步降低激活显存。但序列并行必须和 TP 一起用,单独开没用——因为需要那个 f 和 f' 的通信融合,否则会多一次冗余的 AllReduce。
实践中我建议,如果版本支持,做 TP 时务必开启 --sequence-parallel。实测能降低约 30% 的激活显存,代价是增加一点通信量,但划算。30% 的显存,够你塞多一组 batch 了。
另外,CodeGeeX 的例子之所以 TP size 设成 8,是因为单机 8 卡,NVLink 带宽够高,TP 引起的通信延迟相对小。如果跨机做 TP(比如 TP=16,两台 8 卡机器),就要考虑跨机带宽瓶颈。我试过一次,跨机 TP 的 AllReduce 速度比机内慢 3 倍,直接拖慢了整体吞吐。所以除非你带宽逆天,否则 TP 别跨机。这句话值很多 GPU 小时。
第五关:源码里哪些函数最值得读?我帮你画重点
如果要读 Megatron 的模型并行代码,我推荐三个必看的地方,按重要性排序:
1. megatron/core/tensor_parallel/layers.py:ColumnParallelLinear 和 RowParallelLinear 的实现,以及 LinearWithGradAccumulationAndAsyncCommunication。这是 TP 的核心,看懂这两个类就懂了 80% 的张量并行。看的时候要盯着 weight 怎么初始化、forward 怎么切分。
2. megatron/core/tensor_parallel/mappings.py:Scatter、Gather、ReduceScatter 这些通信原语的封装。注意里面有个 _reduce 函数,是经过优化带 async communication 的版本,不走简单 AllReduce。这东西看起来简单,但细节决定了训练速度。
3. megatron/core/parallel_state.py:进程组初始化。我建议你在这文件里多打几个 print,打印每个进程的 rank、TP group 里的 rank、DP group 里的 rank,比你读任何文章都直观。相信我,看完打印输出,你一下就能从“概念懂”变成“代码也懂”。
还有一个非核心但值得看的:Transformer 层的前向函数。在 megatron/model/transformer.py 或 megatron/model/gpt_model.py 里,你会看到 TP 切分的具体调用点。看完这个,整个并行是怎么拼起来的基本就清楚了。
最后,几个让你再拍大腿的问题
为什么有些框架的 TP 叫“Model Parallelism”,而 Megatron 把 TP 和 PP 分开叫? 因为历史原因。早期 GPipe 把 PP 叫 model parallelism,Megatron 论文把 TP 叫 model parallelism。现在业界倾向用“模型并行”作为统称,包含 TP、PP 以及两者的混合。所以你在不同文章里看到同样的词,意思可能完全不同。阅读时一定要看上下文,否则容易掉坑。
TP size 设多大合适? 没有标准答案。通常和单机 GPU 数一致(8),因为 NVLink 连接。如果用了 NVSwitch,可以扩大到 16 或更多,但要实测。我测过一个 72B 模型,TP 从 8 升到 16,通信开销增加 25%,但单卡显存降低 50%,整体吞吐反而提升了 10%。所以你需要的不是定理,而是 你的硬件和模型。
为什么 Megatron 不开源 PP 的动态调度? 我也想问。目前 Megatron 的 PP 调度策略是 1F1B(一个微批 forward,一个微批 backward),但代码里是编译时写死的。想用 PipeDream 那种异步调度?自己改。我认识一个朋友花了两周给 Megatron 加了一个 PP 调度器,最后发现性能还不如 1F1B,放弃了。这水很深,劝你别瞎折腾。
读源码我最常用的调试方法? 加断点,然后打印 torch.distributed.get_rank() 和 group 的 size。你用 IDE 的 debugger 远不如直接 print 来得快。因为多进程 debug 太痛苦了,一个 print 就能让你看到全局。
尾声:带着这句话走
好了,这篇就先到这儿。如果你在读完源码后发现自己还是没完全懂,别急——我当时也读了三四遍才搞通。记住一个原则,它能让你少怀疑自己智商一百次:
先把进程组映射搞明白,再去看模型定义,最后看训练循环。顺序错了,你会一直问自己‘为什么我的通信炸了’。
模型并行的本质,不是切模型,而是切通信。谁掌握了进程组,谁就掌握了分布式训练的灵魂。
现在,翻代码去吧。你会感谢自己的。🔥
读者评论 3