PyTorch
Spanish
English
Mixture of Experts
mamba
ssm
reasoning
chain-of-thought
Aethelred-7B-v2 / router.py
J4HDx's picture
Create router.py
59fa7ba verified
Raw
History Blame Contribute Delete
2.45 kB
import torch
import torch.nn as nn
import torch.nn.functional as F
class AdaptiveHybridRouter(nn.Module):
"""
Router that uses Gumbel-Softmax to decide between Mamba-2 (SSD)
and Sliding Window Attention (Flash-3).
"""
def __init__(self, d_model: int):
super().__init__()
self.router_net = nn.Linear(d_model, 2) # 0: Mamba, 1: SWA
def forward(self, x, tau=1.0, hard=False):
"""
x: (batch, seq_len, d_model)
returns: (batch, seq_len, 2) weights
"""
logits = self.router_net(x)
# Gumbel-Softmax for differentiable switching
weights = F.gumbel_softmax(logits, tau=tau, hard=hard, dim=-1)
return weights
class ExpertChoiceRouterV2(nn.Module):
"""
Expert Choice (EC) v2 Router.
In EC, experts choose top-k tokens instead of tokens choosing top-k experts.
This guarantees perfect load balancing.
v2 includes entropy-based regularization to avoid expert collapse.
"""
def __init__(self, d_model: int, n_experts: int, active_experts: int, capacity_factor: float = 1.2):
super().__init__()
self.d_model = d_model
self.n_experts = n_experts
self.active_experts = active_experts
self.capacity_factor = capacity_factor
self.gate = nn.Linear(d_model, n_experts, bias=False)
def forward(self, x):
"""
x: (batch, seq_len, d_model)
returns:
topk_indices: (n_experts, capacity) - indices of tokens chosen by each expert
topk_scores: (n_experts, capacity)
router_probs: (batch * seq_len, n_experts)
"""
batch_size, seq_len, _ = x.shape
num_tokens = batch_size * seq_len
x_flat = x.view(num_tokens, self.d_model)
# Router scores: (num_tokens, n_experts)
router_logits = self.gate(x_flat)
router_probs = F.softmax(router_logits, dim=-1)
# Expert Choice: each expert chooses its favorite k tokens
# capacity = (num_tokens * active_experts) // n_experts
capacity = int((num_tokens * self.active_experts / self.n_experts) * self.capacity_factor)
capacity = max(1, capacity)
# Scores: (n_experts, num_tokens)
scores_t = router_probs.transpose(0, 1)
# Top-k tokens for each expert
topk_scores, topk_indices = torch.topk(scores_t, k=capacity, dim=1)
return topk_indices, topk_scores, router_probs