feat: [1/2] [DeepEP] Fuse shared expert into MoE dispatch under EP (#20089)
Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com> Co-authored-by: AichenF <aichenf@nvidia.com>
This commit is contained in:
co-authored by
Claude Sonnet 4.6
AichenF
parent
7965573eb4
commit
57ffc55fb6
@@ -48,7 +48,7 @@ from sglang.srt.layers.moe.topk import (
|
|||||||
TopKOutput,
|
TopKOutput,
|
||||||
TopKOutputChecker,
|
TopKOutputChecker,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.moe.utils import RoutingMethodType
|
from sglang.srt.layers.moe.utils import RoutingMethodType, is_deepep_class_backend
|
||||||
from sglang.srt.layers.quantization.base_config import (
|
from sglang.srt.layers.quantization.base_config import (
|
||||||
FusedMoEMethodBase,
|
FusedMoEMethodBase,
|
||||||
QuantizationConfig,
|
QuantizationConfig,
|
||||||
@@ -197,10 +197,20 @@ class FusedMoE(torch.nn.Module):
|
|||||||
self.moe_ep_rank = get_moe_expert_parallel_rank()
|
self.moe_ep_rank = get_moe_expert_parallel_rank()
|
||||||
self.moe_tp_size = get_moe_tensor_parallel_world_size()
|
self.moe_tp_size = get_moe_tensor_parallel_world_size()
|
||||||
self.moe_tp_rank = get_moe_tensor_parallel_rank()
|
self.moe_tp_rank = get_moe_tensor_parallel_rank()
|
||||||
assert (num_experts - num_fused_shared_experts) % self.moe_ep_size == 0
|
|
||||||
self.num_local_experts = (
|
# DeepEP: each rank has its own shared expert slot, so total shared
|
||||||
num_experts - num_fused_shared_experts
|
# weight slots = num_fused_shared_experts * ep_size.
|
||||||
) // self.moe_ep_size + num_fused_shared_experts
|
# AMD/Standard: shared experts are global, slots = num_fused_shared_experts.
|
||||||
|
if num_fused_shared_experts > 0 and is_deepep_class_backend():
|
||||||
|
num_shared_slots = num_fused_shared_experts * self.moe_ep_size
|
||||||
|
else:
|
||||||
|
num_shared_slots = num_fused_shared_experts
|
||||||
|
|
||||||
|
assert (num_experts - num_shared_slots) % self.moe_ep_size == 0
|
||||||
|
self._num_global_routed = num_experts - num_shared_slots
|
||||||
|
self._num_local_routed = self._num_global_routed // self.moe_ep_size
|
||||||
|
self.num_local_experts = self._num_local_routed + num_fused_shared_experts
|
||||||
|
self._has_fused_shared = num_fused_shared_experts > 0
|
||||||
|
|
||||||
self.expert_mask_gpu = None
|
self.expert_mask_gpu = None
|
||||||
|
|
||||||
@@ -555,18 +565,12 @@ class FusedMoE(torch.nn.Module):
|
|||||||
expert_data.copy_(loaded_weight)
|
expert_data.copy_(loaded_weight)
|
||||||
|
|
||||||
def _map_global_expert_id_to_local_expert_id(self, expert_id: int) -> int:
|
def _map_global_expert_id_to_local_expert_id(self, expert_id: int) -> int:
|
||||||
num_global_routed_experts = self.num_experts - self.num_fused_shared_experts
|
start_idx = self.moe_ep_rank * self._num_local_routed
|
||||||
num_local_routed_experts = (
|
end_idx = start_idx + self._num_local_routed
|
||||||
self.num_local_experts - self.num_fused_shared_experts
|
|
||||||
)
|
|
||||||
start_idx = self.moe_ep_rank * num_local_routed_experts
|
|
||||||
end_idx = (self.moe_ep_rank + 1) * num_local_routed_experts
|
|
||||||
if start_idx <= expert_id < end_idx:
|
if start_idx <= expert_id < end_idx:
|
||||||
return expert_id - start_idx
|
return expert_id - start_idx
|
||||||
elif (
|
elif self._has_fused_shared and expert_id >= self._num_global_routed:
|
||||||
self.num_fused_shared_experts > 0 and expert_id >= num_global_routed_experts
|
return expert_id - self._num_global_routed + self._num_local_routed
|
||||||
):
|
|
||||||
return expert_id - num_global_routed_experts + num_local_routed_experts
|
|
||||||
else:
|
else:
|
||||||
return -1
|
return -1
|
||||||
|
|
||||||
@@ -611,7 +615,7 @@ class FusedMoE(torch.nn.Module):
|
|||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
if expert_id >= self.num_experts - self.num_fused_shared_experts:
|
if self._has_fused_shared and expert_id >= self._num_global_routed:
|
||||||
# This is a shared expert.
|
# This is a shared expert.
|
||||||
physical_expert_ids = [expert_id]
|
physical_expert_ids = [expert_id]
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -36,7 +36,11 @@ try:
|
|||||||
except ImportError:
|
except ImportError:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
from sglang.srt.distributed import get_tp_group
|
from sglang.srt.distributed import (
|
||||||
|
get_moe_expert_parallel_rank,
|
||||||
|
get_moe_expert_parallel_world_size,
|
||||||
|
get_tp_group,
|
||||||
|
)
|
||||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||||
use_symmetric_memory,
|
use_symmetric_memory,
|
||||||
)
|
)
|
||||||
@@ -49,6 +53,7 @@ from sglang.srt.eplb.expert_location_dispatch import (
|
|||||||
from sglang.srt.layers.dp_attention import is_allocation_symmetric
|
from sglang.srt.layers.dp_attention import is_allocation_symmetric
|
||||||
from sglang.srt.layers.moe import get_moe_runner_backend
|
from sglang.srt.layers.moe import get_moe_runner_backend
|
||||||
from sglang.srt.layers.moe.routed_experts_capturer import get_global_experts_capturer
|
from sglang.srt.layers.moe.routed_experts_capturer import get_global_experts_capturer
|
||||||
|
from sglang.srt.layers.moe.utils import is_deepep_class_backend
|
||||||
from sglang.srt.layers.utils import MultiPlatformOp
|
from sglang.srt.layers.utils import MultiPlatformOp
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
@@ -789,15 +794,16 @@ def biased_grouped_topk_gpu(
|
|||||||
num_experts // num_expert_group if num_expert_group else num_experts
|
num_experts // num_expert_group if num_expert_group else num_experts
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# topk for routed experts only (shared experts are appended separately below)
|
||||||
|
topk_routed = topk - num_fused_shared_experts
|
||||||
if (
|
if (
|
||||||
_is_cuda
|
_is_cuda
|
||||||
and fused_topk_deepseek is not None
|
and fused_topk_deepseek is not None
|
||||||
and num_fused_shared_experts == 0
|
|
||||||
and is_power_of_two(num_experts)
|
and is_power_of_two(num_experts)
|
||||||
# flashinfer constraints
|
# flashinfer constraints (applied to routed experts only)
|
||||||
and topk <= 8
|
and topk_routed <= 8
|
||||||
and topk_group <= num_expert_group
|
and topk_group <= num_expert_group
|
||||||
and topk_group * num_expert_group >= topk
|
and topk_group * num_expert_group >= topk_routed
|
||||||
and (
|
and (
|
||||||
(experts_per_group <= 32 and experts_per_group * topk_group <= 128)
|
(experts_per_group <= 32 and experts_per_group * topk_group <= 128)
|
||||||
if num_expert_group > 1
|
if num_expert_group > 1
|
||||||
@@ -806,10 +812,10 @@ def biased_grouped_topk_gpu(
|
|||||||
):
|
):
|
||||||
# Pre-allocate output tensors (flashinfer mutates them in-place)
|
# Pre-allocate output tensors (flashinfer mutates them in-place)
|
||||||
topk_weights = torch.empty(
|
topk_weights = torch.empty(
|
||||||
(num_tokens, topk), dtype=torch.float32, device=gating_output.device
|
(num_tokens, topk_routed), dtype=torch.float32, device=gating_output.device
|
||||||
)
|
)
|
||||||
topk_ids = torch.empty(
|
topk_ids = torch.empty(
|
||||||
(num_tokens, topk), dtype=torch.int32, device=gating_output.device
|
(num_tokens, topk_routed), dtype=torch.int32, device=gating_output.device
|
||||||
)
|
)
|
||||||
|
|
||||||
# flashinfer always applies the scaling_factor internally
|
# flashinfer always applies the scaling_factor internally
|
||||||
@@ -823,13 +829,25 @@ def biased_grouped_topk_gpu(
|
|||||||
correction_bias,
|
correction_bias,
|
||||||
num_expert_group,
|
num_expert_group,
|
||||||
topk_group,
|
topk_group,
|
||||||
topk,
|
topk_routed,
|
||||||
scaling_factor,
|
scaling_factor,
|
||||||
topk_weights,
|
topk_weights,
|
||||||
topk_ids,
|
topk_ids,
|
||||||
True,
|
True,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if num_fused_shared_experts > 0:
|
||||||
|
# Append shared expert columns: ID = num_experts (first shared slot),
|
||||||
|
# weight = sum(routed) / scaling_factor (matching biased_grouped_topk_impl).
|
||||||
|
# DeepEP fusion will overwrite both in _remap_topk_ids_for_deepep_fusion.
|
||||||
|
topk_ids = F.pad(topk_ids, (0, num_fused_shared_experts), value=num_experts)
|
||||||
|
topk_weights = F.pad(topk_weights, (0, num_fused_shared_experts))
|
||||||
|
if routed_scaling_factor is not None:
|
||||||
|
topk_weights[:, topk_routed:] = (
|
||||||
|
topk_weights[:, :topk_routed].sum(dim=-1, keepdim=True)
|
||||||
|
/ routed_scaling_factor
|
||||||
|
)
|
||||||
|
|
||||||
return topk_weights, topk_ids
|
return topk_weights, topk_ids
|
||||||
|
|
||||||
elif (
|
elif (
|
||||||
@@ -938,6 +956,50 @@ else:
|
|||||||
fused_topk_native = fused_topk_torch_native
|
fused_topk_native = fused_topk_torch_native
|
||||||
|
|
||||||
|
|
||||||
|
def _remap_topk_for_deepep(
|
||||||
|
topk_ids: torch.Tensor,
|
||||||
|
topk_weights: torch.Tensor,
|
||||||
|
num_fused_shared_experts: int,
|
||||||
|
n_routed_experts: int,
|
||||||
|
topk_config: TopKConfig,
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""Remap TopK output to DeepEP interleaved expert layout.
|
||||||
|
|
||||||
|
DeepEP dispatch needs each rank's shared expert at a unique ID so tokens
|
||||||
|
route to the correct rank. The layout interleaves shared slots among
|
||||||
|
routed experts: [routed_0..L-1, shared, routed_L..2L-1, shared, ...].
|
||||||
|
|
||||||
|
Routed IDs: e -> e + e // num_local_routed
|
||||||
|
Shared IDs: ep_rank * num_local_experts + num_local_routed
|
||||||
|
Shared weight: 1 / routed_scaling_factor (compensates post-MoE scaling)
|
||||||
|
"""
|
||||||
|
if topk_ids.shape[0] == 0:
|
||||||
|
return topk_ids, topk_weights
|
||||||
|
|
||||||
|
ep_size = get_moe_expert_parallel_world_size()
|
||||||
|
ep_rank = get_moe_expert_parallel_rank()
|
||||||
|
num_local_routed = n_routed_experts // ep_size
|
||||||
|
num_local_experts = num_local_routed + num_fused_shared_experts
|
||||||
|
|
||||||
|
# Remap routed IDs: insert gaps for shared expert slots (single fused op)
|
||||||
|
routed = topk_ids[:, :-num_fused_shared_experts]
|
||||||
|
topk_ids[:, :-num_fused_shared_experts] = routed + routed // num_local_routed
|
||||||
|
|
||||||
|
# Set shared expert IDs to route to home rank (vectorized)
|
||||||
|
topk_ids[:, -num_fused_shared_experts:] = (
|
||||||
|
ep_rank * num_local_experts
|
||||||
|
+ num_local_routed
|
||||||
|
+ torch.arange(num_fused_shared_experts, device=topk_ids.device)
|
||||||
|
)
|
||||||
|
|
||||||
|
# Override shared weight: 1/routed_scaling_factor so net contribution = 1.0
|
||||||
|
routed_scaling_factor = topk_config.routed_scaling_factor
|
||||||
|
if routed_scaling_factor is not None and routed_scaling_factor != 0:
|
||||||
|
topk_weights[:, -num_fused_shared_experts:] = 1.0 / routed_scaling_factor
|
||||||
|
|
||||||
|
return topk_ids, topk_weights
|
||||||
|
|
||||||
|
|
||||||
def _post_process_topk_ids(
|
def _post_process_topk_ids(
|
||||||
topk_ids: torch.Tensor,
|
topk_ids: torch.Tensor,
|
||||||
topk_weights: torch.Tensor,
|
topk_weights: torch.Tensor,
|
||||||
@@ -980,6 +1042,18 @@ def _post_process_topk_ids(
|
|||||||
scale_factor,
|
scale_factor,
|
||||||
N, # base id for shared experts
|
N, # base id for shared experts
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# DeepEP: remap to interleaved expert layout where each rank's shared
|
||||||
|
# expert has a unique ID for dispatch routing.
|
||||||
|
if num_fused_shared_experts > 0 and is_deepep_class_backend():
|
||||||
|
topk_ids, topk_weights = _remap_topk_for_deepep(
|
||||||
|
topk_ids,
|
||||||
|
topk_weights,
|
||||||
|
num_fused_shared_experts,
|
||||||
|
router_logits.shape[1],
|
||||||
|
topk_config,
|
||||||
|
)
|
||||||
|
|
||||||
return topk_ids, topk_weights
|
return topk_ids, topk_weights
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -254,6 +254,12 @@ def is_sbo_enabled() -> bool:
|
|||||||
return IS_SBO_ENABLED
|
return IS_SBO_ENABLED
|
||||||
|
|
||||||
|
|
||||||
|
def is_deepep_class_backend() -> bool:
|
||||||
|
"""Check if the MoE backend is DeepEP-family (DeepEP, Mooncake, or Mori)."""
|
||||||
|
b = get_moe_a2a_backend()
|
||||||
|
return b.is_deepep() or b.is_mooncake() or b.is_mori()
|
||||||
|
|
||||||
|
|
||||||
def get_tbo_token_distribution_threshold() -> float:
|
def get_tbo_token_distribution_threshold() -> float:
|
||||||
global TBO_TOKEN_DISTRIBUTION_THRESHOLD
|
global TBO_TOKEN_DISTRIBUTION_THRESHOLD
|
||||||
if TBO_TOKEN_DISTRIBUTION_THRESHOLD is None:
|
if TBO_TOKEN_DISTRIBUTION_THRESHOLD is None:
|
||||||
|
|||||||
@@ -103,6 +103,9 @@ from sglang.srt.layers.moe.topk import TopK, TopKOutputFormat
|
|||||||
from sglang.srt.layers.moe.utils import (
|
from sglang.srt.layers.moe.utils import (
|
||||||
RoutingMethodType,
|
RoutingMethodType,
|
||||||
filter_moe_weight_param_global_expert,
|
filter_moe_weight_param_global_expert,
|
||||||
|
is_deepep_class_backend,
|
||||||
|
is_sbo_enabled,
|
||||||
|
is_tbo_enabled,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||||
from sglang.srt.layers.quantization.fp8 import Fp8Config
|
from sglang.srt.layers.quantization.fp8 import Fp8Config
|
||||||
@@ -364,11 +367,34 @@ class DeepseekV2MoE(nn.Module):
|
|||||||
self.moe_ep_size = get_moe_expert_parallel_world_size()
|
self.moe_ep_size = get_moe_expert_parallel_world_size()
|
||||||
self.routed_scaling_factor = config.routed_scaling_factor
|
self.routed_scaling_factor = config.routed_scaling_factor
|
||||||
self.n_shared_experts = config.n_shared_experts
|
self.n_shared_experts = config.n_shared_experts
|
||||||
self.num_fused_shared_experts = (
|
|
||||||
0
|
n_shared_experts = (
|
||||||
if get_global_server_args().disable_shared_experts_fusion
|
0 if config.n_shared_experts is None else int(config.n_shared_experts)
|
||||||
else config.n_shared_experts
|
|
||||||
)
|
)
|
||||||
|
_fusion_disabled = get_global_server_args().disable_shared_experts_fusion
|
||||||
|
|
||||||
|
# num_fused_shared_experts drives weight remapping in deepseek_weight_loader:
|
||||||
|
# mlp.shared_experts → mlp.experts.256 when > 0.
|
||||||
|
self.num_fused_shared_experts = 0 if _fusion_disabled else n_shared_experts
|
||||||
|
|
||||||
|
# DeepEP shared expert fusion: shared expert is fused into the same MoE kernel
|
||||||
|
# as a local expert at the home EP rank. Expert layout is expanded from 256
|
||||||
|
# routed to 256+EP_size (e.g. 272 for EP=16). TopK handles interleaving.
|
||||||
|
_is_deepep_fusion = (
|
||||||
|
is_deepep_class_backend() and self.num_fused_shared_experts > 0
|
||||||
|
)
|
||||||
|
|
||||||
|
if _is_deepep_fusion:
|
||||||
|
# 256 routed + EP_size shared slots = 272 experts total (for EP=16)
|
||||||
|
num_experts_for_moe = config.n_routed_experts + self.moe_ep_size
|
||||||
|
top_k_for_moe = config.num_experts_per_tok + 1 # 8 routed + 1 shared
|
||||||
|
# Interleaving for DeepEP dispatch is handled by TopK internally.
|
||||||
|
else:
|
||||||
|
num_experts_for_moe = (
|
||||||
|
config.n_routed_experts + self.num_fused_shared_experts
|
||||||
|
)
|
||||||
|
top_k_for_moe = config.num_experts_per_tok + self.num_fused_shared_experts
|
||||||
|
|
||||||
self.config = config
|
self.config = config
|
||||||
self.layer_id = layer_id
|
self.layer_id = layer_id
|
||||||
self.alt_stream = alt_stream
|
self.alt_stream = alt_stream
|
||||||
@@ -394,19 +420,24 @@ class DeepseekV2MoE(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# scaling factor for fused shared experts on AMD-platform.
|
# scaling factor for fused shared experts on AMD-platform.
|
||||||
|
# DeepEP doesn't need this: shared expert is only computed on home rank
|
||||||
|
# (not all-reduced), so no 1/ep_size correction is needed.
|
||||||
fused_shared_experts_scaling_factor = None
|
fused_shared_experts_scaling_factor = None
|
||||||
if self.moe_ep_size > 1 and self.num_fused_shared_experts > 0:
|
if (
|
||||||
|
self.moe_ep_size > 1
|
||||||
|
and self.num_fused_shared_experts > 0
|
||||||
|
and not _is_deepep_fusion
|
||||||
|
):
|
||||||
# if enable_ep_moe tp_szie == ep_size, every gpu get shared experts gemm output
|
# if enable_ep_moe tp_szie == ep_size, every gpu get shared experts gemm output
|
||||||
# so we scale with 1 / self.moe_ep_size in ep mode which will make it equalation as in tp mode
|
# so we scale with 1 / self.moe_ep_size in ep mode which will make it equalation as in tp mode
|
||||||
# with fused_shared_experts
|
# with fused_shared_experts
|
||||||
fused_shared_experts_scaling_factor = 1.0 / float(self.moe_ep_size)
|
fused_shared_experts_scaling_factor = 1.0 / float(self.moe_ep_size)
|
||||||
|
|
||||||
self.experts = get_moe_impl_class(quant_config)(
|
self.experts = get_moe_impl_class(quant_config)(
|
||||||
num_experts=config.n_routed_experts
|
num_experts=num_experts_for_moe
|
||||||
+ self.num_fused_shared_experts
|
|
||||||
+ get_global_server_args().ep_num_redundant_experts,
|
+ get_global_server_args().ep_num_redundant_experts,
|
||||||
num_fused_shared_experts=self.num_fused_shared_experts,
|
num_fused_shared_experts=self.num_fused_shared_experts,
|
||||||
top_k=config.num_experts_per_tok + self.num_fused_shared_experts,
|
top_k=top_k_for_moe,
|
||||||
hidden_size=config.hidden_size,
|
hidden_size=config.hidden_size,
|
||||||
intermediate_size=config.moe_intermediate_size,
|
intermediate_size=config.moe_intermediate_size,
|
||||||
layer_id=self.layer_id,
|
layer_id=self.layer_id,
|
||||||
@@ -444,7 +475,14 @@ class DeepseekV2MoE(nn.Module):
|
|||||||
self.shared_experts_is_int8 = False
|
self.shared_experts_is_int8 = False
|
||||||
self.shared_experts_is_fp8 = False
|
self.shared_experts_is_fp8 = False
|
||||||
self.shared_experts_weight_block_size = None
|
self.shared_experts_weight_block_size = None
|
||||||
if config.n_shared_experts is not None and self.num_fused_shared_experts == 0:
|
# Shared experts: skip when fused into MoE kernel (self.num_fused_shared_experts > 0)
|
||||||
|
# or when DeepEP fusion is enabled (shared expert is local slot 16 in FusedMoE, no separate MLP).
|
||||||
|
if (
|
||||||
|
config.n_shared_experts is not None
|
||||||
|
and config.n_shared_experts > 0
|
||||||
|
and self.num_fused_shared_experts == 0
|
||||||
|
and not _is_deepep_fusion
|
||||||
|
):
|
||||||
intermediate_size = config.moe_intermediate_size * config.n_shared_experts
|
intermediate_size = config.moe_intermediate_size * config.n_shared_experts
|
||||||
# disable tp for shared experts when enable deepep moe, or with fp4 allgather
|
# disable tp for shared experts when enable deepep moe, or with fp4 allgather
|
||||||
self.shared_experts = DeepseekV2MLP(
|
self.shared_experts = DeepseekV2MLP(
|
||||||
@@ -775,7 +813,7 @@ class DeepseekV2MoE(nn.Module):
|
|||||||
if hidden_states.shape[0] > 0:
|
if hidden_states.shape[0] > 0:
|
||||||
# router_logits: (num_tokens, n_experts)
|
# router_logits: (num_tokens, n_experts)
|
||||||
router_logits = self.gate(hidden_states, forward_batch=forward_batch)
|
router_logits = self.gate(hidden_states, forward_batch=forward_batch)
|
||||||
if not sbo_enabled_flag:
|
if not sbo_enabled_flag and self.num_fused_shared_experts == 0:
|
||||||
if self.alt_stream is not None:
|
if self.alt_stream is not None:
|
||||||
self.alt_stream.wait_stream(torch.cuda.current_stream())
|
self.alt_stream.wait_stream(torch.cuda.current_stream())
|
||||||
with torch.cuda.stream(self.alt_stream):
|
with torch.cuda.stream(self.alt_stream):
|
||||||
@@ -794,6 +832,16 @@ class DeepseekV2MoE(nn.Module):
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
topk_output = self.topk.empty_topk_output(hidden_states.device)
|
topk_output = self.topk.empty_topk_output(hidden_states.device)
|
||||||
|
if is_deepep_class_backend() and self.num_fused_shared_experts > 0:
|
||||||
|
n = self.num_fused_shared_experts
|
||||||
|
topk_output = topk_output._replace(
|
||||||
|
topk_ids=topk_output.topk_ids.new_empty(
|
||||||
|
(0, topk_output.topk_ids.shape[-1] + n)
|
||||||
|
),
|
||||||
|
topk_weights=topk_output.topk_weights.new_empty(
|
||||||
|
(0, topk_output.topk_weights.shape[-1] + n)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
if sbo_overlap_dispatch_flag:
|
if sbo_overlap_dispatch_flag:
|
||||||
shared_output = None
|
shared_output = None
|
||||||
@@ -951,6 +999,7 @@ class DeepseekV2MoE(nn.Module):
|
|||||||
if (
|
if (
|
||||||
hidden_states.shape[0] > 0
|
hidden_states.shape[0] > 0
|
||||||
and not sbo_enabled_flag
|
and not sbo_enabled_flag
|
||||||
|
and self.num_fused_shared_experts == 0
|
||||||
and self.alt_stream is not None
|
and self.alt_stream is not None
|
||||||
):
|
):
|
||||||
torch.cuda.current_stream().wait_event(shared_event)
|
torch.cuda.current_stream().wait_event(shared_event)
|
||||||
@@ -1923,12 +1972,7 @@ class DeepseekV2Model(nn.Module):
|
|||||||
for i in range(len(self.layers)):
|
for i in range(len(self.layers)):
|
||||||
if isinstance(self.layers[i].mlp, DeepseekV2MoE):
|
if isinstance(self.layers[i].mlp, DeepseekV2MoE):
|
||||||
# tp_size = get_tensor_model_parallel_world_size()
|
# tp_size = get_tensor_model_parallel_world_size()
|
||||||
a2a_backend = get_moe_a2a_backend()
|
is_a2a_moe = is_deepep_class_backend()
|
||||||
is_a2a_moe = (
|
|
||||||
a2a_backend.is_deepep()
|
|
||||||
or a2a_backend.is_mori()
|
|
||||||
or a2a_backend.is_mooncake()
|
|
||||||
)
|
|
||||||
tp_size = (
|
tp_size = (
|
||||||
1 if is_a2a_moe else get_tensor_model_parallel_world_size()
|
1 if is_a2a_moe else get_tensor_model_parallel_world_size()
|
||||||
)
|
)
|
||||||
@@ -2176,12 +2220,28 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
|
|||||||
self, architecture: str = "DeepseekV3ForCausalLM"
|
self, architecture: str = "DeepseekV3ForCausalLM"
|
||||||
):
|
):
|
||||||
self.num_fused_shared_experts = 0
|
self.num_fused_shared_experts = 0
|
||||||
if get_global_server_args().disable_shared_experts_fusion:
|
server_args = get_global_server_args()
|
||||||
|
|
||||||
|
if server_args.disable_shared_experts_fusion:
|
||||||
return
|
return
|
||||||
|
|
||||||
# Only Deepseek V3/R1 can use shared experts fusion optimization now.
|
# DeepEP + enforce: the only path that enables fusion under DeepEP.
|
||||||
|
if is_deepep_class_backend() and server_args.enforce_shared_experts_fusion:
|
||||||
|
log_info_on_rank0(
|
||||||
|
logger,
|
||||||
|
"DeepEP shared expert fusion: fusing shared expert into MoE kernel "
|
||||||
|
"at home EP rank local slot (--enforce-shared-experts-fusion).",
|
||||||
|
)
|
||||||
|
self.num_fused_shared_experts = self.config.n_shared_experts
|
||||||
|
return
|
||||||
|
|
||||||
|
# Check all conditions that disable fusion.
|
||||||
disable_reason = None
|
disable_reason = None
|
||||||
if (
|
if is_sbo_enabled() or is_tbo_enabled():
|
||||||
|
disable_reason = "SBO/TBO enabled: incompatible with fusing shared expert into MoE kernel."
|
||||||
|
elif is_deepep_class_backend():
|
||||||
|
disable_reason = "DeepEP: fusion off by default (use --enforce-shared-experts-fusion to enable)."
|
||||||
|
elif (
|
||||||
self.config.architectures[0] != architecture
|
self.config.architectures[0] != architecture
|
||||||
or self.config.n_routed_experts != 256
|
or self.config.n_routed_experts != 256
|
||||||
or self.config.n_shared_experts != 1
|
or self.config.n_shared_experts != 1
|
||||||
@@ -2197,16 +2257,15 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
|
|||||||
elif get_moe_expert_parallel_world_size() > 1 and (
|
elif get_moe_expert_parallel_world_size() > 1 and (
|
||||||
not _is_hip or torch.cuda.get_device_capability("cuda") < (9, 4)
|
not _is_hip or torch.cuda.get_device_capability("cuda") < (9, 4)
|
||||||
):
|
):
|
||||||
disable_reason = "Only Deepseek V3/R1 on AMD-platform with capability >= gfx942(MI30x) can use shared experts fusion optimization under expert parallelism."
|
disable_reason = (
|
||||||
elif disable_reason is None and (
|
"Only Deepseek V3/R1 on AMD-platform with capability >= gfx942(MI30x) "
|
||||||
get_moe_a2a_backend().is_deepep() or get_moe_a2a_backend().is_mori()
|
"can use shared experts fusion optimization under expert parallelism."
|
||||||
):
|
)
|
||||||
disable_reason = "Deepseek V3/R1 cannot use shared experts fusion optimization under deepep expert parallelism."
|
|
||||||
elif self.quant_config and self.quant_config.get_name() == "w4afp8":
|
elif self.quant_config and self.quant_config.get_name() == "w4afp8":
|
||||||
disable_reason = "Deepseek V3/R1 W4AFP8 model uses different quant method for routed experts and shared experts."
|
disable_reason = "Deepseek V3/R1 W4AFP8 model uses different quant method for routed experts and shared experts."
|
||||||
|
|
||||||
if disable_reason is not None:
|
if disable_reason is not None:
|
||||||
get_global_server_args().disable_shared_experts_fusion = True
|
server_args.disable_shared_experts_fusion = True
|
||||||
self.num_fused_shared_experts = 0
|
self.num_fused_shared_experts = 0
|
||||||
log_info_on_rank0(
|
log_info_on_rank0(
|
||||||
logger,
|
logger,
|
||||||
|
|||||||
@@ -663,6 +663,7 @@ class ServerArgs:
|
|||||||
enable_custom_logit_processor: bool = False
|
enable_custom_logit_processor: bool = False
|
||||||
flashinfer_mla_disable_ragged: bool = False
|
flashinfer_mla_disable_ragged: bool = False
|
||||||
disable_shared_experts_fusion: bool = False
|
disable_shared_experts_fusion: bool = False
|
||||||
|
enforce_shared_experts_fusion: bool = False
|
||||||
disable_chunked_prefix_cache: bool = False
|
disable_chunked_prefix_cache: bool = False
|
||||||
disable_fast_image_processor: bool = False
|
disable_fast_image_processor: bool = False
|
||||||
keep_mm_feature_on_device: bool = False
|
keep_mm_feature_on_device: bool = False
|
||||||
@@ -5833,6 +5834,12 @@ class ServerArgs:
|
|||||||
action="store_true",
|
action="store_true",
|
||||||
help="Disable shared experts fusion optimization for deepseek v3/r1.",
|
help="Disable shared experts fusion optimization for deepseek v3/r1.",
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--enforce-shared-experts-fusion",
|
||||||
|
action="store_true",
|
||||||
|
help="Enforce shared experts fusion even when it would normally be disabled (e.g. under DeepEP). "
|
||||||
|
"Mutually exclusive with --disable-shared-experts-fusion.",
|
||||||
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--disable-chunked-prefix-cache",
|
"--disable-chunked-prefix-cache",
|
||||||
action="store_true",
|
action="store_true",
|
||||||
|
|||||||
Reference in New Issue
Block a user