
我常说自己现在是研究训练 infra 的但回顾自己在个人博客等平台发表的文章一直没有在这个方向有所产出。所以借着这次契机来分享和学习一下我们团队最近发表的一篇关于专家并行不均衡问题的文章吧。背景论文地址https://arxiv.org/pdf/2608.03676我常说自己现在是研究训练 infra 的但回顾自己在个人博客等平台发表的文章一直没有在这个方向有所产出。所以借着这次契机来分享和学习一下我们团队最近发表的一篇关于专家并行不均衡问题的文章吧。背景论文地址https://arxiv.org/pdf/2608.03676翻译一下TAOT面向混合专家模型训练中动态专家副本放置的拓扑感知最优传输方法这篇论文名称看似复杂实际上还是解决MoE一个老生常谈的负载不均衡问题即我们在做EP并行的时候每张卡能处理多少token完全由路由动态分配如果我们不进行人为的干预那么就会造成某些卡一直在计算而某些卡一直或者不间断的空闲摸鱼。所以这个问题要治其实业界和学术界有许多应对思路在算法层面像DeepSeek等MoE系列论文就在不停创新一些辅助损失函数方法来追求统计上的长期平衡而在系统方面则是保持路由策略不变通过动态调整计算资源的分配来解决负载不均。例如如果某个专家的计算量过高它会临时复制一份权重放到空闲卡上由这张卡分担一部分计算。这就是动态复制专家 (Hot-Expert Replication)也是本篇论文的聚焦点。本篇论文站在系统方法的肩膀上发现这些方法几乎只关注“如何平衡负载”却忽略了“平衡负载本身的成本”。一般来说我们实现的都是大规模EP并行那就不仅仅要在机内通信还要在机与机之间。而同一台机器内的GPU通过NVLink/PCIe互联跨机器之间则通过InfiniBand。两者之间的传输速度不同。这意味着把一个专家从同服务器的GPU复制过来成本很低但从另一台服务器复制过来成本则高昂得多。现有的动态复制方法为了追求极致的负载均衡可能会把一个热门专家复制到远在另一台服务器的GPU上。这种跨节点的通信开销很可能超过了复制专家所带来的计算收益最终导致训练反而更慢。方法那 TAOT 是怎么解的呢TAOT在追求负载平衡的同时把通信开销也一起算进去。具体是使用一个三阶段规划器从 rank 级到专家级再到 token 级由粗到细。之所以要拆成三段是因为把“均衡”和“通信代价”塞进同一个联合优化问题里规模会随 EP 爆炸一次性求解并不现实。Phase 1Sinkhorn-Knopp 拓扑感知流规划rank 级这一步的做法是把每个 rank 超出来的那部分负载当成“供给”把空闲 rank 的剩余容量当成“需求”再配上一个表示机内便宜、跨机贵的拓扑代价矩阵WWW最终求一个最优传输OT方案。原始 OT 要跑线性规划对GPU计算不友好。但是加一个负熵正则项做松弛之后最优解会呈现 Gibbs 核结构Tdiag(u) M diag(v)T\mathrm{diag}(u)\,M\,\mathrm{diag}(v)Tdiag(u)Mdiag(v)只需要交替做几轮 GEMV 迭代也就是 Sinkhorn-Knopp就能收敛全程矩阵向量乘u(t1)sM v(t),v(t1)dM⊤u(t1)\mathbf{u}^{(t1)} \frac{\mathbf{s}}{M\,\mathbf{v}^{(t)}}, \quad \mathbf{v}^{(t1)} \frac{\mathbf{d}}{M^\top \mathbf{u}^{(t1)}}u(t1)Mv(t)s,v(t1)M⊤u(t1)d这里再把正则系数直接取成跨节点代价λ\lambdaλ之后节点内与跨节点的核值之比天然大于 1直接可以来表示机内便宜、跨机贵的软拓扑偏好。Phase 2列优先迭代匹配专家级Phase 1 算出来的只是“这张卡该往那张卡挪多少负载”这种连续的量但真到放副本的时候没法这么灵活一份专家权重要么整个复制过去要么就不复制没有中间状态而且一张空闲卡上最多也只能放 K 份副本。更麻烦的是好几张空闲卡可能同时想接同一个计算量最大的专家。所以这一步要把那张流量表落成一个明确答案哪个专家复制到哪张卡上。判断标准是下面这个三级打分均衡收益、拓扑、Phase 1 给的 OT 提示依次让位前一项打平了才轮到后一项说话scoreermin(rem_spille, rem_sparer)⏟主分均衡改善量αBer⏟次分拓扑0.1α(Ter)norm⏟第三OT 提示\text{score}_{er} \underbrace{\min(\text{rem\_spill}_e,\ \text{rem\_spare}_r)}_{\text{主分均衡改善量}} \alpha \underbrace{B_{er}}_{\text{次分拓扑}} 0.1\alpha \underbrace{(T_{\text{er}})_{\text{norm}}}_{\text{第三OT 提示}}scoreer主分均衡改善量min(rem_spille,rem_sparer)α次分拓扑Ber0.1α第三OT提示(Ter)norm比公式更关键的其实是遍历顺序因为scoreer\text{score}_{er}scoreer的行是专家、列是卡而贪心一次只能敲定一格先算谁就等于把选择权先给谁。如果按专家做外层循环行优先那么排在前面的几个计算量最大的专家会一路挑下去直到自己多出来的 token 被消化完才轮到下一个空闲卡上的位置就这样被头部几个专家先占掉了而排在后面那批同样过载的专家一点缓解都拿不到它们所在的卡照旧要算到最后。所以 TAOT 换成了站在空闲卡视角的列优先匹配每一轮让每张还有空位的卡各自挑一个最合适的专家撞车了再仲裁没挑到的卡进下一轮继续挑这样副本会自然摊到更多专家身上而且还能把“就近”这件事交给最清楚自己位置的卡去判断。Phase 3拉格朗日拍卖 token 分配token 级Phase 2 只定了“专家一共给某张卡分多少 token”但专家的 token 本来就散在好几张过载卡上还得决定每张过载卡各发多少。按比例硬切有两个毛病一是浮点截断误差二是完全没用上拓扑信息。TAOT 的做法是给每张卡的容量约束挂一个拉格朗日乘子把它当作“价格”每轮各方按“拓扑收益减去当前价格”的净收益去竞价谁中标谁的价格就往上涨下一轮竞争力自然衰减。本质上就是一场拍卖会抢手的卡越抢越贵负载被自动摊平就近原则也顺手融进了分配过程而且迭代次数固定天然兼容 CUDA Graph。算法之外两道工程上的坎方法讲完了但要真把它塞进训练流程还有两件绕不开的事。第一道坎kernel太碎了上面的过程其实并不轻量而且它还得逐 microbatch 实时跑。如果规划本身开销过高通信侧省下来的收益很容易被原地抵消。论文给自己定的线是规划耗时不到一次前反向FB的 1%所以还要做算子级工程。问题主要出在小规模配置上。EP8、EP16 这种配置下规划的真实 GPU 计算量本来就小得可怜可原始 PyTorch 写法每调用一次都会触发一堆小 kernel重建张量、跑 50 轮 Sinkhorn、循环 K 轮 Phase 2加起来几百次 launch。这时候 host 端的启动开销已经远超 GPU 实际计算的时间。所以 TAOT 用 Triton 重写了这套逻辑从三个方向把 launch 次数砍下来静态张量缓存。拓扑代价矩阵MMM、每个专家原本所在的卡、拓扑偏好这些量只跟拓扑配置R、E、每节点卡数、跨节点代价有关跟 token 怎么分毫无关系整个训练过程中都不变。原始实现每次都重新构建白搭约 15 次 kernel launch缓存进一个进程级字典之后首次之后直接命中这 15 次就省了。Phase 1 单 CTA Sinkhorn kernel。50 轮 Sinkhorn-Knopp 原本是 50 次torch.mv加 clamp共约 150 次 launch。因为 R ≤ 64整个MMM、uuu、vvv都塞得进寄存器于是把 50 轮迭代整个塞进单个 CTA、编译期完全展开150 次 launch 变 1 次中途还省掉了 CPU-GPU 同步。Phase 2 的 K 轮融合。把 K 轮分配融成一次 launchused[E,R]掩码全程留在寄存器里、跨轮不回写显存最后才落盘。这里有个反直觉的取舍内层那 R 次迭代故意不做编译期展开因为一展开[E,R]临时张量就按 R 倍复制EP16 下寄存器直接爆掉、SM 占用率崩盘改用动态循环反倒让编译器能跨轮复用寄存器占用率明显回升。第二道坎计算-通信重叠跨节点通信降下来之后TAOT 还想再贪一点。一张卡上其实有两类专家权重本来就在这张卡上的叫 home 专家从别的卡复制过来的那份副本叫 guest 专家。既然 guest 专家的权重总要搬一趟那就如下图所示把这趟搬运塞进同一张卡上 home 专家的计算里。FC1/FC2 两个阶段里home 专家的 GEMM 在计算流上跑guest 专家的 expert_dispatch 在通信流上并行搬运靠一次同步把两者对齐两类专家之间就实现了计算-通信重叠。反向阶段同理guest 专家算出的权重梯度通过一次反向 All-to-All 送回原本所在的卡累加。这样 guest 专家机制引入的额外通信基本都被计算掩盖掉了。结果在包含32张NVIDIA A800 GPU 的集群和 Qwen3-30B-A3B 上TAOT实现了最高1.43倍的端到端训练加速。在达到与最先进方法同等甚至更优的负载均衡效果的同时其专家通信成本最高降低了74%。代码获取上述方法开源在 https://github.com/baidu-baige/LoongForge 大家可以结合AI把LoongForge 跑起来来学习。如果也是训练infra的同学可以多多关注我们的框架。因为不仅仅本文的方法我们也有多个在LLM/VLM/Diffusion/Embodied models上的优化哦