一个训练日志里的异常
假设你在训练一个稀疏 MoE 语言模型,语料混合了代码、法律合同、医学文献和客服对话。训练到中期,日志里出现一个熟悉又令人不安的现象:某几个专家的累计被选次数持续上升,其余专家几乎不再被路由命中。梯度仍然在流动,总损失也在下降,但模型的参数利用率越来越低——这就是专家塌缩。
直觉上,路由器应该学会把代码 token 送给擅长代码的专家,把合同条款送给擅长法务的专家。但训练早期路由权重接近随机,少量专家因为初始化或数据顺序先获得更多样本,被优化得更充分,于是下一轮更容易被选中。这个正反馈一旦启动,就会把其余专家饿死。辅助负载均衡损失(load balance loss,LBL)正是为打断这个正反馈而加入的正则项。代价是:它在防止塌缩的同时,也在持续给路由器施加一种“平均分配”的压力,这种压力与“让专家分化”的目标存在张力。
专家塌缩的成因与辅助损失的介入点
MoE 层由路由器和一个专家池组成。对 Transformer 而言,每个专家通常是一个前馈网络。以 Mixtral 8x7B 为例,每一层有 8 个前馈专家,路由器为每个 token 选择 2 个专家并加权合并输出,因此每个 token 能访问约 47B 参数,但推理时只激活约 13B 参数。这种“参数多、计算少”的特性来自稀疏激活,而稀疏激活的稳定性完全依赖路由决策的质量。
Switch Transformer 的工作把 MoE 路由简化为每个 token 只选一个专家,并明确指出训练不稳定是 MoE 广泛落地的主要障碍之一。更早的 Sparsely-Gated MoE 论文则用可训练门控网络在数千个子网络中选择稀疏组合,把模型容量推到千亿参数级别。两篇工作共同的隐含前提是:路由器必须被约束,否则它会退化成只使用极少数专家。
辅助损失的形式并不复杂。记专家 i 在一个统计窗口内的激活频率为 f_i,路由器分配给它的平均门控分数为 p_i,则常用的负载均衡损失大致是专家数乘以 f_i 与 p_i 乘积之和。工程上通常这样理解这两项:f_i 反映“实际被选中多少次”,p_i 反映“路由器想选它的倾向有多强”。损失同时惩罚两者,既纠正已经发生的偏斜,也抑制正在形成的偏好。
这里有一个容易被忽略的细节:f_i 是在哪个范围上统计的。如果统计窗口是单个 micro-batch,损失就要求“这一小批数据也要均匀分给所有专家”。而一个 micro-batch 的数据往往来自同一领域。阿里云通义 Qwen 团队在《Demons in the Detail》中正是抓住这一点:局部均衡会强迫同一领域的输入也均匀分配,从而阻碍专家形成领域层次的分化。
局部均衡与全局均衡的机制差异
把均衡范围从 micro-batch 扩大到 global-batch,是这篇工作提出的核心修正。做法本身很轻:专家激活频率只是一个长度为专家数的向量,在所有 micro-batch 之间同步这个向量,用全局统计的 f_i 参与每个局部位置的损失计算,再对损失做聚合。由于同步的数据量很小,通信开销有限,还可以用计算掩盖进一步隐藏。
论文给出的等价关系值得注意:用全局激活频率参与局部计算后再平均,等价于直接计算全局均衡损失。对于需要梯度累积的场景,他们还提出缓存机制,把各个累积步的激活频率累计起来,使得计算节点较少、单次通信覆盖范围有限时,也能逐步逼近全局统计。
为什么扩大范围有效?论文用一组对照实验排除了“只是统计方差变小”的解释。他们构造了 shuffled batch balance:从 global batch 中随机抽取一个与 micro-batch 等大的子集来统计激活频率。这个设置与 micro-batch balance 拥有相同的 token 数量,与 global-batch balance 拥有相同的 token 分布。结果显示 shuffled batch balance 与 global batch balance 表现几乎一致,都明显好于 micro-batch balance。这说明提升的首要原因不是样本量,而是统计集合中领域信息的多样性。
论文在 3.4B 激活 0.6B、15B 激活 2.54B、43B 激活 6.6B 三种规模上训练 120B 和 400B tokens,观察到均衡范围从常见框架实现的 4、8、16 增大到 128 以上后,Benchmark 指标和 PPL 都有明显改善;在 3.4B 激活 0.6B 训练 400B tokens 的设置上,balance BSZ 从 2 到 128 时 PPL 快速下降,128 之后逐渐饱和。这些数字来自论文,不应外推到其他数据配比或模型结构。
贯穿场景:多领域语料下的路由决策流
把上面的机制放进一个具体场景。假设我们有一批混合语料,micro-batch 大小为 1,每个 batch 恰好是一段代码或一份合同。局部均衡损失会要求这段代码的每个 token 尽量分散到所有专家;全局均衡损失则允许代码 token 集中流向少数专家,只要在整个 global batch 的统计上,各专家总负载大致均衡。
下面的流程图描述一次训练步中,token 从进入 MoE 层到损失回传的路径。关键转折点在“统计激活频率”这一步:统计窗口决定了均衡压力作用在哪个粒度上。
flowchart TD
A[混合语料 micro-batch] --> B[路由器计算门控分数]
B --> C[TopK 选择专家]
C --> D[专家容量因子裁剪]
D --> E[专家前向计算]
E --> F[加权合并输出]
C --> G[统计激活频率]
G --> H{统计窗口}
H -->|micro-batch| I[局部均衡损失]
H -->|global-batch| J[全局均衡损失]
I --> K[总损失]
J --> K
E --> L[路由器 logit 正则 z-loss]
L --> K
K --> M[反向传播更新路由器与专家]
这张图里有两个容易被忽视的环节。一是专家容量因子:它为每个专家设置可接收 token 的上限,超出部分会被丢弃或绕过。容量因子过小会让热门专家溢出,token 被送到次优专家;过大则浪费显存和计算。二是 z-loss:它作用在路由器 logit 上,惩罚 logit 幅度过大,目的是稳定训练数值,与负载均衡损失作用在不同对象上。
均衡强度与路由偏差的权衡
负载均衡损失是一个加权项,权重系数决定均衡压力有多强。系数太小,塌缩风险回升;系数太大,路由器被迫追求形式上的平均,专家分化被压制,模型质量可能下降。Wang 等人的工作把这种关系描述为负载均衡损失与语言模型损失之间的杠杆:两者优化目标并不一致,因此他们提出用基于专家选择频率更新的偏置项来平衡选择,而不改变路由分数,从而去掉辅助损失。
Qwen 团队的实验提供了另一个角度的证据:只使用全局均衡会降低局部均衡程度,进而影响 MoE 的计算效率;在主要使用全局均衡的前提下,额外加入权重为全局 LBL 1% 的局部均衡损失,每个更新步耗时从 1.64 秒降到 1.59 秒,同时模型效果几乎不受影响。这说明局部均衡并非完全无用,它在计算效率上有独立价值,只是不应作为主导约束。
下表对比几种常见做法在均衡范围、专家分化、计算效率和实现复杂度上的差异。表中结论来自资料支持与通用工程经验,不含未经验证的性能数字。
| 方案 | 均衡统计范围 | 对专家分化的影响 | 计算/通信开销 | 主要风险 |
|---|---|---|---|---|
| 纯局部均衡 | micro-batch | 压制领域分化,专家趋同 | 低 | 专家塌缩缓解但特异性弱 |
| 纯全局均衡 | global-batch | 允许领域层次分化 | 需同步激活频率向量,可掩盖 | 局部负载不均,推理效率下降 |
| 全局为主加少量局部 | global 加 1% 局部 | 分化基本保留 | 略高于纯全局 | 权重设置需调参 |
| 辅助损失无关方案 | 基于选择频率的偏置 | 不改变路由分数 | 依赖频率统计更新 | 与均衡范围的关系未经充分比较 |
从工程决策看,如果训练语料领域高度混杂且希望专家形成可解释的领域分工,全局均衡是更合理的主导约束;如果推理侧对单卡负载均衡极度敏感,保留少量局部均衡能减少热点专家带来的排队。
失败模式、可观测指标与适用边界
路由偏差不只表现为塌缩,还有几种更隐蔽的失败形态。
路由器抖动是其中之一。当均衡压力过强,路由器在相邻训练步之间频繁改变选择,专家参数无法稳定收敛。工程上通常观察每个专家的激活频率在滑动窗口内的方差,以及路由器输出分布在连续步之间的变化幅度。抖动加剧时,即使总损失仍在下降,验证集指标也可能停滞。
容量因子与均衡损失会相互放大问题。均衡损失推动均匀分配,容量因子则硬性截断超限 token。如果容量因子设得偏小,被截断的 token 会被送到非首选专家,这些 token 的梯度又反过来影响路由器,可能形成新的偏斜。观察指标包括每个专家的溢出 token 比例和绕过路径的占比。
z-loss 与均衡损失作用对象不同,但会共同影响路由器。z-loss 抑制 logit 幅度,均衡损失调整选择分布。如果 z-loss 权重过大,路由器 logit 被压平,门控分数接近均匀,专家区分度下降;如果过小,训练后期可能出现 logit 数值离群,影响低精度训练的稳定性。Switch Transformer 的工作提到其训练技术帮助模型首次以 bfloat16 低精度训练大规模稀疏模型,这提示数值稳定性与路由约束需要一起考虑。
评估偏差也需要注意。如果只在整体验证集上看 PPL,可能掩盖某些领域被少数专家垄断、其他领域专家退化的问题。更细的做法是按领域分别统计专家激活分布,检查是否存在某个领域几乎只使用固定几个专家,而另一些专家在该领域从未被激活。
适用边界方面,全局均衡的收益依赖一个前提:global batch 内确实包含多样化的领域信息。如果训练数据本身单一,或 global batch 很小,扩大均衡范围带来的多样性增益有限,此时收益可能不足以抵消同步开销。论文也指出,其实验主要集中在语言任务上,向其他模态或任务迁移需要重新验证。
部署侧还需要回答的问题
训练时的均衡约束会直接影响推理时的负载分布。全局均衡允许专家分化,但分化后的激活频率在真实请求分布下是否仍然均衡,取决于线上流量与训练语料的匹配程度。如果线上请求集中在某个领域,该领域的专家会成为热点,其他专家闲置,MoE 的稀疏计算优势被削弱。
一个实用的做法是把训练侧统计的专家激活分布作为基线,在推理服务中持续采集每个专家的实际请求量,对比两者的偏移。当偏移超过阈值时,可能意味着流量分布发生了变化,或者路由对某些输入模式产生了训练时未见的偏好。
仍未解决的问题包括:均衡范围应该随模型规模和语料多样性如何缩放;全局均衡与辅助损失无关方案在不同并行策略下的相对优劣;以及如何在训练早期就识别出正在形成的路由偏差,而不是等到塌缩已经发生。这些问题目前没有统一的答案,需要在具体训练配置下通过消融实验确定。