mixture-of-experts
PyTorch版稀疏门控MoE层实现,复现Google经典论文,支持Top-K专家激活与负载均衡辅助
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
PyTorch版稀疏门控MoE层实现,复现Google经典论文,支持Top-K专家激活与负载均衡辅助
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
2017年,Google Brain 的 Noam Shazeer 等人发表了一篇足以写进深度学习教材的论文——Outrageously Large Neural Networks: The Sparsely-Gated Mixture-of-Experts Layer。论文的核心洞察是:传统神经网络中,所有参数对每一个输入都会参与计算——这是巨大的浪费。一个1000亿参数的模型,处理一条"今天天气怎么样"的文本时,可能只有1%的参数真正做出了贡献。
Shazeer 的思路如同一个医院分诊台:为每类疾病分配专科医生,患者进门后由智能分诊员判断"你该去哪个诊室",每个医生只处理自己擅长的情况。这既保证了能力上限(专科医生够专业),又控制了计算成本(不是所有人都要全楼上岗)。这篇论文奠定了现代 MoE 的理论基础。
David Rau 在 2019 年将 Google 原版 TensorFlow 实现"翻译"为 PyTorch 版本,便诞生了这个项目——davidmrau/mixture-of-experts。代码仅 6 个文件(moe.py、example.py、requirements.py 等),却完整复现了论文中稀疏门控的核心机制。此后 FastMoE 等知名项目都将它作为单 GPU 训练的标准参考实现。
传统 Dense 层像一个全勤医院:所有科室医生每天都来上班,不管有没有患者。稀疏门控则是一个按需调配的弹性团队——只有被选中的专家才会处理当前输入。
以本项目为例:num_experts=10、k=4 表示共有 10 个专家网络,每次前向传播只会激活(dispatch)其中 Top-4 个。输入 batch 中每条样本的分配决策由一个轻量级的 Gating Network(即 w_gate 参数矩阵)决定,输出形状为 [batch_size, num_experts] 的概率分布。
这是项目最精妙的设计。SparseDispatcher 负责两件事:
output[b] = Σᵢ(gate[b,i] × expert_i(input[b]))
通过索引和 torch.split/torch.index_add,整个过程完全在 GPU 上向量化,无需 Python 循环遍历专家,效率极高。
如果门控完全确定性(直接取 Top-k),会导致"富者愈富"问题:热门专家持续被选中、越来越忙,冷门专家越来越闲、梯度越来越少,最终模型退化为普通层(只有少数专家真正被训练)。
解决方案是在门控层加入可学习的噪声(_prob_in_top_k 函数),模拟一个"抽签"机制:即使某个专家 logits 稍低,也有一定概率被选中,使负载分布更均衡。这是 2017 年论文的原创设计。
项目仅 4 个 Python 文件,总计约 600 行,结构清晰:
| 文件 | 职责 |
|---|---|
moe.py (~350行) | 核心实现:SparseDispatcher、MLP、MoE 三个类 |
example.py (~80行) | Dummy 数据集训练/评估演示 |
cifar10_example.py (~120行) | CIFAR-10 真实数据集示例 |
requirements.py | 唯一依赖:torch |
class MoE(nn.Module):
def __init__(self, input_size, output_size, num_experts, hidden_size,
noisy_gating=True, k=4):
# 10个独立的 MLP expert
self.experts = nn.ModuleList([
MLP(input_size, output_size, hidden_size)
for _ in range(num_experts)
])
# 可学习的门控网络(核心参数)
self.w_gate = nn.Parameter(torch.zeros(input_size, num_experts))
self.w_noise = nn.Parameter(torch.zeros(input_size, num_experts))
门控网络和专家网络是联合训练的,但门控的随机性可能导致某些专家长期不被调用。MoE 的 forward 返回两个值:预测输出 + 辅助损失(Auxiliary Loss),由两部分组成:
def forward(self, x, loss_coef=1e-2):
gates, load = self.noisy_top_k_gating(x, self.training)
importance = gates.sum(0) # 各专家被调用概率之和
aux_loss = self.cv_squared(importance) + self.cv_squared(load)
# ... dispatcher dispatch + expert forward + combine ...
return y, aux_loss * loss_coef
总损失 = 任务损失 + loss_coef × aux_loss,loss_coef 默认 0.01 是经验值——太小均衡效果弱,太大则任务损失被稀释。
README 坦言:CIFAR-10 示例仅用随机超参数跑了跑,准确率 39%。这并不意味着 MoE 在 CIFAR-10 上效果差,而是说明这是一个研究原型代码库,不追求某个具体任务的 SOTA。FastMoE 等项目将其作为单 GPU baseline,后续做分布式扩展才更有意义。
from moe import MoE
import torch
model = MoE(input_size=1000, output_size=20, num_experts=10,
hidden_size=66, k=4, noisy_gating=True)
x = torch.rand(32, 1000)
y_hat, aux_loss = model(x) # y_hat: [32, 20], aux_loss: scalar
total_loss = task_loss + aux_loss
这个代码库是 MoE 研究的"Hello World",后续扩展方向包括:
这个 2019 年的小项目,是现代 MoE 浪潮的重要节点:
| 年份 | 里程碑事件 | 与本项目的关系 |
|---|---|---|
| 2017 | Google 论文发布 | 原始理论奠基 |
| 2019 | David Rau PyTorch 重实现 | 本项目,降低使用门槛 |
| 2021 | FastMoE 论文(北大) | 将本项目作为单 GPU baseline |
| 2022 | Google Switch Transformer | 首次将 MoE 扩展到万亿参数 |
| 2023 | Mistral Mixtral 8x7B | 工业级 MoE 大模型爆发 |
| 2024 | GPT-4o、DBRX、Grok-1 | MoE 成为大模型标配架构 |
作为 PyTorch 生态中最早期的 MoE 参考实现之一,本项目的价值不在于刷榜,而在于降低学术复现门槛——无数研究者通过阅读它的代码理解了稀疏门控的核心逻辑,进而在此基础上做创新。