Outlier-10B-V3.2 / modeling_outlier_moe.py
ur-dad-matt's picture
fix: rename Qwen2MoE → OutlierMoE to avoid transformers 4.46+ Qwen2MoeForCausalLM namespace collision (Mac-launch-006) [modeling]
0cbc2ad verified
Raw
History Blame Contribute Delete
3.89 kB
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,
)