Support DeepSeek V4 DeepEP Waterfill (#25391)

This commit is contained in:
xutizhou
2026-05-25 21:04:26 -07:00
committed by GitHub
parent 3142278c5f
commit 59cad671e2
4 changed files with 58 additions and 16 deletions
+30 -2
View File
@@ -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])
+5 -10
View File
@@ -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:
@@ -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
+16
View File
@@ -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,