| 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) |
|
|
| 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) |
| |
| 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_logits = self.gate(x_flat) |
| router_probs = F.softmax(router_logits, dim=-1) |
|
|
| |
| |
| capacity = int((num_tokens * self.active_experts / self.n_experts) * self.capacity_factor) |
| capacity = max(1, capacity) |
|
|
| |
| scores_t = router_probs.transpose(0, 1) |
|
|
| |
| topk_scores, topk_indices = torch.topk(scores_t, k=capacity, dim=1) |
|
|
| return topk_indices, topk_scores, router_probs |