以下代码实现了 MoE(混合专家模型)的路由机制和稀疏激活,请阅读代码并填写空缺部分。 [1] 和 [2] 处应分别填入:
import torch
import torch.nn as nn
import torch.nn.functional as F
class MoELayer(nn.Module):
def __init__(self, input_dim, hidden_dim, num_experts, top_k=2):
super().__init__()
self.num_experts = num_experts
self.top_k = top_k
# 路由网络:将输入映射到专家概率分布
self.gate = nn.Linear(input_dim, num_experts, bias=False)
# 多个专家网络
self.experts = nn.ModuleList([
nn.Sequential(
nn.Linear(input_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, input_dim)
) for _ in range(num_experts)
])
def forward(self, x):
# x: (batch_size, seq_len, input_dim)
batch_size, seq_len, dim = x.shape
x_flat = x.view(-1, dim) # (batch*seq, dim)
# [1] 通过路由网络计算每个 token 对应各专家的分数,并用 softmax 归一化
router_logits = self.gate(x_flat) # (batch*seq, num_experts)
____________________________________________________________
# [2] 选择 Top-K 个专家
_______________________________________________________________
# 归一化 top-k 权重
top_k_weights = top_k_weights / top_k_weights.sum(dim=-1, keepdim=True)
# 初始化输出张量
final_output = torch.zeros_like(x_flat)
# 对每个专家处理对应的 token
for k in range(self.top_k):
expert_weight = top_k_weights[:, k:k+1] # (batch*seq, 1)
expert_idx = top_k_indices[:, k] # (batch*seq,)
# 对每个专家分别处理
for expert_id in range(self.num_experts):
mask = (expert_idx == expert_id)
if mask.any():
expert_input = x_flat[mask]
expert_output = self.experts[expert_id](expert_input)
final_output[mask] += expert_weight[mask] * expert_output
return final_output.view(batch_size, seq_len, dim)
[1] router_probs = F.softmax(router_logits, dim=-1)
[2] top_k_weights, top_k_indices = torch.topk(router_probs, self.top_k, dim=-1)
[1] router_probs = F.sigmoid(router_logits)
[2] top_k_weights, top_k_indices = torch.topk(router_probs, self.top_k, dim=-1)
[1] router_probs = F.softmax(router_logits, dim=0)
[2] top_k_weights, top_k_indices = torch.sort(router_probs, dim=-1)
[1] router_probs = F.relu(router_logits)
[2] top_k_weights, top_k_indices = torch.topk(router_probs, self.num_experts, dim=-1)