1. 从“训练也PD分离”说起:一个被低估的Scaling思路
第一次看到“训练也PD分离”这个说法,我脑子里蹦出来的其实是推理侧那套已经玩得很熟的Prefill-Decode分离架构。做LLM推理优化的人对PD分离肯定不陌生:Prefill阶段计算密集、序列并行度高,Decode阶段访存密集、逐token生成,两者对硬件资源的诉求完全不同,混在一起跑就会互相拖累。把这两个阶段拆到不同的机器、不同的并行策略上,吞吐和延迟都能明显改善。
但“训练也PD分离”这个提法有意思的地方在于,它把同样的解耦思想搬到了训练侧。这里的P和D不是Prefill和Decode,而是训练过程中的两种性质截然不同的计算负载。具体指什么,不同团队的理解略有差异,但核心逻辑是一致的:训练过程中存在计算特征差异极大的阶段或模块,把它们强行绑在同一套并行策略和同一批硬件上,是一种隐性的浪费。这个思路和KITE、Transformer、MoE这几个热词放在一起看,指向就很清晰了——它讨论的是大规模模型训练中,如何根据计算负载的性质做更细粒度的资源调度和并行策略拆分。
我自己在过去一年多的时间里,先后参与过几个千卡级别的训练任务调优,从Dense Transformer到MoE架构都踩过坑。最开始我对“PD分离”这个说法是有点怀疑的:训练不就是前向加反向,还能怎么分离?但真正把训练过程中的计算剖面拆开看之后,我发现这个思路不仅成立,而且在MoE和长序列场景下,收益比想象中大得多。这篇文章我就把自己对这件事的理解、实操中的具体做法、以及踩过的坑完整梳理一遍,适合正在做大规模训练调优、或者对Scaling效率感兴趣的朋友参考。不管你是刚接触分布式训练的新手,还是已经调过几百张卡的老手,应该都能从中找到一些可以直接抄作业的东西。
2. 训练PD分离到底在分离什么
2.1 训练负载的异质性:被忽视的效率杀手
要理解训练PD分离,首先得承认一个事实:一次完整的训练迭代里,不同计算模块的硬件诉求差异极大。我拿一个典型的MoE Transformer层来举例。一个MoE层里通常包含这么几块计算:Attention部分(QKV投影、注意力计算、输出投影)、路由网络(Router/Gate)、专家网络(Expert FFN)、以及各种归一化和残差连接。这几块计算的特征完全不同。
Attention部分在长序列场景下是典型的计算密集型,矩阵乘法的规模随序列长度平方增长,GPU的Tensor Core利用率可以打得很高。而路由网络是个很小的门控网络,参数量可能只有几百万,但它需要做全局的token到专家的分配,涉及all-to-all通信,是典型的通信密集型加访存密集型。专家网络则是参数量的大头,但每个token只激活其中一小部分专家,计算量和通信量取决于路由的均衡程度。
如果你把这三块绑在同一套并行策略上,会发生什么?我实测过一个具体案例:在一个64卡的MoE训练任务里,Attention部分用TP=8、EP=1的配置跑得很舒服,但专家部分因为参数量太大,必须用EP=8才能把显存放下来。结果就是Attention部分被迫跟着用EP=8的通信组,每次前向都要多做一轮不必要的all-to-all,整体MFU直接掉了将近15个百分点。这就是典型的“一刀切”并行策略带来的浪费。
2.2 PD分离的核心思想:让合适的计算跑在合适的策略上
训练PD分离的核心思想其实很朴素:识别出训练过程中计算特征不同的阶段或模块,给它们分别配置最合适的并行策略、通信组和资源配比。这里的P和D,我倾向于把它理解为两种典型的计算范式——一种是计算密集、适合高并行度的(比如Attention和稠密FFN),另一种是通信密集或访存密集、需要特殊调度的(比如MoE的专家路由和专家计算)。
这个思路和推理侧的PD分离在哲学上是一脉相承的,但实现难度高了一个量级。推理侧Prefill和Decode是串行执行的,拆开相对干净;训练侧前向和反向是耦合的,而且不同模块之间有梯度依赖,拆分的边界需要非常小心。我见过一些团队尝试把整个前向和反向拆到不同机器上,结果通信开销直接把收益吃光了。所以训练PD分离的可行做法,不是粗粒度地拆阶段,而是在模块级别做细粒度的策略分离。
具体来说,我总结下来有三个可操作的分离维度。第一个是并行策略分离:Attention用TP+SP,专家用EP,路由用DP,各走各的通信组。第二个是计算精度分离:Attention和专家计算用BF16,路由和归一化用FP32,减少数值误差。第三个是资源配比分离:给计算密集的模块分配更多高算力卡,给通信密集的模块优化网络拓扑。这三个维度可以组合使用,效果叠加。
2.3 为什么现在才被重视:Scaling瓶颈倒逼架构创新
这个思路其实不算新,早在Megatron-LM早期版本里就有类似的设计,把Attention和FFN用不同的TP策略。但为什么最近“训练PD分离”又被拿出来讨论?我觉得核心原因是Scaling的边际收益在下降,大家被迫从架构和调度层面找效率。
过去两年,模型规模从几十亿涨到几千亿,大家发现单纯堆卡堆参数的收益越来越不明显。一个千亿参数的Dense模型,训练MFU能到40%就算不错了,大部分时间都浪费在通信和等待上。MoE架构本来是为了解决这个问题,用稀疏激活降低计算量,但MoE引入了更复杂的通信模式,路由不均衡、专家负载倾斜、all-to-all开销大,这些问题让MoE的训练效率反而可能比Dense还低。
KITE这个工作我关注过一段时间,它讨论的正是如何在MoE训练中做更精细的通信和计算调度。虽然KITE的具体实现细节我没有完整复现过,但它的核心洞察和训练PD分离是一致的:训练效率的瓶颈已经从“算力不够”变成了“调度不优”。在这个背景下,把不同性质的计算负载分开调度,就成了一个必然的选择。Transformer架构的模块化特性,恰好为这种分离提供了天然的边界。
3. 核心细节拆解:从Attention到MoE的分离实操
3.1 Attention模块的并行策略选择与参数计算
Attention模块的并行策略选择,核心是看序列长度和头数。我拿一个具体配置来算:假设模型有64个头,每个头维度128,隐藏维度8192,序列长度4096,batch size 8。单卡放不下整个Attention计算,需要做TP切分。
TP切分的基本单位是头。64个头,如果TP=8,每张卡分到8个头,每个头的QKV投影矩阵是[8192, 3*128],参数量约3M,8个头就是24M参数,显存放得下。但如果TP=16,每张卡只有4个头,通信量会增加,因为每次Attention计算后需要做all-reduce来合并结果。我实测下来,TP=8在这个配置下是甜点,再大通信收益就递减了。
序列并行(SP)是另一个维度。当序列长度超过8K时,单卡的激活值显存会成为瓶颈。SP的做法是把序列维度切到不同卡上,每张卡只计算部分序列的Attention,然后通过ring attention或者all-gather来交换KV。我试过在序列长度16K时开SP=4,激活值显存从每卡48G降到14G,效果立竿见影。但SP的通信模式比较复杂,需要和TP配合使用,配置错了容易出现死锁。
这里有个实操心得:Attention的TP和SP配置,最好和模型的头数成整除关系。比如64个头,TP=8、SP=2,每张卡实际处理4个头、一半序列,计算和通信比较均衡。如果TP=6这种非整除配置,会出现负载不均,部分卡空转。我踩过一次TP=6的坑,MFU直接掉了8个点,排查了半天才发现是头数分配不均导致的。
3.2 MoE专家层的EP并行与通信优化
MoE的专家层是训练PD分离收益最大的地方。专家层的参数量通常是Attention的几倍甚至十几倍,必须用EP(Expert Parallel)来切分。EP的核心是把不同的专家放到不同的卡上,每个token根据路由结果被发送到对应的专家卡上计算,算完再发回来。
EP的配置有几个关键参数。第一个是EP size,也就是专家切分到多少张卡上。假设有64个专家,EP=8,每张卡放8个专家。第二个是专家容量因子(capacity factor),控制每个专家最多处理多少token。容量因子设太小,token会被丢弃,影响模型效果;设太大,显存浪费严重。我一般从1.25开始调,根据路由的均衡程度微调。
all-to-all通信是EP的瓶颈。每次前向,token需要从原来的卡发送到专家卡,算完再发回来,两次all-to-all。如果EP=8,通信组是8张卡,all-to-all的延迟还可以接受。但如果EP=64,跨节点通信,延迟会显著增加。我实测过一个EP=32的配置,all-to-all占了整个前向时间的35%,非常夸张。
优化all-to-all有几个手段。一是用分层all-to-all,先在节点内做,再跨节点做,减少跨节点流量。二是把路由计算和all-to-all重叠起来,路由算完一部分就先发一部分,不用等全部算完。三是调整专家放置策略,把热门专家分散到不同节点,避免单节点流量过大。这几个手段组合使用,我最多把all-to-all的占比从35%压到18%。
3.3 路由网络的精度与调度细节
路由网络虽然小,但它是MoE训练里最敏感的部分。路由的输出是一个softmax分布,决定每个token去哪些专家。如果路由的数值精度不够,容易出现路由崩塌——所有token都涌向少数几个专家,其他专家饿死。
我的做法是路由网络全程用FP32计算,包括门控线性层、softmax、以及top-k选择。虽然这会增加一点计算量,但相比路由崩塌带来的训练失败,这点开销完全值得。我见过一个团队为了省显存把路由也改成BF16,结果训练到一半路由熵急剧下降,模型效果直接崩了,回滚重训浪费了一周。
路由的调度还有一个细节是负载均衡损失的系数。MoE训练通常会加一个auxiliary loss来鼓励专家负载均衡,系数一般设在0.01到0.1之间。系数太小,负载不均衡;系数太大,路由被强行拉平,模型表达能力受损。我一般从0.01开始,观察专家负载的基尼系数,如果超过0.3就调大系数。这个调参过程需要盯着训练日志看,不能设完就不管。
另外,路由的top-k选择也有讲究。top-1路由计算量最小,但负载最容易不均衡;top-2路由计算量翻倍,但负载更均衡,模型效果通常也更好。我实测下来,在专家数超过32时,top-2的收益明显大于开销。如果专家数少,top-1也够用。
4. 完整实操流程:从单卡验证到千卡扩展
4.1 小规模验证:单机8卡跑通PD分离配置
任何大规模训练之前,我都会先在单机8卡上把PD分离的配置跑通。这一步的目的是验证并行策略的正确性,以及测量各个模块的实际开销占比。
具体步骤是这样的。第一步,写一个最小化的MoE Transformer模型,层数设2层,隐藏维度1024,专家数8,序列长度512。这个规模在单卡上都能跑,但为了验证并行,还是用8卡。第二步,配置并行策略:Attention用TP=2、SP=1,专家用EP=4,路由用DP=8。第三步,跑100个step,用profiler记录每个模块的耗时和通信量。
我一般用PyTorch的profiler,重点看三个指标:Attention的计算时间、专家的all-to-all时间、路由的计算时间。如果all-to-all占比超过30%,说明EP配置需要调整;如果Attention的计算时间远大于通信时间,说明TP可以再大一点。这个阶段的调优目标是让三个模块的耗时尽量均衡,避免某个模块成为瓶颈。
这里有个小技巧:在单机验证阶段就把通信组固定下来。比如Attention的TP组是[0,1]、[2,3]、[4,5]、[6,7],专家的EP组是[0,1,2,3]、[4,5,6,7],路由的DP组是全8卡。这样到了大规模训练时,通信组的拓扑结构可以直接复用,减少调试成本。我见过有人单机验证时随便配通信组,到了千卡环境发现通信组和网络拓扑不匹配,又得重新调,浪费了很多时间。
4.2 中等规模调优:64卡下的参数扫描
单机验证通过后,下一步是64卡的中等规模调优。这个规模足够暴露大部分通信和负载问题,但又不至于调一次要等太久。
64卡环境下,我一般会做一轮参数扫描。扫描的维度包括:TP size(4、8、16)、EP size(8、16、32)、序列长度(2K、4K、8K)、以及是否开SP。每个配置跑50个step,记录MFU和显存占用。这个扫描大概需要一天时间,但能帮你找到大致的甜点区域。
我实测过的一组数据是这样的:在64卡、序列长度4K、专家数64的配置下,TP=8、EP=16、不开SP的MFU是38.2%;TP=8、EP=16、开SP=2的MFU是41.5%;TP=16、EP=16、开SP=2的MFU是39.8%。可以看到SP的收益很明显,但TP从8加到16反而掉了,因为通信开销增加超过了计算收益。
这个阶段的另一个重点是验证路由的负载均衡。我会在训练日志里打印每个专家的token数,算基尼系数。如果基尼系数超过0.4,说明负载严重不均,需要调大auxiliary loss系数或者调整专家初始化。我遇到过一次基尼系数0.6的情况,排查发现是专家初始化时用了相同的随机种子,导致专家之间的区分度不够,路由倾向于选同一个专家。改成不同种子后,基尼系数降到0.25。
4.3 大规模部署:千卡环境的通信拓扑与容错
千卡以上的规模,通信拓扑就成了决定性因素。我参与过的一个千卡MoE训练,用的是分层all-to-all加节点内NVLink的方案。具体来说,节点内8卡通过NVLink全互联,节点间通过RDMA网络通信。EP的通信组尽量放在节点内,减少跨节点流量。
部署时有个关键决策是专家放置策略。如果专家均匀放在所有卡上,跨节点all-to-all的流量会很大。我的做法是把专家分成两组,一组放在前半数节点,一组放在后半数节点,路由时优先把token发到同组的专家。这样跨节点流量能减少一半。代价是专家利用率可能略低,但整体吞吐是提升的。
容错也是千卡环境必须考虑的。训练过程中难免有卡挂掉,如果每次挂卡都重启整个任务,浪费的时间太多。我的做法是在PD分离的框架下做模块级容错:如果挂的是专家卡,只重启专家部分的通信组,Attention和路由继续跑;如果挂的是Attention卡,同理。这需要训练框架支持动态通信组重建,实现起来有点复杂,但收益很大。我实测过一次挂卡恢复,模块级容错只花了3分钟,而全量重启花了25分钟。
还有一个细节是checkpoint的保存策略。PD分离后,不同模块的参数量差异很大,如果每次都保存全量checkpoint,IO开销很可观。我的做法是专家部分保存频率低一点(比如每1000 step),Attention和路由保存频率高一点(每200 step)。恢复时先加载专家,再加载其他部分。这样既保证了恢复的完整性,又减少了IO压力。
5. 常见问题与排查技巧实录
5.1 训练不稳定:路由崩塌与梯度爆炸的排查
路由崩塌是MoE训练最常见的问题,表现是训练到某个step后loss突然飙升,或者路由熵急剧下降。我排查这个问题的第一步是看路由熵的曲线。正常训练时路由熵应该缓慢下降但保持在一定水平,如果出现断崖式下跌,基本就是路由崩塌。
原因通常有三个。一是路由精度不够,前面说过,路由必须用FP32。二是auxiliary loss系数太小,负载不均衡导致部分专家梯度消失。三是学习率太大,路由网络的梯度更新过猛。我的排查顺序是:先确认路由精度,再调大aux loss系数,最后降学习率。大部分情况前两步就能解决。
梯度爆炸在PD分离配置下也有特殊性。因为不同模块用了不同的并行策略,梯度在all-reduce时的数值范围可能差异很大。我遇到过一次专家部分的梯度范数是Attention部分的100倍,导致梯度裁剪失效。解决办法是对每个模块单独做梯度裁剪,而不是全局裁剪。具体来说,给专家部分设一个更小的裁剪阈值,Attention部分设大一点。这个改动很小,但效果立竿见影。
5.2 通信瓶颈:all-to-all延迟高的定位方法
all-to-all延迟高是EP并行的老大难问题。定位方法我一般分三步。第一步,用NCCL的调试日志看通信时间,确认是all-to-all本身慢还是等待慢。第二步,用网络监控工具看跨节点流量,如果跨节点流量远大于节点内流量,说明专家放置策略有问题。第三步,用profiler看all-to-all和其他计算的重叠情况,如果完全没有重叠,说明调度有问题。
优化手段前面提过分层all-to-all和重叠调度,这里补充一个通信压缩的技巧。all-to-all传输的是token的隐藏状态,可以用FP8或者INT8压缩后再传,接收端解压。我实测过FP8压缩,通信量减少一半,精度损失在可接受范围内(loss曲线几乎无差异)。但要注意,压缩和解压本身有计算开销,如果通信量本来就不大,压缩反而得不偿失。我一般只在跨节点all-to-all时开压缩,节点内不压缩。
还有一个容易忽视的点是通信组的创建顺序。NCCL在创建通信组时,如果顺序不一致,可能导致通信组之间的干扰。我的做法是在训练开始前,按照固定的顺序创建所有通信组,并且给每个通信组分配独立的stream。这样能减少通信组之间的资源竞争。这个细节在文档里很少提,但实测能提升5%左右的通信效率。
5.3 显存不足:激活值重计算与专家卸载的取舍
显存不足在PD分离配置下更复杂,因为不同模块的显存压力不同。Attention部分主要是激活值占显存,专家部分主要是参数占显存,路由部分显存压力最小。
对于Attention的激活值,我一般用选择性重计算。只重计算Attention矩阵,不重计算QKV投影。这样能省下30%左右的激活值显存,计算开销增加不到10%。如果还不够,就上full重计算,但计算开销会增加30%以上,需要权衡。
对于专家的参数,如果显存放不下,可以用专家卸载,把不活跃的专家参数放到CPU内存,需要时再加载。但卸载的延迟很高,我实测过,卸载比例超过20%后,训练速度会下降一半以上。所以卸载是最后的手段,优先还是调EP size或者用更小的专家。
这里有个经验:显存优化要按模块分别做,不要全局一刀切。我见过有人全局开重计算,结果Attention部分显存是够了,但计算开销大增,MFU掉了10个点。正确的做法是只对显存压力大的模块开重计算,其他模块保持原样。
5.4 常见问题速查表
| 问题现象 | 可能原因 | 排查方法 | 解决手段 |
|---|---|---|---|
| loss突然飙升 | 路由崩塌 | 看路由熵曲线 | 路由改FP32、调大aux loss、降学习率 |
| MFU低于30% | 通信瓶颈 | NCCL日志、网络监控 | 分层all-to-all、通信压缩、调整专家放置 |
| 显存OOM | 激活值或参数过大 | 分模块看显存占用 | 选择性重计算、专家卸载、调EP size |
| 训练速度波动大 | 负载不均衡 | 看专家token数基尼系数 | 调aux loss系数、调整专家初始化 |
| 挂卡恢复慢 | 全量重启 | 看恢复日志 | 模块级容错、分级checkpoint |
| 梯度裁剪失效 | 模块间梯度范数差异大 | 分模块看梯度范数 | 分模块梯度裁剪 |
6. 我踩过的坑与实操心得
6.1 并行策略不是越多越好:过度分离的反效果
我最开始做PD分离时,恨不得把每个模块都拆开用不同的并行策略。Attention用TP=8,专家用EP=32,路由用DP=64,结果训练速度反而比不分离还慢。排查后发现,通信组的数量太多,NCCL的资源竞争严重,而且不同通信组之间的同步等待时间很长。
后来我总结了一个原则:并行策略的维度不要超过3个。比如TP+EP+DP是合理的,再加SP就要慎重。如果非要加,尽量让SP和TP共用通信组,减少通信组数量。我现在的配置一般是Attention用TP+SP,专家用EP,路由用DP,总共3个通信组,效果比较均衡。
另一个反效果是分离粒度太细。有人把Attention里的QKV投影、注意力计算、输出投影都拆开用不同策略,结果通信开销爆炸。我的经验是,分离粒度到模块级别就够了,模块内部保持一致的策略。模块内部的子计算通常特征相似,拆开收益很小。
6.2 精度配置的坑:BF16不是万能的
BF16是现在训练的主流精度,但在PD分离配置下,有些地方不能用BF16。除了前面说的路由必须用FP32,还有几个地方要注意。
归一化层建议用FP32。LayerNorm或者RMSNorm的数值范围比较敏感,BF16的精度不够,容易导致训练不稳定。我实测过,归一化用BF16时,训练到后期loss会有轻微震荡,改成FP32后震荡消失。
损失函数的计算建议用FP32。特别是MoE的aux loss,涉及多个专家的负载统计,BF16的累加误差会比较大。我一般把aux loss的计算单独拎出来用FP32,其他部分用BF16。
优化器的状态建议用FP32。Adam的动量和方差对精度敏感,BF16存储会导致优化器状态失真。现在大部分框架默认优化器状态用FP32,但如果你手动改过,记得改回来。
6.3 监控与日志:看不见的指标才是关键
训练PD分离配置时,常规的loss和MFU监控不够,还需要加一些模块级的指标。
每个模块的耗时占比。我一般每100 step打印一次,看Attention、专家、路由的耗时比例。如果某个模块占比超过50%,说明它是瓶颈,需要优化。
all-to-all的通信量。这个指标能反映EP的效率。如果通信量远大于理论值,说明有冗余通信,需要检查通信组配置。
专家负载的基尼系数。前面提过,这个指标反映路由的均衡程度。我一般每500 step算一次,超过0.4就告警。
梯度范数的分模块统计。这个指标能提前发现梯度爆炸的苗头。如果某个模块的梯度范数突然增大,及时干预。
这些指标我一般用TensorBoard或者WandB记录,设置告警阈值。训练过程中不用一直盯着,但出了问题能快速定位。
6.4 一个具体的调优案例:从32%到47%的MFU提升
最后分享一个我实际做过的调优案例。一个64卡的MoE训练任务,初始配置是TP=8、EP=8、不开SP,MFU只有32%。我做了以下几轮优化。
第一轮,开SP=2,MFU提升到36%。激活值显存下降,batch size可以开大一点。
第二轮,调整专家放置策略,把热门专家分散到不同节点,all-to-all占比从30%降到22%,MFU提升到40%。
第三轮,路由改FP32,aux loss系数从0.01调到0.03,专家负载基尼系数从0.45降到0.28,MFU提升到43%。
第四轮,开通信压缩(FP8),跨节点all-to-all通信量减半,MFU提升到45%。
第五轮,分模块梯度裁剪,训练稳定性提升,可以用更大的学习率,MFU最终到47%。
这个案例里,每一轮优化的收益都不大,但累积起来很可观。关键是不要指望一次调优就到位,要迭代着来。每轮优化后跑一段时间,确认稳定了再做下一轮。我见过有人一次性改一堆配置,结果出了问题不知道是哪个改动导致的,排查成本很高。
另外,这个案例里的收益主要来自通信优化和负载均衡,而不是计算优化。这也印证了前面的判断:现在训练效率的瓶颈主要在调度和通信,不在算力。PD分离的价值,正是通过更精细的调度把通信和计算的效率榨出来。这个方向还有很多可以挖的地方,比如动态调整并行策略、根据训练阶段自动切换配置,都是值得尝试的。我接下来打算试试在训练不同阶段用不同的EP size,前期用大EP快速收敛,后期用小EP精细调优,有结果再分享。