File size: 5,677 Bytes
d91766b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 | from __future__ import annotations
import os
import torch
from einops import rearrange, repeat
from diffulex.moe.topk.base import TopKRouter
from diffulex.moe.topk.output import TopKOutput
from diffulex_kernel import fused_group_limited_topk
def _tile_kernels_group_limited_topk(
scores: torch.Tensor,
*,
top_k: int,
n_group: int,
topk_group: int,
) -> torch.Tensor | None:
try:
import tile_kernels # type: ignore
except Exception:
return None
if (
n_group <= 0
or topk_group <= 0
or n_group > int(scores.shape[-1])
or int(scores.shape[-1]) % n_group != 0
):
try:
return tile_kernels.moe.topk_gate(scores, top_k)
except Exception:
return None
experts_per_group = int(scores.shape[-1]) // n_group
scores_by_group = rearrange(scores, "token (group expert) -> token group expert", group=n_group)
num_group_sum_topk = min(2, experts_per_group)
num_topk_groups = min(topk_group, n_group)
try:
group_idx = tile_kernels.moe.topk_sum_and_topk_group_idx(
scores_by_group,
num_group_sum_topk,
num_topk_groups,
)
group_mask = torch.zeros(
(scores.shape[0], n_group),
dtype=torch.bool,
device=scores.device,
)
group_mask.scatter_(1, group_idx.to(torch.int64), True)
score_mask = repeat(group_mask, "token group -> token (group expert)", expert=experts_per_group)
masked_scores = scores.masked_fill(~score_mask, float("-inf"))
return tile_kernels.moe.topk_gate(masked_scores, top_k)
except Exception:
return None
class GroupLimitedTopKRouter(TopKRouter):
def __init__(
self,
top_k: int,
*,
num_experts: int,
n_group: int,
topk_group: int,
routed_scaling_factor: float = 1.0,
renormalize: bool = True,
scoring_func: str = "sigmoid",
expert_bias_getter,
) -> None:
super().__init__(top_k=top_k, renormalize=renormalize, scoring_func=scoring_func)
if scoring_func != "sigmoid":
raise NotImplementedError("GroupLimitedTopKRouter currently supports sigmoid scoring only.")
self.num_experts = num_experts
self.n_group = n_group
self.topk_group = topk_group
self.routed_scaling_factor = routed_scaling_factor
self._expert_bias_getter = expert_bias_getter
def _group_limited_topk(self, scores: torch.Tensor) -> torch.Tensor:
if (
self.n_group <= 0
or self.topk_group <= 0
or self.n_group > self.num_experts
or self.num_experts % self.n_group != 0
):
return torch.topk(scores, k=self.top_k, dim=-1, sorted=False).indices
experts_per_group = self.num_experts // self.n_group
scores_by_group = rearrange(
scores,
"token (group expert) -> token group expert",
group=self.n_group,
)
group_scores = scores_by_group.topk(
min(2, experts_per_group),
dim=-1,
).values.sum(dim=-1)
group_idx = torch.topk(group_scores, k=min(self.topk_group, self.n_group), dim=-1, sorted=False).indices
group_mask = torch.zeros_like(group_scores, dtype=torch.bool)
group_mask.scatter_(1, group_idx, True)
score_mask = repeat(
group_mask,
"token group -> token (group expert)",
expert=experts_per_group,
)
masked_scores = scores.masked_fill(~score_mask, float("-inf"))
return torch.topk(masked_scores, k=self.top_k, dim=-1, sorted=False).indices
def _forward_naive(self, router_logits: torch.Tensor) -> TopKOutput:
scores = torch.sigmoid(router_logits.float()).to(router_logits.dtype)
expert_bias = self._expert_bias_getter().to(scores.device, dtype=scores.dtype)
rank_scores = scores + expert_bias
topk_ids = None
if (
router_logits.is_cuda
and os.getenv("DIFFULEX_MOE_TOPK_IMPL", "").lower() in {"tile", "tilekernels", "tile_kernels"}
):
topk_ids = _tile_kernels_group_limited_topk(
rank_scores,
top_k=self.top_k,
n_group=self.n_group,
topk_group=self.topk_group,
)
if topk_ids is None:
topk_ids = self._group_limited_topk(rank_scores)
topk_ids = topk_ids.to(torch.int64)
topk_weights = torch.gather(scores, dim=-1, index=topk_ids)
if self.renormalize and self.top_k > 1:
topk_weights = topk_weights / (topk_weights.sum(dim=-1, keepdim=True) + 1e-20)
topk_weights = topk_weights * self.routed_scaling_factor
return TopKOutput(weights=topk_weights, ids=topk_ids, router_logits=router_logits)
def forward(self, router_logits: torch.Tensor) -> TopKOutput:
if not router_logits.is_cuda or os.getenv("DIFFULEX_REFERENCE_MOE_ROUTER", "0") == "1":
return self._forward_naive(router_logits)
expert_bias = self._expert_bias_getter().to(router_logits.device, dtype=router_logits.dtype)
topk_weights, topk_ids = fused_group_limited_topk(
router_logits=router_logits,
expert_bias=expert_bias,
top_k=self.top_k,
n_group=self.n_group,
topk_group=self.topk_group,
routed_scaling_factor=self.routed_scaling_factor,
renormalize=self.renormalize,
)
return TopKOutput(weights=topk_weights, ids=topk_ids, router_logits=router_logits)
|