import torch import torch.nn as nn import torch.nn.functional as F from transformers.models.qwen2.modeling_qwen2 import Qwen2ForCausalLM, Qwen2MLP class Qwen2ExpertMLP(nn.Module): def __init__(self, config): super().__init__() self.hidden_size = config.hidden_size self.intermediate_size = config.intermediate_size self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False) self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False) self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False) self.act_fn = F.silu nn.init.zeros_(self.gate_proj.weight) nn.init.zeros_(self.up_proj.weight) nn.init.zeros_(self.down_proj.weight) def forward(self, x): return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x)) class Qwen2MoEMLP(nn.Module): def __init__(self, config, shared_mlp: Qwen2MLP, num_experts: int, top_k: int, alpha_values): super().__init__() object.__setattr__(self, "shared_mlp", shared_mlp) self.hidden_size = config.hidden_size self.intermediate_size = config.intermediate_size self.num_experts = num_experts self.top_k = top_k # Keep the dense/shared MLP weights under their original parameter names. self.gate_proj = shared_mlp.gate_proj self.up_proj = shared_mlp.up_proj self.down_proj = shared_mlp.down_proj self.act_fn = shared_mlp.act_fn self.router = nn.Linear(self.hidden_size, self.num_experts, bias=False) nn.init.zeros_(self.router.weight) self.experts = nn.ModuleList([Qwen2ExpertMLP(config) for _ in range(self.num_experts)]) alpha_tensor = torch.tensor(alpha_values, dtype=torch.float32) self.register_buffer("alpha_values", alpha_tensor, persistent=True) def forward(self, x): original_shape = x.shape x_flat = x.reshape(-1, original_shape[-1]) shared_out = self.shared_mlp(x_flat) logits = self.router(x_flat).float() probs = F.softmax(logits, dim=-1) top_k = min(self.top_k, probs.shape[-1]) weights, indices = torch.topk(probs, k=top_k, dim=-1) weights = weights / weights.sum(dim=-1, keepdim=True).clamp_min(1e-9) weights = weights.to(shared_out.dtype) expert_out = torch.zeros_like(shared_out) for expert_idx in torch.unique(indices).tolist(): assignment = indices == expert_idx token_mask = assignment.any(dim=-1) if not token_mask.any(): continue selected_weights = (weights * assignment.to(weights.dtype)).sum(dim=-1, keepdim=True) expert_result = self.experts[expert_idx](x_flat[token_mask]) alpha = self.alpha_values[expert_idx].to(expert_result.dtype) expert_out[token_mask] += selected_weights[token_mask] * alpha * expert_result return (shared_out + expert_out).reshape(original_shape) class OutlierMoEForCausalLM(Qwen2ForCausalLM): def __init__(self, config): super().__init__(config) moe_layer_indices = [int(layer) for layer in getattr(config, "moe_layer_indices", [])] experts_per_layer = int(getattr(config, "experts_per_layer", getattr(config, "n_experts", 0))) top_k = int(getattr(config, "top_k", 2)) alpha_values = getattr(config, "alpha_values", {}) for layer_idx in moe_layer_indices: layer = self.model.layers[layer_idx] per_layer_alpha = alpha_values.get(str(layer_idx), [0.0] * experts_per_layer) layer.mlp = Qwen2MoEMLP( config=config, shared_mlp=layer.mlp, num_experts=experts_per_layer, top_k=top_k, alpha_values=per_layer_alpha, )