
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部分用TP8、EP1的配置跑得很舒服但专家部分因为参数量太大必须用EP8才能把显存放下来。结果就是Attention部分被迫跟着用EP8的通信组每次前向都要多做一轮不必要的all-to-all整体MFU直接掉了将近15个百分点。这就是典型的“一刀切”并行策略带来的浪费。2.2 PD分离的核心思想让合适的计算跑在合适的策略上训练PD分离的核心思想其实很朴素识别出训练过程中计算特征不同的阶段或模块给它们分别配置最合适的并行策略、通信组和资源配比。这里的P和D我倾向于把它理解为两种典型的计算范式——一种是计算密集、适合高并行度的比如Attention和稠密FFN另一种是通信密集或访存密集、需要特殊调度的比如MoE的专家路由和专家计算。这个思路和推理侧的PD分离在哲学上是一脉相承的但实现难度高了一个量级。推理侧Prefill和Decode是串行执行的拆开相对干净训练侧前向和反向是耦合的而且不同模块之间有梯度依赖拆分的边界需要非常小心。我见过一些团队尝试把整个前向和反向拆到不同机器上结果通信开销直接把收益吃光了。所以训练PD分离的可行做法不是粗粒度地拆阶段而是在模块级别做细粒度的策略分离。具体来说我总结下来有三个可操作的分离维度。第一个是并行策略分离Attention用TPSP专家用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序列长度4096batch size 8。单卡放不下整个Attention计算需要做TP切分。TP切分的基本单位是头。64个头如果TP8每张卡分到8个头每个头的QKV投影矩阵是[8192, 3*128]参数量约3M8个头就是24M参数显存放得下。但如果TP16每张卡只有4个头通信量会增加因为每次Attention计算后需要做all-reduce来合并结果。我实测下来TP8在这个配置下是甜点再大通信收益就递减了。序列并行SP是另一个维度。当序列长度超过8K时单卡的激活值显存会成为瓶颈。SP的做法是把序列维度切到不同卡上每张卡只计算部分序列的Attention然后通过ring attention或者all-gather来交换KV。我试过在序列长度16K时开SP4激活值显存从每卡48G降到14G效果立竿见影。但SP的通信模式比较复杂需要和TP配合使用配置错了容易出现死锁。这里有个实操心得Attention的TP和SP配置最好和模型的头数成整除关系。比如64个头TP8、SP2每张卡实际处理4个头、一半序列计算和通信比较均衡。如果TP6这种非整除配置会出现负载不均部分卡空转。我踩过一次TP6的坑MFU直接掉了8个点排查了半天才发现是头数分配不均导致的。3.2 MoE专家层的EP并行与通信优化MoE的专家层是训练PD分离收益最大的地方。专家层的参数量通常是Attention的几倍甚至十几倍必须用EPExpert Parallel来切分。EP的核心是把不同的专家放到不同的卡上每个token根据路由结果被发送到对应的专家卡上计算算完再发回来。EP的配置有几个关键参数。第一个是EP size也就是专家切分到多少张卡上。假设有64个专家EP8每张卡放8个专家。第二个是专家容量因子capacity factor控制每个专家最多处理多少token。容量因子设太小token会被丢弃影响模型效果设太大显存浪费严重。我一般从1.25开始调根据路由的均衡程度微调。all-to-all通信是EP的瓶颈。每次前向token需要从原来的卡发送到专家卡算完再发回来两次all-to-all。如果EP8通信组是8张卡all-to-all的延迟还可以接受。但如果EP64跨节点通信延迟会显著增加。我实测过一个EP32的配置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用TP2、SP1专家用EP4路由用DP8。第三步跑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 size4、8、16、EP size8、16、32、序列长度2K、4K、8K、以及是否开SP。每个配置跑50个step记录MFU和显存占用。这个扫描大概需要一天时间但能帮你找到大致的甜点区域。我实测过的一组数据是这样的在64卡、序列长度4K、专家数64的配置下TP8、EP16、不开SP的MFU是38.2%TP8、EP16、开SP2的MFU是41.5%TP16、EP16、开SP2的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分离后不同模块的参数量差异很大如果每次都保存全量checkpointIO开销很可观。我的做法是专家部分保存频率低一点比如每1000 stepAttention和路由保存频率高一点每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用TP8专家用EP32路由用DP64结果训练速度反而比不分离还慢。排查后发现通信组的数量太多NCCL的资源竞争严重而且不同通信组之间的同步等待时间很长。后来我总结了一个原则并行策略的维度不要超过3个。比如TPEPDP是合理的再加SP就要慎重。如果非要加尽量让SP和TP共用通信组减少通信组数量。我现在的配置一般是Attention用TPSP专家用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训练任务初始配置是TP8、EP8、不开SPMFU只有32%。我做了以下几轮优化。第一轮开SP2MFU提升到36%。激活值显存下降batch size可以开大一点。第二轮调整专家放置策略把热门专家分散到不同节点all-to-all占比从30%降到22%MFU提升到40%。第三轮路由改FP32aux loss系数从0.01调到0.03专家负载基尼系数从0.45降到0.28MFU提升到43%。第四轮开通信压缩FP8跨节点all-to-all通信量减半MFU提升到45%。第五轮分模块梯度裁剪训练稳定性提升可以用更大的学习率MFU最终到47%。这个案例里每一轮优化的收益都不大但累积起来很可观。关键是不要指望一次调优就到位要迭代着来。每轮优化后跑一段时间确认稳定了再做下一轮。我见过有人一次性改一堆配置结果出了问题不知道是哪个改动导致的排查成本很高。另外这个案例里的收益主要来自通信优化和负载均衡而不是计算优化。这也印证了前面的判断现在训练效率的瓶颈主要在调度和通信不在算力。PD分离的价值正是通过更精细的调度把通信和计算的效率榨出来。这个方向还有很多可以挖的地方比如动态调整并行策略、根据训练阶段自动切换配置都是值得尝试的。我接下来打算试试在训练不同阶段用不同的EP size前期用大EP快速收敛后期用小EP精细调优有结果再分享。