AI 技术
#MoE#专家容量#Token丢弃#负载均衡#稀疏训练

MoE 专家容量因子与 Token 丢弃:容量溢出如何影响训练质量与推理一致性

训练稀疏 MoE 时,专家容量因子决定每个专家能接收的 token 上限,超出的 token 会被丢弃或经残差绕过。本文以多语言翻译训练为贯穿场景,解释容量因子、负载均衡损失与丢弃率的耦合关系,对比 GShard 与 Switch Transformer 的容量设计,并讨论丢弃对长尾 token 学习、推理期行为差异和可观测指标的影响。

一个被丢弃的稀有语种 token

假设你在训练一个覆盖上百种语言的多语言翻译模型,骨干是稀疏 MoE。某个 batch 里混入了大量英语和中文句子,同时夹着少量冰岛语和斯瓦希里语的短句。路由网络给这些句子打分后,英语 token 大量涌向少数几个“通用语义专家”,而冰岛语 token 因为训练样本少,路由器对它们的偏好还不稳定,可能也被推向同一批热门专家。热门专家的待处理队列迅速膨胀,冷门专家却几乎空闲。

如果系统允许热门专家无限接收 token,显存和计算时间就会由最忙的专家决定,训练步会被拖长。如果给每个专家设一个上限,超出的 token 就必须有去处。这个上限就是专家容量,而超出部分被丢弃或绕过残差连接,正是本文要讨论的机制。

容量因子不是超参数列表里的一个孤立数字。它同时牵动负载均衡损失、token 丢弃率、长尾语种的学习效果,以及推理阶段同一输入在不同 batch 组成下是否得到一致输出。

容量从哪里来:从路由到固定形状张量

稀疏 MoE 的核心思路是条件计算:每个 token 只激活一部分专家,从而在不按比例增加计算量的前提下扩大参数量。Shazeer 等人的 Sparsely-Gated Mixture-of-Experts 层用可训练的 gating network 为每个 example 选择稀疏的专家组合,在语言建模和机器翻译任务上验证了这一方向。

训练时,一个 batch 的 token 需要被分发到不同专家。如果每个专家接收的 token 数完全由路由结果决定,张量形状就会随 batch 内容变化,GPU 上的分组矩阵乘法难以高效执行。工程上的常见做法是:为每个专家分配一个固定的槽位数,即容量。容量通常写成

capacity = capacity_factor × (tokens_per_batch / num_experts)

这里的 tokens_per_batch 是当前 batch 送入 MoE 层的 token 总数,num_experts 是专家数量。括号里的部分可以理解为“如果完全均匀,每个专家应处理的平均 token 数”。容量因子大于 1 时,每个专家能接收的 token 多于平均值;小于 1 时则少于平均值。

这个公式的工程含义是:容量把动态的路由结果压成静态的张量形状。每个专家拿到一个固定大小的输入缓冲区,分组 GEMM 可以按这个形状编译执行。代价是,缓冲区之外的 token 无法进入该专家。

GShard 在 2020 年把这一思路扩展到超过 6000 亿参数的稀疏 MoE 多语言翻译模型,并在 2048 块 TPU v3 上训练。它的贡献之一是提供轻量注解 API 和 XLA 编译器扩展,让专家分片、容量约束和 all-to-all 通信可以由编译器处理,而不是手写每处并行逻辑。

溢出之后:丢弃、残差与“位置在变、语义没变”

当某个专家收到的 token 超过容量,系统必须决定哪些 token 进入、哪些被排除。常见策略是按路由权重排序,保留权重最高的 token,其余标记为溢出。

被排除的 token 有两种处理方式。一种是直接丢弃,该 token 在这一层不经过任何专家,输出为零向量或直接跳过。另一种是经残差连接绕过,token 的原始表示继续传到下一层,只是没有获得专家子网络的变换。

这两种方式对训练信号的影响不同。直接丢弃意味着该 token 在这一层完全没有梯度贡献,路由器和专家都收不到关于它的学习信号。残差绕过保留了 token 的表示,但专家部分没有更新。工程上通常认为,残差绕过比硬丢弃温和,因为它至少不破坏主干信息流;但它仍然让该 token 错过了专家层的容量扩展。

容量因子越小,溢出越频繁。极端情况下,如果容量因子远小于 1,大量 token 被排除,MoE 层退化成接近恒等映射,参数量的优势无法转化为质量提升。容量因子越大,溢出越少,但固定缓冲区越大,显存占用和计算浪费越多,因为冷门专家的槽位大部分是空的。

负载均衡损失与容量因子的耦合

路由网络天然倾向于把 token 发给少数几个它认为“好”的专家。如果没有约束,热门专家会越来越热,冷门专家得不到足够梯度,进一步被路由器冷落。这种正反馈就是专家塌缩。

GShard 和 Switch Transformer 都引入辅助负载均衡损失来对抗这一趋势。它的基本形式是让每个专家被选中的概率分布接近均匀分布,从而给冷门专家更多机会。Switch Transformer 进一步把路由简化为每个 token 只选一个专家,减少了通信和计算开销,并在 T5-Base 和 T5-Large 上报告了相同计算资源下最高 7 倍的预训练加速。

负载均衡损失和容量因子作用在不同环节。均衡损失在训练时调整路由器的偏好,希望从源头让 token 分布更均匀;容量因子在分发时截断每个专家的接收量,是事后的硬约束。两者配合时,均衡损失降低溢出概率,容量因子限制最坏情况下的队列长度。

但均衡损失本身会引入路由偏差。它鼓励路由器把 token 分给冷门专家,即使某个热门专家在当前上下文下确实更合适。容量因子越小,这种偏差的代价越大:被均衡损失推向冷门专家的 token 可能因为冷门专家容量已满而被丢弃,或者反过来,热门专家因为容量截断丢失了本该由它处理的 token。

训练中的丢弃率:一个被低估的诊断信号

在多语言翻译训练中,丢弃率不是均匀分布的。高频语种的 token 表示在训练早期就趋于稳定,路由器对它们的偏好明确,容易集中到少数专家。低频语种的 token 表示稀疏,路由器打分方差大,可能在不同专家之间摇摆。

如果容量因子设置得偏小,高频语种的大量 token 会挤占热门专家的槽位,低频语种的 token 即使被路由到正确专家,也可能因为排序权重低而被挤出容量。结果是:模型在英语到中文等高频方向上继续改善,但在冰岛语到英语等长尾方向上,专家层几乎没有为这些 token 提供有效变换。

训练日志里可以观察几个信号。每个专家的接收 token 数与容量之比,反映该专家的填充率;全局丢弃率反映溢出总量;按语种分组的丢弃率则揭示长尾是否被系统性牺牲。如果某个语种的丢弃率显著高于平均,而该语种的验证集指标停滞,容量因子或均衡损失权重就需要重新审视。

Switch Transformer 的论文提到,他们用训练技巧来抑制稀疏模型的不稳定性,并首次展示了大型稀疏模型可以用 bfloat16 低精度格式训练。容量约束和丢弃策略是这些训练技巧的一部分,但论文没有给出所有语种上丢弃率与质量关系的完整分解。

推理期:同样的 token,不同的 batch 命运

训练时的丢弃影响参数更新,推理时的丢弃影响输出一致性。

推理阶段,请求通常按 batch 组织。一个 batch 里如果混入了大量同质请求,比如都是英语长句,热门专家的负载会远高于平均值。容量感知推理的研究把这种现象称为 Straggler Effect:负载轻的专家早早算完,但必须等待负载重的专家,因为专家并行下各设备之间有同步屏障,最忙的专家决定整体延迟。

Capacity-Aware Inference 的工作提出对高负载专家执行 token drop,强制容量上限,丢弃超出部分。在 OLMoE 上,他们报告了约 30% 的加速,性能下降约 0.9%。他们还发现,即使高负载专家被截断,一些低负载专家仍远低于容量,于是进一步提出 Expanded Drop,让溢出 token 在候选集中考虑额外的本地专家,以提高冷门专家的利用率。在 Mixtral-8×7B-Instruct 上,他们报告了平均性能提升约 0.2% 和约 1.85 倍的推理加速。

这些数字来自该论文的实验设置,不能直接外推到所有模型和硬件。但它们说明一个关键点:推理期的容量约束不仅影响延迟,也影响输出。同一个 token,如果单独请求时进入某个专家,在混合 batch 中可能因为容量已满而被丢弃或改派给其他专家,两次输出就可能不一致。

这种不一致在训练时也存在,只是被平均掉了。训练时每个 batch 的组成不同,同一个 token 在不同 step 可能被丢弃或保留,梯度信号因此带有噪声。容量因子越小,噪声越大。

容量因子、均衡损失与丢弃率的三角关系

下面的表格对比三种典型配置在训练和推理中的表现。表中的判断基于 GShard、Switch Transformer 和 Capacity-Aware Inference 论文所描述的设计,以及稀疏 MoE 的通用工程经验;具体数值依赖模型、数据和硬件,不能直接套用。

配置容量因子均衡损失典型丢弃率训练质量影响推理延迟适用场景
紧容量偏小较强高长尾 token 学习受损,高频方向仍可收敛低,最忙专家队列短推理吞吐优先,训练资源受限
宽容量偏大中等低长尾 token 获得更多专家变换,但显存和计算浪费高,冷门专家槽位空置训练质量优先,显存充足
均衡容量接近 1适中中等高频与长尾折中,依赖均衡损失抑制塌缩中等通用多语言或混合负载
无容量约束不设上限弱无热门专家过载,训练步长由最忙专家决定极高,Straggler Effect 明显小规模实验或专家数很少

这张表的核心信息是:容量因子不是越大越好,也不是越小越好。它是在训练质量、显存、推理延迟和长尾覆盖之间做取舍。均衡损失可以降低对容量因子的依赖,但它本身会引入路由偏差,不能替代容量设计。

一次训练步里 token 的完整路径

下面的流程图展示一个 batch 中 token 从路由到专家、再到溢出处理的完整路径。场景是前面提到的多语言翻译训练:一个 batch 里混合了高频语种和低频语种的句子。

flowchart TD
    A[输入 batch: 混合语种 token] --> B[路由器为每个 token 打分]
    B --> C[按专家分组, 统计各专家接收量]
    C --> D{专家接收量是否超过容量}
    D -- 否 --> E[token 进入该专家]
    D -- 是 --> F[按路由权重排序, 保留高权重 token]
    F --> G[溢出 token 处理]
    G --> H[直接丢弃: 该层无专家变换]
    G --> I[残差绕过: 保留原始表示]
    E --> J[专家子网络计算]
    H --> K[汇总到下一层]
    I --> K
    J --> K
    K --> L[计算任务损失与负载均衡损失]
    L --> M[反向传播更新路由器与专家]

图中最关键的分支是 D 到 F。容量判断发生在所有 token 的路由分数已知之后,但专家计算之前。排序保留了路由权重最高的 token,这意味着被丢弃的往往是路由器置信度较低的 token。在多语言场景中,低频语种 token 因为表示不稳定,路由器置信度偏低,更容易落入被丢弃的一侧。

另一个关键点是 K 到 L:被丢弃的 token 仍然参与任务损失计算,因为它们的主干表示经过残差连接传到了下一层。但专家部分的参数没有收到这些 token 的梯度。如果丢弃率长期偏高,专家层的有效训练样本数会少于名义 batch size,收敛速度可能下降。

失败模式与可观测指标

容量因子设置不当会以几种方式暴露。

第一种是长尾语种指标停滞。全局损失继续下降,但低频语种的验证集 BLEU 或准确率不再改善。检查按语种分组的丢弃率,如果长尾语种的丢弃率显著高于平均,说明容量截断正在系统性排除它们。

第二种是专家利用率两极分化。少数专家的填充率长期接近或达到容量,多数专家填充率很低。这说明均衡损失不足以抵消路由器的偏好,或者容量因子太小,使得热门专家频繁溢出,而冷门专家的空余容量没有被利用。Capacity-Aware Inference 的 Expanded Drop 正是针对这一现象:让溢出 token 在候选集中考虑额外专家,把冷门专家的空余容量用起来。

第三种是推理输出不一致。同一个 prompt 在不同 batch 组成下得到不同结果,差异集中在那些处于容量边界附近的 token。如果业务对一致性敏感,推理期需要固定 batch 组成,或者提高容量因子减少溢出,或者记录每个请求的丢弃 token 位置以便复现。

第四种是训练不稳定。Switch Transformer 的论文提到稀疏模型存在训练不稳定问题,并用训练技巧来抑制。容量因子过小会放大这种不稳定,因为每个 step 被丢弃的 token 集合变化较大,梯度噪声增加。

可观测指标包括:全局丢弃率、按专家分组的填充率、按语种或按 token 类型分组的丢弃率、路由权重分布、负载均衡损失值、专家梯度范数。训练时如果发现某个专家的梯度范数长期接近零,同时它的填充率很低,说明它正在被路由器冷落,需要检查均衡损失权重和容量因子。

与替代方案的比较

容量约束和 token 丢弃不是处理专家负载不均的唯一方法。

一种替代方案是不设容量,让每个专家处理实际路由给它的所有 token。这在小规模实验或专家数很少时可行,但在专家并行下会导致 Straggler Effect,最忙的专家决定整体延迟。Capacity-Aware Inference 的测量显示,最高负载专家处理的 token 数可以超过平均负载的七倍。

另一种替代方案是复制高负载专家,把副本部署到不同设备上分担负载。DeepSeek-V3 采用了这种思路。它的好处是不丢弃 token,质量损失小;代价是额外的显存和部署复杂度,而且副本之间需要同步或至少保持一致。

还有一种方案是调整路由算法本身,比如用更强的均衡约束或不同的路由打分方式,从源头减少不均衡。这与容量因子正交,可以叠加使用。但更强的均衡约束会引入更大的路由偏差,可能把 token 推给并不最适合的专家。

容量因子加丢弃的优势是实现简单、张量形状固定、对现有训练框架改动小。它的代价是丢弃带来的质量损失和推理不一致。选择哪种方案,取决于业务对质量、延迟、显存和一致性的优先级。

尚未解决的问题

容量因子应该设多大,目前没有通用公式。它依赖专家数量、路由算法、数据分布、batch 组成和硬件拓扑。工程上通常从小规模实验开始,观察丢弃率和长尾指标,再逐步调整。

丢弃对长尾 token 的长期影响也缺少系统研究。被丢弃的 token 如果反复出现,模型是否会在后续 step 中学会把它们路由到容量充足的专家?还是说路由器因为收不到这些 token 的梯度而永远无法修正?这取决于残差绕过保留了多少信息,以及均衡损失是否足够强。

推理期的一致性保障同样没有标准做法。在 batch 组成动态变化的生产环境中,同一个请求可能因为并发请求的不同而得到不同输出。记录丢弃 token 的位置和原因,是复现和诊断的起点,但如何在延迟约束下保证关键 token 不被丢弃,仍然是一个开放问题。

资料来源

  1. Switch Transformers: Scaling to Trillion Parameter Models with Simple and Efficient Sparsity
  2. GShard: Scaling Giant Models with Conditional Computation and Automatic Sharding
  3. Outrageously Large Neural Networks: The Sparsely-Gated Mixture-of-Experts Layer
  4. Capacity-Aware Inference: Mitigating the Straggler Effect in Mixture of Experts