Support DeepSeek V4 DeepEP Waterfill (#25391)
This commit is contained in:
@@ -35,6 +35,20 @@ class HashTopK(nn.Module):
|
|||||||
apply_routed_scaling_factor_on_output=False,
|
apply_routed_scaling_factor_on_output=False,
|
||||||
):
|
):
|
||||||
super().__init__()
|
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.num_experts = num_experts
|
||||||
self.topk = topk
|
self.topk = topk
|
||||||
self.routed_scaling_factor = routed_scaling_factor
|
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_weights = torch.empty((0, topk), dtype=torch.float32, device=device)
|
||||||
topk_ids = torch.full((0, topk), -1, dtype=torch.int32, 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)
|
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(
|
def _forward_torch(
|
||||||
self, router_logits: torch.Tensor, input_ids: torch.Tensor
|
self, router_logits: torch.Tensor, input_ids: torch.Tensor
|
||||||
@@ -152,4 +180,4 @@ class HashTopK(nn.Module):
|
|||||||
topk_output = StandardTopKOutput(
|
topk_output = StandardTopKOutput(
|
||||||
topk_weights=topk_weights, topk_ids=topk_ids, router_logits=router_logits
|
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])
|
||||||
|
|||||||
@@ -339,17 +339,12 @@ class TopK(MultiPlatformOp):
|
|||||||
assert num_expert_group is not None and topk_group is not None
|
assert num_expert_group is not None and topk_group is not None
|
||||||
|
|
||||||
self.layer_id = layer_id
|
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 = (
|
self.enable_deepep_waterfill = (
|
||||||
get_global_server_args().enable_deepep_waterfill
|
num_fused_shared_experts > 0
|
||||||
|
and get_global_server_args().enable_deepep_waterfill
|
||||||
)
|
)
|
||||||
except ValueError:
|
|
||||||
self.enable_deepep_waterfill = False
|
|
||||||
else:
|
|
||||||
self.enable_deepep_waterfill = False
|
|
||||||
|
|
||||||
self.deepep_waterfill_balancer = None
|
self.deepep_waterfill_balancer = None
|
||||||
if self.enable_deepep_waterfill:
|
if self.enable_deepep_waterfill:
|
||||||
|
|||||||
@@ -121,6 +121,7 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
set_is_extend_in_batch,
|
set_is_extend_in_batch,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
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.moe.topk import TopK
|
||||||
from sglang.srt.layers.pooler import EmbeddingPoolerOutput
|
from sglang.srt.layers.pooler import EmbeddingPoolerOutput
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
|
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
|
||||||
@@ -1432,7 +1433,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
num_prepared = 0
|
num_prepared = 0
|
||||||
num_routed_experts = None
|
num_routed_experts = None
|
||||||
for module in self.model.modules():
|
for module in self.model.modules():
|
||||||
if not isinstance(module, TopK):
|
if not isinstance(module, (TopK, HashTopK)):
|
||||||
continue
|
continue
|
||||||
if (
|
if (
|
||||||
not module.enable_deepep_waterfill
|
not module.enable_deepep_waterfill
|
||||||
@@ -1459,15 +1460,17 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
num_physical_routed_experts = (
|
num_physical_routed_experts = (
|
||||||
num_routed_experts + self.server_args.ep_num_redundant_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(
|
module.deepep_waterfill_balancer = balancer_cls(
|
||||||
num_routed_experts=num_physical_routed_experts,
|
num_routed_experts=num_physical_routed_experts,
|
||||||
world_size=self.moe_ep_size,
|
world_size=self.moe_ep_size,
|
||||||
rank=self.moe_ep_rank,
|
rank=self.moe_ep_rank,
|
||||||
layer_id=module.layer_id,
|
layer_id=module.layer_id,
|
||||||
routed_scaling_factor=(
|
routed_scaling_factor=(
|
||||||
module.topk_config.routed_scaling_factor
|
routed_scaling_factor if routed_scaling_factor is not None else 1.0
|
||||||
if module.topk_config.routed_scaling_factor is not None
|
|
||||||
else 1.0
|
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
num_prepared += 1
|
num_prepared += 1
|
||||||
|
|||||||
@@ -1396,6 +1396,22 @@ class DeepseekV4ForCausalLM(nn.Module):
|
|||||||
if get_global_server_args().disable_shared_experts_fusion:
|
if get_global_server_args().disable_shared_experts_fusion:
|
||||||
return
|
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
|
get_global_server_args().disable_shared_experts_fusion = True
|
||||||
log_info_on_rank0(
|
log_info_on_rank0(
|
||||||
logger,
|
logger,
|
||||||
|
|||||||
Reference in New Issue
Block a user