diff --git a/python/sglang/srt/layers/moe/hash_topk.py b/python/sglang/srt/layers/moe/hash_topk.py index 0902403e6..1e13881dc 100644 --- a/python/sglang/srt/layers/moe/hash_topk.py +++ b/python/sglang/srt/layers/moe/hash_topk.py @@ -35,6 +35,20 @@ class HashTopK(nn.Module): apply_routed_scaling_factor_on_output=False, ): super().__init__() + self.layer_id = None + from sglang.srt.server_args import get_global_server_args + + self.enable_deepep_waterfill = ( + num_fused_shared_experts > 0 + and get_global_server_args().enable_deepep_waterfill + ) + self.deepep_waterfill_balancer = None + + if self.enable_deepep_waterfill: + # Waterfill appends the shared expert after EPLB maps routed IDs. + topk -= num_fused_shared_experts + num_fused_shared_experts = 0 + self.num_experts = num_experts self.topk = topk self.routed_scaling_factor = routed_scaling_factor @@ -70,7 +84,21 @@ class HashTopK(nn.Module): topk_weights = torch.empty((0, topk), dtype=torch.float32, device=device) topk_ids = torch.full((0, topk), -1, dtype=torch.int32, device=device) router_logits = torch.empty((0, topk), dtype=torch.float32, device=device) - return StandardTopKOutput(topk_weights, topk_ids, router_logits) + return self._apply_deepep_waterfill( + StandardTopKOutput(topk_weights, topk_ids, router_logits), + num_tokens=0, + ) + + def _apply_deepep_waterfill( + self, topk_output: StandardTopKOutput, num_tokens: int + ) -> StandardTopKOutput: + if self.enable_deepep_waterfill and self.deepep_waterfill_balancer is None: + raise RuntimeError( + "DeepEP waterfill HashTopK must be prepared by ModelRunner before forward." + ) + if self.deepep_waterfill_balancer is None: + return topk_output + return self.deepep_waterfill_balancer.expand_topk(topk_output, num_tokens) def _forward_torch( self, router_logits: torch.Tensor, input_ids: torch.Tensor @@ -152,4 +180,4 @@ class HashTopK(nn.Module): topk_output = StandardTopKOutput( topk_weights=topk_weights, topk_ids=topk_ids, router_logits=router_logits ) - return topk_output + return self._apply_deepep_waterfill(topk_output, hidden_states.shape[0]) diff --git a/python/sglang/srt/layers/moe/topk.py b/python/sglang/srt/layers/moe/topk.py index ca716dc33..3c073f370 100644 --- a/python/sglang/srt/layers/moe/topk.py +++ b/python/sglang/srt/layers/moe/topk.py @@ -339,17 +339,12 @@ class TopK(MultiPlatformOp): assert num_expert_group is not None and topk_group is not None self.layer_id = layer_id - if num_fused_shared_experts > 0: - from sglang.srt.server_args import get_global_server_args + from sglang.srt.server_args import get_global_server_args - try: - self.enable_deepep_waterfill = ( - get_global_server_args().enable_deepep_waterfill - ) - except ValueError: - self.enable_deepep_waterfill = False - else: - self.enable_deepep_waterfill = False + self.enable_deepep_waterfill = ( + num_fused_shared_experts > 0 + and get_global_server_args().enable_deepep_waterfill + ) self.deepep_waterfill_balancer = None if self.enable_deepep_waterfill: diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index d8245c06f..fad8a8868 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -121,6 +121,7 @@ from sglang.srt.layers.dp_attention import ( set_is_extend_in_batch, ) from sglang.srt.layers.logits_processor import LogitsProcessorOutput +from sglang.srt.layers.moe.hash_topk import HashTopK from sglang.srt.layers.moe.topk import TopK from sglang.srt.layers.pooler import EmbeddingPoolerOutput from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype @@ -1432,7 +1433,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): num_prepared = 0 num_routed_experts = None for module in self.model.modules(): - if not isinstance(module, TopK): + if not isinstance(module, (TopK, HashTopK)): continue if ( not module.enable_deepep_waterfill @@ -1459,15 +1460,17 @@ class ModelRunner(ModelRunnerKVCacheMixin): num_physical_routed_experts = ( num_routed_experts + self.server_args.ep_num_redundant_experts ) + if isinstance(module, TopK): + routed_scaling_factor = module.topk_config.routed_scaling_factor + else: + routed_scaling_factor = module.routed_scaling_factor module.deepep_waterfill_balancer = balancer_cls( num_routed_experts=num_physical_routed_experts, world_size=self.moe_ep_size, rank=self.moe_ep_rank, layer_id=module.layer_id, routed_scaling_factor=( - module.topk_config.routed_scaling_factor - if module.topk_config.routed_scaling_factor is not None - else 1.0 + routed_scaling_factor if routed_scaling_factor is not None else 1.0 ), ) num_prepared += 1 diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 6e8af891c..da6845ca0 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -1396,6 +1396,22 @@ class DeepseekV4ForCausalLM(nn.Module): if get_global_server_args().disable_shared_experts_fusion: return + # Waterfill needs shared-experts fusion so it can dispatch shared + # expert tokens to least-loaded EP ranks. + if get_global_server_args().enable_deepep_waterfill: + if self.config.n_shared_experts != 1: + raise ValueError( + "DeepEP Waterfill for DeepSeek V4 expects exactly one shared " + f"expert, but got n_shared_experts={self.config.n_shared_experts}." + ) + self.num_fused_shared_experts = self.config.n_shared_experts + log_info_on_rank0( + logger, + "DeepSeek V4: --enable-deepep-waterfill set; KEEP shared-experts " + "fusion enabled so waterfill can rebalance shared expert dispatch.", + ) + return + get_global_server_args().disable_shared_experts_fusion = True log_info_on_rank0( logger,