From e193d5f9628fb2a7fb30951da7a2511004300f63 Mon Sep 17 00:00:00 2001 From: iLeGend <824040212@qq.com> Date: Fri, 24 Jul 2026 18:41:27 +0800 Subject: [PATCH] Enable support for Tutel when k != 1 for shared moe Signed-off-by: iLeGend <824040212@qq.com> --- deepspeed/moe/sharded_moe.py | 40 ++++++++++++++++++++++++++++-------- tests/unit/moe/test_moe.py | 29 ++++++++++++++++++++++++++ 2 files changed, 61 insertions(+), 8 deletions(-) diff --git a/deepspeed/moe/sharded_moe.py b/deepspeed/moe/sharded_moe.py index d2a6c089e8e7..341991ec0959 100644 --- a/deepspeed/moe/sharded_moe.py +++ b/deepspeed/moe/sharded_moe.py @@ -293,7 +293,8 @@ def top2gating(logits: Tensor, min_capacity: int, drop_tokens: bool = True, ep_group: Union[torch.distributed.ProcessGroup, None] = None, - top2_2nd_expert_sampling: bool = True) -> Tuple[Tensor, Tensor, Tensor, Tensor]: + top2_2nd_expert_sampling: bool = True, + use_tutel: bool = False) -> Tuple[Tensor, Tensor, Tensor, Tensor]: """Implements Top2Gating on logits.""" # everything is in fp32 in this function gates = F.softmax(logits, dim=1) @@ -359,6 +360,12 @@ def top2gating(logits: Tensor, gates1_s /= denom_s gates2_s /= denom_s + if use_tutel: + indices1_s = torch.where(mask1.sum(dim=1).bool(), indices1_s, torch.full_like(indices1_s, -1)) + indices2_s = torch.where(mask2.sum(dim=1).bool(), indices2_s, torch.full_like(indices2_s, -1)) + return l_aux, capacity, num_experts, [indices1_s, indices2_s], [locations1_s, + locations2_s], [gates1_s, gates2_s], exp_counts + # Calculate combine_weights and dispatch_mask gates1 = einsum("s,se->se", gates1_s, mask1_float) gates2 = einsum("s,se->se", gates2_s, mask2_float) @@ -380,6 +387,7 @@ def topkgating( drop_tokens: bool = True, ep_group: Union[torch.distributed.ProcessGroup, None] = None, drop_policy: str = "probs", + use_tutel: bool = False, ) -> Tuple[Tensor, Tensor, Tensor, Tensor]: """Implements TopKGating on logits.""" @@ -439,6 +447,20 @@ def topkgating( if locations is None: raise ValueError(f"Locations is not set: {locations}") + + if use_tutel: + indices_ = [] + locations_ = [] + gates_ = [] + for route in range(k): + indices_s = top_idx[:, route] + route_mask = F.one_hot(indices_s, num_classes=num_experts).bool() & mask + indices_s = torch.where(route_mask.any(dim=1), indices_s, torch.full_like(indices_s, -1)) + indices_.append(indices_s) + locations_.append(torch.sum(locations * route_mask, dim=1)) + gates_.append(torch.sum(gates_masked * route_mask, dim=1)) + return l_aux, capacity, num_experts, indices_, locations_, gates_, exp_counts + # dispatch_mask locations_sc = _one_hot_to_float((locations * mask), capacity) @@ -520,11 +542,16 @@ def forward(self, elif self.k == 2: gate_output = top2gating(logits, self.capacity_factor if self.training else self.eval_capacity_factor, - self.min_capacity, self.drop_tokens, self.ep_group, self.top2_2nd_expert_sampling) + self.min_capacity, self.drop_tokens, self.ep_group, self.top2_2nd_expert_sampling, + use_tutel) else: - gate_output = topkgating(logits, self.k, + gate_output = topkgating(logits, + self.k, self.capacity_factor if self.training else self.eval_capacity_factor, - self.min_capacity, self.drop_tokens, self.ep_group) + self.min_capacity, + self.drop_tokens, + self.ep_group, + use_tutel=use_tutel) if self.wall_clock_breakdown: self.timers(TOPK_GATE_TIMER).stop() @@ -571,16 +598,13 @@ def __init__(self, self.timers = SynchronizedWallClockTimer() self.wall_clock_breakdown = False - self.use_tutel = use_tutel and TUTEL_INSTALLED and gate.k == 1 + self.use_tutel = use_tutel and TUTEL_INSTALLED if self.use_tutel: logger.info('Using Tutel optimizations.') elif use_tutel and not TUTEL_INSTALLED: logger.warning("Tutel optimization requested but not installed. " "Proceeding without Tutel.") - elif use_tutel and TUTEL_INSTALLED and gate.k != 1: - logger.warning("To enable Tutel optimization, use top-1 instead of top-2 gate. " - "Proceeding without Tutel.") def _set_ep_group(self, ep_group): self.ep_group = ep_group diff --git a/tests/unit/moe/test_moe.py b/tests/unit/moe/test_moe.py index 6283007d3e8e..6b250d74302e 100644 --- a/tests/unit/moe/test_moe.py +++ b/tests/unit/moe/test_moe.py @@ -300,6 +300,35 @@ def fail_collective(*args, **kwargs): class TestTopkGate(DistributedTest): + @staticmethod + def _assert_tutel_routes_match(combine_weights, tutel_output): + _, capacity, num_experts, indices_, locations_, gates_, _ = tutel_output + reconstructed = torch.zeros_like(combine_weights) + for indices_s, locations_s, gates_s in zip(indices_, locations_, gates_): + route_mask = indices_s.ge(0) + safe_indices = indices_s.clamp_min(0) + expert_mask = torch.nn.functional.one_hot(safe_indices, num_classes=num_experts) + expert_mask *= route_mask.unsqueeze(1) + location_mask = torch.nn.functional.one_hot(locations_s, num_classes=int(capacity)) + reconstructed += torch.einsum("s,se,sc->sec", gates_s, expert_mask, location_mask) + torch.testing.assert_close(reconstructed, combine_weights) + + def test_tutel_top2_routing(self): + logits = torch.tensor([[0.1, 0.9, 0.2], [0.8, 0.1, 0.3], [0.4, 0.6, 0.5], [0.7, 0.2, 0.1]]) + native_output = top2gating(logits, 0.5, 0, top2_2nd_expert_sampling=False) + tutel_output = top2gating(logits, 0.5, 0, top2_2nd_expert_sampling=False, use_tutel=True) + + assert len(tutel_output[3]) == 2 + self._assert_tutel_routes_match(native_output[1], tutel_output) + + def test_tutel_topk_routing(self): + logits = torch.tensor([[0.1, 0.9, 0.2, 0.3], [0.8, 0.1, 0.3, 0.2], [0.4, 0.6, 0.5, 0.7], [0.7, 0.2, 0.1, 0.8]]) + native_output = topkgating(logits, 3, 0.5, 0) + tutel_output = topkgating(logits, 3, 0.5, 0, use_tutel=True) + + assert len(tutel_output[3]) == 3 + self._assert_tutel_routes_match(native_output[1], tutel_output) + def test(self): def check_equal(logits, cap, sparse_truth, res):