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