Qwen 全局批次负载均衡用于 MoE LLM 训练

全局批次负载均衡提升 MoE 性能和专家专业化

Qwen 提出了一种用于训练混合专家(MoE)大型语言模型(LLM)的全局批次负载均衡损失方法。通过将负载均衡损失的计算从微批次层面转移到全局批次层面,模型在各种任务上实现了更好的性能,并在专家之间促进了显著的领域专业化,解决了微批次平衡常常阻止专家在特定数据领域进行专业化的限制。

微批次负载均衡的局限性

大多数现有的 MoE 训练框架,例如 Megatron-core,在微批次层面实现负载均衡损失($L_{ ext{balance}}$)。这意味着损失在每个微批次内部计算,然后在全局批次上取平均值。

当微批次缺乏多样化数据时,这种方法会出现问题。例如,如果一个微批次仅包含代码数据,微批次层面的损失会迫使路由器将这些代码标记均匀分布到所有专家上。这会阻止专家在特定领域(如代码)进行专业化,并可能损害整体模型性能。由于单个微批次的数据在 LLM 训练中往往来自同一领域,这解释了为什么许多开源 MoE 模型缺乏显著的专家专业化。

实现全局批次负载均衡

为了解决这个问题,Qwen 提议在全局批次层面计算负载均衡损失($L_{ ext{global}}$)。这是通过以下三个步骤实现的:

  1. 同步 所有并行组之间的专家选择频率($f_{i}$)
  2. 计算 每个并行组中的负载均衡损失(例如,在单个 GPU 上)
  3. 聚合 所有微批次之间的损失。

因为专家选择频率是一个小向量(维度等于专家数量),在微批次之间同步这些数据在计算上成本低廉,使得实现几乎是“免费”的。

性能提升和专家专业化

Qwen 在三种 MoE 配置(3.4B 激活 0.6B,15B 激活 2.5B,以及 43B 激活 6.6B)和两种数据配置(120B 和 400B 标记)上测试了该方法。

模型有效性提升

与微批次层面的损失相比,全局批次方法在所有测试设置中实现了更好的性能,包括所有模型规模、数据量和任务。

增强的领域专业化

全局批次平衡使专家能够进行专业化。虽然微批次平衡导致专家无论领域如何都被均匀激活,但全局批次平衡允许特定专家被特定领域频繁激活,从而展现出明显的专业化。

平衡批次大小的影响

在 3.4B 模型(激活 0.6B)上的实验表明,随着平衡批次大小从 2 增加到 128,ptr-training PPL(困惑度)迅速下降,在 128 后趋于饱和。这凸显了全局方法的重要性,因为主流 MoE 框架通常对较大模型使用 8 到 16 之间的平衡批次大小。

平衡效率和效果

虽然全局批次平衡可以提升模型质量,但它可能导致微批次平衡下降,从而对计算效率产生负面影响。

为了缓解这一点,Qwen 尝试在全局批次平衡损失之上添加一个微批次平衡损失(使用全局批次损失的常量权重 0.01)。这种混合方法将更新步骤的速度从 1.64 秒提升到 1.59 秒每更新步骤,而模型的有效性基本不受影响。

结论

通过实施全局批次负载均衡,Qwen 解决了 MoE 训练中的一个关键限制,使得专家能够表现出更高的性能和专业化。此方法为优化 MoE 模型提供了一种新视角,特别是在跨各种领域扩展到更大规模和更专业的 MoE 模型方面。

Sources