跳到主要内容
极客日志极客日志面向AI+效率的开发者社区
首页博客我的书AI学习GitHub 精选镜像AI 生图工具UI配色美学关于
搜索内容 / 工具 / 仓库 / 镜像...⌘K搜索
注册
博客列表
PythonAI算法

混合专家网络 MOE 技术原理与代码实战

混合专家网络 MOE 通过门控网络对多个专家进行加权或选择,解决多任务学习中的平衡与泛化问题。文章解析了 MoE 核心原理,对比传统 DNN 优缺点,并结合推荐场景展示业务建模思路。提供基于 PyTorch 的完整代码实现,涵盖模型构建、训练循环及推理测试,帮助开发者落地 MoE 技术至大模型或推荐系统。

FlinkHero发布于 2026/3/22更新于 2026/10/791 浏览
混合专家网络 MOE 技术原理与代码实战

文章配图

一、引言

经历了大模型技术的快速发展,MoE(Mixture-of-Experts)作为核心架构之一,在 DeepSeek-v3 等大模型中展现了极低的推理成本与优异的效果。

1.1 本文侧重点

本文重点在于从代码级认识 MoE 混合专家网络技术,而非单纯讨论训练与推理细节。目标是带大家实现一个 MoE 网络,了解其构建方式,以便根据业务场景创新性地构建自己的专家网络。

1.2 技术洞察—MoE

MoE(Mixture-of-Experts)在近 7-8 年间已广泛应用于推荐系统多任务学习,以 MMoE(Google, 2018)、PLE(腾讯,2020)为基石,通过门控网络为多个专家网络加权平均,解决多目标、多场景问题。近 1-2 年间,基于 MoE 思想构建的大模型层出不穷,如 DeepSeekMoE、Mixtral 8x7B、Flan-MoE 等,通过路由网络对多个专家进行选择,提升推理效率。

二、MoE(Mixture-of-Experts,混合专家网络)

2.1 技术原理

MoE 全称为混合专家网络,主要由多个专家网络、多个任务塔、门控网络构成。核心原理如下:

  1. 样本数据输入:分别输入 num_experts 个专家网络进行推理。每个专家网络实际上是一个前馈神经网络(MLP),输入维度为 x,输出维度为 output_experts_dim。
  2. 门控网络:样本数据同时输入门控网络(也是 MLP),输出为 num_experts 个专家的概率分布,维度为 num_experts。采用 softmax 将输出归一化,各个维度加起来和为 1。
  3. 加权平均:将每个专家网络的输出,基于 gate 门控网络的 softmax 加权平均,作为 Task 的输入。Task 的输入统一维度均为 output_experts_dim。
  4. 参数更新:在每次反向传播迭代时,对 Gate 和 num_experts 个专家参数进行更新,Gate 和专家网络的参数受任务 A、B 共同影响。

文章配图

**专家网络:**样本数据分别输入 num_experts 个专家网络进行推理,每个专家网络实际上是一个前馈神经网络(MLP),输入维度为 x,输出维度为 output_experts_dim。 **门控网络:**样本数据输入门控网络,门控网络也是一个 MLP,输出为 num_experts 个 experts 专家的概率分布,维度为 num_experts(采用 softmax 将输出归一化,各个维度加起来和为 1)。 **任务网络:**将每个专家网络的输出,基于 gate 门控网络的 softmax 加权平均,作为 Task 的输入,Task 的输入统一维度均为 output_experts_dim。

2.2 技术优缺点

相较于传统的 DNN 网络,MoE 的本质是通过多个专家网络对预估任务共同决策,引入 Gate 作为专家的裁判,给每一个专家打分,判定哪个专家更加权威。(DeepSeekMoE 的 Router 与 Gate 类似,区别是 Gate 为每一个专家赋分,加权平均,Router 对专家进行选择,推理速度更快)。

**优点:**多个 DNN 专家网络投票共同决定推理结果,相较于单个 DNN 网络泛化性更好,准确率更高。Gate 网络基于多个 Task 任务进行反馈收敛,可以学到多个 Task 任务数据的平衡性。

**缺点:**朴素的 MoE 仅使用了一个 Gate 网络,虽然 Gate 网络由多个 Task 任务共同收敛学习得到,具有一定的平衡性,但对于每个 Task 的个性化能力仍然不足。(Google 针对此缺点发布了 MMoE)。底层多个专家网络均为共享专家,输入均为样本数据,参数的差异主要由初始化的不同得到,并不具备特异性。(腾讯针对此缺点发布了 PLE)。输入 Input 均为全部样本数据,学不出不同场景任务的差异性,需要在输入层对场景特征进行拆分(阿里针对此缺点发布了 ESMM)。

2.3 业务代码实践

2.3.1 业务场景与建模

我们以小红书推荐场景为例,用户在一级发现页场景中停留并点击了'误杀 3'中的一个视频笔记,在二级场景视频播放页中观看并点赞了视频。

我们构建一个 100 维特征输入,4 个 experts 专家网络,2 个 task 任务的,1 个门控的 MoE 网络,用于建模跨场景多任务学习问题,模型架构图如下:

文章配图

文章配图

如架构图所示,其中有几个注意的点:

**num_experts:**门控 gate 的输出维度和专家数相同,均为 num_experts,因为 gate 的用途是对专家网络最后一层进行加权平均,gate 维度与专家数是直接对应关系。 **output_experts_dim:**专家网络的输出维度和 task 网络的输入维度相同,task 网络承接的是专家网络各维度的加权平均值,experts 网络与 task 网络是直接对应关系。 **Softmax:**Gate 门控网络对最后一层采用 Softmax 归一化,保证专家网络加权平均后值域相同。

2.3.2 模型代码实现

基于 PyTorch,实现上述网络架构,如下:

import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader, TensorDataset

class MoEModel(nn.Module):
    def __init__(self, input_dim, experts_hidden1_dim, experts_hidden2_dim, output_experts_dim, task_hidden1_dim, task_hidden2_dim, output_task1_dim, output_task2_dim, gate_hidden1_dim, gate_hidden2_dim, num_experts):
        super(MoEModel, self).__init__()
        self.num_experts = num_experts
        self.output_experts_dim = output_experts_dim
        
        # 初始化多个专家网络
        self.experts = nn.ModuleList([
            nn.Sequential(
                nn.Linear(input_dim, experts_hidden1_dim),
                nn.ReLU(),
                nn.Linear(experts_hidden1_dim, experts_hidden2_dim),
                nn.ReLU(),
                nn.Linear(experts_hidden2_dim, output_experts_dim),
                nn.ReLU()
            ) for _ in range(num_experts)
        ])
        
        # 定义任务 1 的输出层
        self.task1_head = nn.Sequential(
            nn.Linear(output_experts_dim, task_hidden1_dim),
            nn.ReLU(),
            nn.Linear(task_hidden1_dim, task_hidden2_dim),
            nn.ReLU(),
            nn.Linear(task_hidden2_dim, output_task1_dim),
            nn.Sigmoid()
        )
        
        # 定义任务 2 的输出层
        self.task2_head = nn.Sequential(
            nn.Linear(output_experts_dim, task_hidden1_dim),
            nn.ReLU(),
            nn.Linear(task_hidden1_dim, task_hidden2_dim),
            nn.ReLU(),
            nn.Linear(task_hidden2_dim, output_task2_dim),
            nn.Sigmoid()
        )
        
        # 初始化门控网络
        self.gating_network = nn.Sequential(
            nn.Linear(input_dim, gate_hidden1_dim),
            nn.ReLU(),
            nn.Linear(gate_hidden1_dim, gate_hidden2_dim),
            nn.ReLU(),
            nn.Linear(gate_hidden2_dim, num_experts),
            nn.Softmax(dim=1)
        )

    def forward(self, x):
        # 计算输入数据通过门控网络后的权重
        gates = self.gating_network(x)
        batch_size, _ = x.shape
        task1_inputs = torch.zeros(batch_size, self.output_experts_dim)
        task2_inputs = torch.zeros(batch_size, self.output_experts_dim)
        
        # 计算每个专家的输出并加权求和
        for i in range(self.num_experts):
            expert_output = self.experts[i](x)
            task1_inputs += expert_output * gates[:, i].unsqueeze(1)
            task2_inputs += expert_output * gates[:, i].unsqueeze(1)
            
        task1_outputs = self.task1_head(task1_inputs)
        task2_outputs = self.task2_head(task2_inputs)
        return task1_outputs, task2_outputs

# 实例化模型对象
num_experts = 4
experts_hidden1_dim = 64
experts_hidden2_dim = 32
output_experts_dim = 16
gate_hidden1_dim = 16
gate_hidden2_dim = 8
task_hidden1_dim = 32
task_hidden2_dim = 16
output_task1_dim = 3
output_task2_dim = 2

# 构造虚拟样本数据
torch.manual_seed(42)
input_dim = 10
num_samples = 1024
X_train = torch.randint(0, 2, (num_samples, input_dim)).float()
y_train_task1 = torch.rand(num_samples, output_task1_dim)
y_train_task2 = torch.rand(num_samples, output_task2_dim)

# 创建数据加载器
train_dataset = TensorDataset(X_train, y_train_task1, y_train_task2)
train_loader = DataLoader(train_dataset, batch_size=128, shuffle=True)

model = MoEModel(input_dim, experts_hidden1_dim, experts_hidden2_dim, output_experts_dim, task_hidden1_dim, task_hidden2_dim, output_task1_dim, output_task2_dim, gate_hidden1_dim, gate_hidden2_dim, num_experts)

# 定义损失函数和优化器
criterion_task1 = nn.MSELoss()
criterion_task2 = nn.MSELoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)

# 训练循环
num_epochs = 100
for epoch in range(num_epochs):
    model.train()
    running_loss = 0.0
    for batch_idx, (X_batch, y_task1_batch, y_task2_batch) in enumerate(train_loader):
        outputs_task1, outputs_task2 = model(X_batch)
        loss_task1 = criterion_task1(outputs_task1, y_task1_batch)
        loss_task2 = criterion_task2(outputs_task2, y_task2_batch)
        total_loss = loss_task1 + loss_task2
        optimizer.zero_grad()
        total_loss.backward()
        optimizer.step()
        running_loss += total_loss.item()
    if epoch % 10 == 0:
        print(f'Epoch [{epoch+1}/{num_epochs}], Loss: {running_loss/len(train_loader):.4f}')

print(model)
for param_tensor in model.state_dict():
    print(param_tensor, "\t", model.state_dict()[param_tensor].size())

# 模型预测
model.eval()
with torch.no_grad():
    test_input = torch.randint(0, 2, (1, input_dim)).float()
    pred_task1, pred_task2 = model(test_input)
    print(f'一级场景预测结果:{pred_task1}')
    print(f'二级场景预测结果:{pred_task2}')
2.3.3 模型训练与推理测试

运行上述代码,模型启动训练,Loss 逐渐收敛,测试结果如下:

文章配图

2.3.4 打印模型结构

使用 print(model) 打印模型结构如下:

文章配图

三、总结

本文从代码级讲解了 DeepSeek 大模型、MMoE 推荐模型中的 MoE(Mixture-of-Experts)技术。该技术的主要思想是通过门控(gate)或路由(router)网络,对多个专家进行加权平均或筛选,将一个 DNN 网络裂变为多个 DNN 网络后,投票决定预测结果。相较于单一的 DNN 网络,具有更强的容错性、泛化性与准确性,同时可以提高推理速度,节省推理资源。

技术洞察结论:MoE 技术未来将成为大模型和推荐系统进一步突破的关键技术,该技术为算法基础技术中的 SOTA。通过动手实现一个 MoE,再基于自己的业务场景,对齐专家网络、门控网络、任务网络进行创新,可更好地应用该技术。

目录

  1. 一、引言
  2. 1.1 本文侧重点
  3. 1.2 技术洞察—MoE
  4. 二、MoE(Mixture-of-Experts,混合专家网络)
  5. 2.1 技术原理
  6. 2.2 技术优缺点
  7. 2.3 业务代码实践
  8. 2.3.1 业务场景与建模
  9. 2.3.2 模型代码实现
  10. 实例化模型对象
  11. 构造虚拟样本数据
  12. 创建数据加载器
  13. 定义损失函数和优化器
  14. 训练循环
  15. 模型预测
  16. 2.3.3 模型训练与推理测试
  17. 2.3.4 打印模型结构
  18. 三、总结

更多推荐文章

查看全部
  • Mac Mini 部署 OpenClaw 实战指南
  • AI 大模型迈入应用时代,推动“可控大模型”落地
  • C 语言快速排序详解:从基础到非递归实现
  • 李飞飞最新论文:我们需要什么样的 AI Agent
  • Flutter 集成 genkit 实现鸿蒙端 AI 流式响应与提示词工程
  • Whisper 模型国内镜像源汇总与快速下载指南
  • 树莓派 4B 连接大疆 M300 RTK 无人机 PSDK 开发指南
  • LeetCode 原地复写零:双指针 + 逆向填充实现 O(n) 时间 O(1) 空间
  • LiuJuan Z-Image Generator 本地部署与 8K 人像生成指南
  • 大规模语言模型:从理论到实践的模型训练
  • 机器人表情模拟实现:Arduino 控制面部舵机项目详解
  • C++ 设计模式在面向对象开发中的应用
  • MySQL 8 核心日志与备份恢复详解
  • Flutter 在 OpenHarmony 中使用 nanoid 替代 UUID 生成唯一标识
  • 2026年最新全球AI大模型深度研究报告
  • 深度学习入门实战:从基础概念到手写数字识别
  • Flutter 底部导航与 TabBar 多页切换实战及状态保持
  • 40 道 Python 经典面试题及参考答案
  • Vue2 纯前端对接海康威视摄像头实时视频预览
  • Python 构建跨平台前端界面:Flet 库详解

相关免费在线工具

  • 加密/解密文本

    使用加密算法(如AES、TripleDES、Rabbit或RC4)加密和解密文本明文。 在线工具,加密/解密文本在线工具,online

  • RSA密钥对生成器

    生成新的随机RSA私钥和公钥pem证书。 在线工具,RSA密钥对生成器在线工具,online

  • Mermaid 预览与可视化编辑

    基于 Mermaid.js 实时预览流程图、时序图等图表,支持源码编辑与即时渲染。 在线工具,Mermaid 预览与可视化编辑在线工具,online

  • 随机西班牙地址生成器

    随机生成西班牙地址(支持马德里、加泰罗尼亚、安达卢西亚、瓦伦西亚筛选),支持数量快捷选择、显示全部与下载。 在线工具,随机西班牙地址生成器在线工具,online

  • Gemini 图片去水印

    基于开源反向 Alpha 混合算法去除 Gemini/Nano Banana 图片水印,支持批量处理与下载。 在线工具,Gemini 图片去水印在线工具,online

  • curl 转代码

    解析常见 curl 参数并生成 fetch、axios、PHP curl 或 Python requests 示例代码。 在线工具,curl 转代码在线工具,online