diff --git a/python/sglang/srt/layers/moe/ep_moe/layer.py b/python/sglang/srt/layers/moe/ep_moe/layer.py index f02577643..5cd2e5766 100644 --- a/python/sglang/srt/layers/moe/ep_moe/layer.py +++ b/python/sglang/srt/layers/moe/ep_moe/layer.py @@ -5,7 +5,6 @@ from typing import TYPE_CHECKING, Any, Dict, Optional, Union import torch -from sglang.srt import single_batch_overlap from sglang.srt.layers import deep_gemm_wrapper from sglang.srt.layers.moe import ( get_deepep_mode, @@ -22,7 +21,6 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.fp8 import Fp8Config from sglang.srt.layers.quantization.fp8_kernel import is_fp8_fnuz from sglang.srt.layers.quantization.w4afp8 import W4AFp8Config, W4AFp8MoEMethod -from sglang.srt.single_batch_overlap import DownGemmOverlapArgs from sglang.srt.utils import get_bool_env_var, is_hip, is_npu if TYPE_CHECKING: @@ -145,29 +143,24 @@ class DeepEPMoE(FusedMoE): self, hidden_states: torch.Tensor, topk_output: TopKOutput, - forward_shared_experts=None, - alt_stream=None, - disable_sbo=False, ): if self.deprecate_flag: - assert forward_shared_experts is None - assert alt_stream is None return super().forward( hidden_states, topk_output, ) - # We have to call SBO inside MoE to be compatible with hooks used in offloading - return single_batch_overlap.execute_sbo( - hidden_states=hidden_states, - topk_output=topk_output, - # SBO args - experts=self, - forward_shared_experts=forward_shared_experts, - alt_stream=alt_stream, - disable_sbo=disable_sbo, + # TODO: can we call super().forward here? + dispatch_output = self.dispatcher.dispatch( + hidden_states=hidden_states, topk_output=topk_output ) + combine_input = self.run_moe_core(dispatch_output) + hidden_states = self.dispatcher.combine( + combine_input=combine_input, + ) + + return hidden_states def dispatch( self, @@ -182,11 +175,9 @@ class DeepEPMoE(FusedMoE): def run_moe_core( self, dispatch_output: DispatchOutput, - down_gemm_overlap_args: Optional[DownGemmOverlapArgs] = None, ): if self.deprecate_flag: - assert down_gemm_overlap_args is None return super().run_moe_core( dispatch_output, ) @@ -210,9 +201,7 @@ class DeepEPMoE(FusedMoE): get_moe_runner_backend().is_flashinfer_cutedsl() and self.quant_config.get_name() == "modelopt_fp4" ): - output = self.forward_flashinfer_cutedsl( - dispatch_output, down_gemm_overlap_args=down_gemm_overlap_args - ) + output = self.forward_flashinfer_cutedsl(dispatch_output) elif self.use_w4afp8: output = self.forward_cutlass_w4afp8_masked(dispatch_output) else: @@ -280,7 +269,6 @@ class DeepEPMoE(FusedMoE): def forward_flashinfer_cutedsl( self, dispatch_output: DeepEPLLDispatchOutput, - down_gemm_overlap_args: Optional[DownGemmOverlapArgs], ): hidden_states, hidden_states_scale, _, _, masked_m, _ = dispatch_output assert self.quant_method is not None @@ -291,7 +279,6 @@ class DeepEPMoE(FusedMoE): x=(hidden_states, hidden_states_scale), masked_m=masked_m, moe_runner_config=self.moe_runner_config, - down_gemm_overlap_args=down_gemm_overlap_args, ) return output diff --git a/python/sglang/srt/layers/moe/fused_moe_native.py b/python/sglang/srt/layers/moe/fused_moe_native.py index a3d3a09bf..4a9070fe3 100644 --- a/python/sglang/srt/layers/moe/fused_moe_native.py +++ b/python/sglang/srt/layers/moe/fused_moe_native.py @@ -8,16 +8,19 @@ from torch.nn import functional as F from sglang.srt.layers.activation import GeluAndMul, SiluAndMul from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig -from sglang.srt.layers.moe.token_dispatcher import StandardDispatchOutput +from sglang.srt.layers.moe.token_dispatcher import ( + StandardCombineInput, + StandardDispatchOutput, +) from sglang.srt.layers.moe.topk import StandardTopKOutput def fused_moe_forward_native( layer: torch.nn.Module, dispatch_output: StandardDispatchOutput, -) -> torch.Tensor: +) -> StandardCombineInput: - x, topk_output = dispatch_output + x, x_scale, topk_output = dispatch_output moe_runner_config = layer.moe_runner_config if moe_runner_config.apply_router_weight_on_input: @@ -37,7 +40,10 @@ def fused_moe_forward_native( raise ValueError(f"Unsupported activation: {moe_runner_config.activation=}") x3 = torch.einsum("ti, taoi -> tao", x, w3_weights) expert_outs = torch.einsum("tao, taio -> tai", (x1 * x3), w2_weights) - return torch.einsum("tai,ta -> ti", expert_outs, topk_weights.to(expert_outs.dtype)) + expert_outs = torch.einsum( + "tai,ta -> ti", expert_outs, topk_weights.to(expert_outs.dtype) + ) + return StandardCombineInput(hidden_states=expert_outs) def moe_forward_native( diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py index 4924fca15..aae387734 100644 --- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py +++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py @@ -46,6 +46,7 @@ from sglang.srt.layers.quantization.modelopt_quant import ModelOptNvFp4FusedMoEM from sglang.srt.layers.quantization.unquant import UnquantizedFusedMoEMethod from sglang.srt.model_loader.weight_utils import narrow_padded_param_and_loaded_weight from sglang.srt.server_args import get_global_server_args +from sglang.srt.single_batch_overlap import DownGemmOverlapArgs from sglang.srt.two_batch_overlap import MaybeTboDeepEPDispatcher from sglang.srt.utils import ( cpu_has_amx_support, @@ -253,6 +254,10 @@ class FusedMoE(torch.nn.Module): self.routing_method_type = routing_method_type + # overlap args + self.down_gemm_overlap_args: Optional[DownGemmOverlapArgs] = None + self.meta_overlap_args: Optional[dict] = None + def _load_per_tensor_weight_scale( self, shard_id: str, @@ -856,7 +861,7 @@ class FusedMoE(torch.nn.Module): f"Unsupported weight_name {weight_name} for FusedMoE weight_loader_fused. Nothing is loaded." ) - def forward(self, hidden_states: torch.Tensor, topk_output: TopKOutput, **kwargs): + def forward(self, hidden_states: torch.Tensor, topk_output: TopKOutput): origin_hidden_states_dim = hidden_states.shape[-1] assert self.quant_method is not None @@ -875,7 +880,6 @@ class FusedMoE(torch.nn.Module): combine_input = self.run_moe_core( dispatch_output=dispatch_output, - **kwargs, ) with use_symmetric_memory( @@ -893,12 +897,11 @@ class FusedMoE(torch.nn.Module): return final_hidden_states - def run_moe_core(self, dispatch_output: DispatchOutput, **kwargs) -> CombineInput: + def run_moe_core(self, dispatch_output: DispatchOutput) -> CombineInput: # TODO: consider using symmetric memory return self.quant_method.apply( layer=self, dispatch_output=dispatch_output, - **kwargs, ) @classmethod @@ -993,6 +996,16 @@ class FusedMoE(torch.nn.Module): for shard_id in ["w1", "w2", "w3"] ] + def set_overlap_args( + self, down_gemm_overlap_args: DownGemmOverlapArgs, meta_overlap_args: dict + ): + self.down_gemm_overlap_args = down_gemm_overlap_args + self.meta_overlap_args = meta_overlap_args + + def clear_overlap_args(self) -> None: + self.down_gemm_overlap_args = None + self.meta_overlap_args = None + class FlashInferFusedMoE(FusedMoE): def __init__(self, *args, **kwargs): diff --git a/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py b/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py index d18845f99..b74560910 100644 --- a/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py +++ b/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py @@ -342,7 +342,10 @@ def pre_permute_standard_to_deep_gemm( from sglang.srt.layers.moe.ep_moe.kernels import moe_ep_deepgemm_preprocess - hidden_states, topk_output = dispatch_output + hidden_states, topk_output = ( + dispatch_output.hidden_states, + dispatch_output.topk_output, + ) topk_weights, topk_ids, _ = topk_output hidden_states_shape = hidden_states.shape diff --git a/python/sglang/srt/layers/moe/moe_runner/triton.py b/python/sglang/srt/layers/moe/moe_runner/triton.py index 8c77d7580..25fbcf770 100644 --- a/python/sglang/srt/layers/moe/moe_runner/triton.py +++ b/python/sglang/srt/layers/moe/moe_runner/triton.py @@ -377,7 +377,10 @@ def pre_permute_standard_to_triton( ) from sglang.srt.layers.moe.topk import TopKOutputChecker - hidden_states, topk_output = dispatch_output + hidden_states, topk_output = ( + dispatch_output.hidden_states, + dispatch_output.topk_output, + ) assert TopKOutputChecker.format_is_standard(topk_output) diff --git a/python/sglang/srt/layers/moe/token_dispatcher/base.py b/python/sglang/srt/layers/moe/token_dispatcher/base.py index f4bce4ae8..a5e61c107 100644 --- a/python/sglang/srt/layers/moe/token_dispatcher/base.py +++ b/python/sglang/srt/layers/moe/token_dispatcher/base.py @@ -1,8 +1,20 @@ from __future__ import annotations +import weakref from abc import ABC, abstractmethod from enum import Enum -from typing import TYPE_CHECKING, Protocol, TypeGuard, Union, runtime_checkable +from typing import ( + TYPE_CHECKING, + Any, + Callable, + Optional, + OrderedDict, + Protocol, + Tuple, + TypeGuard, + Union, + runtime_checkable, +) import torch @@ -16,6 +28,91 @@ if TYPE_CHECKING: StandardDispatchOutput, ) from sglang.srt.layers.moe.topk import TopKOutput + from sglang.srt.single_batch_overlap import CombineOverlapArgs + + +# ------------------------------ Dispatcher Hook ------------------------------------- + + +class _RemovableDispatcherHandle: + + next_id = 0 # Global counter for unique IDs + + def __init__(self, hooks_dict: OrderedDict): + self.id = _RemovableDispatcherHandle.next_id + _RemovableDispatcherHandle.next_id += 1 + self.weak_hooks_dict = weakref.ref(hooks_dict) + + def remove(self): + hooks_dict = self.weak_hooks_dict() + if hooks_dict is not None and self.id in hooks_dict: + del hooks_dict[self.id] + + +class _DispatcherBaseHooks: + + def __init__(self): + self.hook_dict = OrderedDict[int, Callable]() + + def register_hook(self, hook_fun: Callable) -> _RemovableDispatcherHandle: + handle = _RemovableDispatcherHandle(self.hook_dict) + self.hook_dict[handle.id] = hook_fun + return handle + + def __call__(self, *args, **kwargs) -> Optional[Any]: + raise NotImplementedError("This method should be overridden by subclasses") + + +class _PreDispatchHooks(_DispatcherBaseHooks): + + def __call__( + self, + dispatcher: BaseDispatcher, + hidden_states: torch.Tensor, + topk_output: TopKOutput, + ) -> Optional[Tuple[torch.Tensor, TopKOutput]]: + for hook_fun in self.hook_dict.values(): + hook_output = hook_fun(dispatcher, hidden_states, topk_output) + if hook_output is not None: + hidden_states, topk_output = hook_output + return hidden_states, topk_output + + +class _PostDispatchHooks(_DispatcherBaseHooks): + + def __call__( + self, dispatcher: BaseDispatcher, dispatch_output: DispatchOutput + ) -> Optional[DispatchOutput]: + for hook_fun in self.hook_dict.values(): + hook_output = hook_fun(dispatcher, dispatch_output) + if hook_output is not None: + dispatch_output = hook_output + return dispatch_output + + +class _PreCombineHooks(_DispatcherBaseHooks): + + def __call__( + self, dispatcher: BaseDispatcher, combine_input: CombineInput + ) -> Optional[CombineInput]: + for hook_fun in self.hook_dict.values(): + hook_output = hook_fun(dispatcher, combine_input) + if hook_output is not None: + combine_input = hook_output + return combine_input + + +class _PostCombineHooks(_DispatcherBaseHooks): + + def __call__( + self, dispatcher: BaseDispatcher, hidden_states: torch.Tensor + ) -> Optional[torch.Tensor]: + for hook_fun in self.hook_dict.values(): + hook_output = hook_fun(dispatcher, hidden_states) + if hook_output is not None: + hidden_states = hook_output + return hidden_states + # ------------------------------ Dispatch Output ------------------------------------- @@ -145,12 +242,112 @@ class BaseDispatcherConfig(ABC): class BaseDispatcher(ABC): """Base class for dispatchers.""" + def __init__(self): + self.quant_config: Optional[dict] = None + + # Overlap args + self.overlap_args: Optional[CombineOverlapArgs] = None + self.meta_overlap_args: Optional[dict] = None + + # Hooks + self._pre_dispatch_hooks: Optional[_PreDispatchHooks] = None + self._post_dispatch_hooks: Optional[_PostDispatchHooks] = None + self._pre_combine_hooks: Optional[_PreCombineHooks] = None + self._post_combine_hooks: Optional[_PostCombineHooks] = None + self._original_dispatch_func: Optional[Callable] = None + self._original_combine_func: Optional[Callable] = None + @abstractmethod def dispatch( - self, hidden_states: torch.Tensor, topk_output: TopKOutput, **kwargs + self, hidden_states: torch.Tensor, topk_output: TopKOutput ) -> DispatchOutput: pass + def _dispatch_with_hook( + self, hidden_states: torch.Tensor, topk_output: TopKOutput + ) -> DispatchOutput: + if self._pre_dispatch_hooks is not None: + hidden_states, topk_output = self._pre_dispatch_hooks( + self, hidden_states, topk_output + ) + dispatch_output = self._original_dispatch_func( + hidden_states=hidden_states, topk_output=topk_output + ) + if self._post_dispatch_hooks is not None: + dispatch_output = self._post_dispatch_hooks(self, dispatch_output) + return dispatch_output + + def _override_dispatch_func(self) -> None: + if self._original_dispatch_func is None: + self._original_dispatch_func = self.dispatch + self.dispatch = self._dispatch_with_hook + @abstractmethod - def combine(self, combine_input: CombineInput, **kwargs) -> torch.Tensor: + def combine(self, combine_input: CombineInput) -> torch.Tensor: pass + + def _combine_with_hook(self, combine_input: CombineInput) -> torch.Tensor: + if self._pre_combine_hooks is not None: + combine_input = self._pre_combine_hooks(self, combine_input) + hidden_states = self._original_combine_func(combine_input=combine_input) + if self._post_combine_hooks is not None: + hidden_states = self._post_combine_hooks(self, hidden_states) + return hidden_states + + def _override_combine_func(self) -> None: + if self._original_combine_func is None: + self._original_combine_func = self.combine + self.combine = self._combine_with_hook + + def register_pre_dispatch_hook( + self, + hook: Callable[ + [BaseDispatcher, torch.Tensor, TopKOutput], + Optional[Tuple[torch.Tensor, TopKOutput]], + ], + ) -> _RemovableDispatcherHandle: + if self._pre_dispatch_hooks is None: + self._pre_dispatch_hooks = _PreDispatchHooks() + self._override_dispatch_func() + handle = self._pre_dispatch_hooks.register_hook(hook) + return handle + + def register_post_dispatch_hook( + self, hook: Callable[[BaseDispatcher, DispatchOutput], Optional[DispatchOutput]] + ) -> _RemovableDispatcherHandle: + if self._post_dispatch_hooks is None: + self._post_dispatch_hooks = _PostDispatchHooks() + self._override_dispatch_func() + handle = self._post_dispatch_hooks.register_hook(hook) + return handle + + def register_pre_combine_hook( + self, hook: Callable[[BaseDispatcher, CombineInput], Optional[CombineInput]] + ) -> _RemovableDispatcherHandle: + if self._pre_combine_hooks is None: + self._pre_combine_hooks = _PreCombineHooks() + self._override_combine_func() + handle = self._pre_combine_hooks.register_hook(hook) + return handle + + def register_post_combine_hook( + self, hook: Callable[[BaseDispatcher, torch.Tensor], Optional[torch.Tensor]] + ) -> _RemovableDispatcherHandle: + if self._post_combine_hooks is None: + self._post_combine_hooks = _PostCombineHooks() + self._override_combine_func() + handle = self._post_combine_hooks.register_hook(hook) + return handle + + def set_quant_config(self, quant_config: dict) -> None: + self.quant_config = quant_config + + def set_overlap_args( + self, combine_overlap_args: CombineOverlapArgs, meta_overlap_args: dict + ) -> None: + self.overlap_args = combine_overlap_args + self.meta_overlap_args = meta_overlap_args + + def clear_overlap_args(self) -> None: + self.overlap_args = None + self.meta_overlap_args = None diff --git a/python/sglang/srt/layers/moe/token_dispatcher/deepep.py b/python/sglang/srt/layers/moe/token_dispatcher/deepep.py index 12fccaacb..24dd42b41 100644 --- a/python/sglang/srt/layers/moe/token_dispatcher/deepep.py +++ b/python/sglang/srt/layers/moe/token_dispatcher/deepep.py @@ -316,6 +316,11 @@ class _DeepEPDispatcherImplBase: self.handle = None + self.quant_config: Optional[dict] = None + + self.overlap_args: Optional[CombineOverlapArgs] = None + self.meta_overlap_args: Optional[dict] = None + def dispatch_a( self, hidden_states: torch.Tensor, @@ -331,7 +336,6 @@ class _DeepEPDispatcherImplBase: hidden_states: torch.Tensor, topk_ids: torch.Tensor, topk_weights: torch.Tensor, - overlap_args: Optional[CombineOverlapArgs] = None, ): raise NotImplementedError @@ -341,6 +345,19 @@ class _DeepEPDispatcherImplBase: def _get_buffer(self): raise NotImplementedError + def set_quant_config(self, quant_config: dict) -> None: + self.quant_config = quant_config + + def set_overlap_args( + self, combine_overlap_args: CombineOverlapArgs, meta_overlap_args: dict + ) -> None: + self.overlap_args = combine_overlap_args + self.meta_overlap_args = meta_overlap_args + + def clear_overlap_args(self) -> None: + self.overlap_args = None + self.meta_overlap_args = None + class _DeepEPDispatcherImplNormal(_DeepEPDispatcherImplBase): def __init__(self, async_finish: bool, **kwargs): @@ -461,7 +478,6 @@ class _DeepEPDispatcherImplNormal(_DeepEPDispatcherImplBase): hidden_states: torch.Tensor, topk_ids: torch.Tensor, topk_weights: torch.Tensor, - overlap_args: Optional[CombineOverlapArgs] = None, ): if deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM or _use_aiter or _is_npu: @@ -503,9 +519,6 @@ class _DeepEPDispatcherImplNormal(_DeepEPDispatcherImplBase): self.num_experts, ) - def set_quant_config(self, quant_config: dict): - self.quant_config = quant_config - class _DeepEPDispatcherImplLowLatency(_DeepEPDispatcherImplBase): def __init__(self, return_recv_hook: bool, **kwargs): @@ -617,17 +630,16 @@ class _DeepEPDispatcherImplLowLatency(_DeepEPDispatcherImplBase): hidden_states: torch.Tensor, topk_ids: torch.Tensor, topk_weights: torch.Tensor, - overlap_args: Optional[CombineOverlapArgs] = None, ): hidden_states, event, hook = self._combine_core( hidden_states, topk_ids, topk_weights, - overlap_args=overlap_args, ) - return hidden_states, event, hook, overlap_args + return hidden_states, event, hook - def combine_b(self, hidden_states, event, hook, overlap_args): + def combine_b(self, hidden_states, event, hook): + overlap_args = self.overlap_args if overlap_args is not None: overlap_args.stream.wait_stream(self.device_module.current_stream()) @@ -643,9 +655,9 @@ class _DeepEPDispatcherImplLowLatency(_DeepEPDispatcherImplBase): hidden_states: torch.Tensor, topk_ids: torch.Tensor, topk_weights: torch.Tensor, - overlap_args: Optional[CombineOverlapArgs] = None, ): buffer = self._get_buffer() + overlap_args = self.overlap_args ctx = nullcontext() if overlap_args is not None: @@ -685,9 +697,6 @@ class _DeepEPDispatcherImplLowLatency(_DeepEPDispatcherImplBase): self.num_experts, ) - def set_quant_config(self, quant_config: dict): - self.quant_config = quant_config - @dataclass class _Stage(Enum): @@ -711,6 +720,8 @@ class DeepEPDispatcher(BaseDispatcher): async_finish: bool = False, return_recv_hook: bool = False, ): + super().__init__() + self.deepep_mode = deepep_mode common_kwargs = dict( @@ -737,8 +748,12 @@ class DeepEPDispatcher(BaseDispatcher): self._stage = _Stage.INITIAL - def dispatch(self, *args, **kwargs) -> DispatchOutput: - self.dispatch_a(*args, **kwargs) + def dispatch( + self, + hidden_states: torch.Tensor, + topk_output: TopKOutput, + ) -> DispatchOutput: + self.dispatch_a(hidden_states, topk_output) ret = self.dispatch_b() return ret @@ -763,16 +778,14 @@ class DeepEPDispatcher(BaseDispatcher): def combine( self, combine_input: CombineInput, - overlap_args: Optional[CombineOverlapArgs] = None, - ) -> Tuple: - self.combine_a(combine_input, overlap_args) + ) -> torch.Tensor: + self.combine_a(combine_input) ret = self.combine_b() return ret def combine_a( self, combine_input: CombineInput, - overlap_args: Optional[CombineOverlapArgs] = None, ): hidden_states, topk_ids, topk_weights = combine_input self._update_stage(_Stage.AFTER_DISPATCH_B, _Stage.AFTER_COMBINE_A) @@ -780,7 +793,6 @@ class DeepEPDispatcher(BaseDispatcher): hidden_states=hidden_states, topk_ids=topk_ids, topk_weights=topk_weights, - overlap_args=overlap_args, ) self._combine_intermediate_state = inner_state @@ -805,7 +817,28 @@ class DeepEPDispatcher(BaseDispatcher): self._stage = new_stage def set_quant_config(self, quant_config: dict): + super().set_quant_config(quant_config) if self.deepep_mode.enable_low_latency(): self._low_latency_dispatcher.set_quant_config(quant_config) if self.deepep_mode.enable_normal(): self._normal_dispatcher.set_quant_config(quant_config) + + def set_overlap_args( + self, combine_overlap_args: CombineOverlapArgs, meta_overlap_args: dict + ): + super().set_overlap_args(combine_overlap_args, meta_overlap_args) + if self.deepep_mode.enable_low_latency(): + self._low_latency_dispatcher.set_overlap_args( + combine_overlap_args, meta_overlap_args + ) + if self.deepep_mode.enable_normal(): + self._normal_dispatcher.set_overlap_args( + combine_overlap_args, meta_overlap_args + ) + + def clear_overlap_args(self): + super().clear_overlap_args() + if self.deepep_mode.enable_low_latency(): + self._low_latency_dispatcher.clear_overlap_args() + if self.deepep_mode.enable_normal(): + self._normal_dispatcher.clear_overlap_args() diff --git a/python/sglang/srt/layers/moe/token_dispatcher/mooncake.py b/python/sglang/srt/layers/moe/token_dispatcher/mooncake.py index d3e7831a5..d5a109072 100644 --- a/python/sglang/srt/layers/moe/token_dispatcher/mooncake.py +++ b/python/sglang/srt/layers/moe/token_dispatcher/mooncake.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from dataclasses import dataclass -from typing import TYPE_CHECKING, NamedTuple, Optional, Tuple +from typing import TYPE_CHECKING, NamedTuple, Optional from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder @@ -302,6 +302,8 @@ class MooncakeEPDispatcher(BaseDispatcher): async_finish: bool = False, return_recv_hook: bool = False, ): + super().__init__() + self.deepep_mode = deepep_mode if self.deepep_mode.enable_low_latency(): @@ -321,8 +323,12 @@ class MooncakeEPDispatcher(BaseDispatcher): self._stage = _Stage.INITIAL - def dispatch(self, *args, **kwargs) -> DispatchOutput: - self.dispatch_a(*args, **kwargs) + def dispatch( + self, + hidden_states: torch.Tensor, + topk_output: TopKOutput, + ) -> DispatchOutput: + self.dispatch_a(hidden_states, topk_output) ret = self.dispatch_b() return ret @@ -347,16 +353,14 @@ class MooncakeEPDispatcher(BaseDispatcher): def combine( self, combine_input: CombineInput, - overlap_args: Optional[CombineOverlapArgs] = None, - ) -> Tuple: - self.combine_a(combine_input, overlap_args) + ) -> torch.Tensor: + self.combine_a(combine_input) ret = self.combine_b() return ret def combine_a( self, combine_input: CombineInput, - overlap_args: Optional[CombineOverlapArgs] = None, ): hidden_states, topk_ids, topk_weights = combine_input self._update_stage(_Stage.AFTER_DISPATCH_B, _Stage.AFTER_COMBINE_A) @@ -364,7 +368,7 @@ class MooncakeEPDispatcher(BaseDispatcher): hidden_states=hidden_states, topk_ids=topk_ids, topk_weights=topk_weights, - overlap_args=overlap_args, + overlap_args=self.overlap_args, ) self._combine_intermediate_state = inner_state @@ -387,6 +391,3 @@ class MooncakeEPDispatcher(BaseDispatcher): def _update_stage(self, old_stage, new_stage): assert self._stage == old_stage self._stage = new_stage - - def set_quant_config(self, quant_config: dict): - pass diff --git a/python/sglang/srt/layers/moe/token_dispatcher/standard.py b/python/sglang/srt/layers/moe/token_dispatcher/standard.py index 26105a27c..2c959c799 100644 --- a/python/sglang/srt/layers/moe/token_dispatcher/standard.py +++ b/python/sglang/srt/layers/moe/token_dispatcher/standard.py @@ -1,12 +1,21 @@ from __future__ import annotations -from typing import TYPE_CHECKING, NamedTuple +from typing import TYPE_CHECKING, NamedTuple, Optional import torch 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 ( + use_symmetric_memory, +) +from sglang.srt.layers.dp_attention import ( + get_dp_global_num_tokens, + get_local_dp_buffer, + is_allocation_symmetric, ) from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig from sglang.srt.layers.moe.token_dispatcher.base import ( @@ -16,9 +25,12 @@ from sglang.srt.layers.moe.token_dispatcher.base import ( DispatchOutput, DispatchOutputFormat, ) -from sglang.srt.layers.moe.topk import TopKOutput, TopKOutputChecker -from sglang.srt.layers.moe.utils import get_moe_runner_backend -from sglang.srt.utils import get_bool_env_var, is_hip +from sglang.srt.layers.moe.topk import StandardTopKOutput, TopKOutput, TopKOutputChecker +from sglang.srt.layers.moe.utils import ( + get_moe_runner_backend, + should_use_flashinfer_cutlass_moe_fp4_allgather, +) +from sglang.srt.utils.common import get_bool_env_var, is_hip, is_sm120_supported _is_hip = is_hip() _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip @@ -27,10 +39,22 @@ if TYPE_CHECKING: from sglang.srt.layers.moe.topk import TopKOutput +try: + if is_sm120_supported(): + from flashinfer import fp4_quantize + else: + from sgl_kernel import scaled_fp4_quant as fp4_quantize + + from flashinfer import fp4_quantize as fp4_quantize_flashinfer +except ImportError: + fp4_quantize = None + + class StandardDispatchOutput(NamedTuple): """Standard dispatch output.""" hidden_states: torch.Tensor + hidden_states_scale: Optional[torch.Tensor] topk_output: TopKOutput @property @@ -57,6 +81,7 @@ assert isinstance(StandardCombineInput, CombineInput) class StandardDispatcher(BaseDispatcher): def __init__(self, moe_runner_config: MoeRunnerConfig): + super().__init__() self.moe_ep_size = get_moe_expert_parallel_world_size() self.enable_flashinfer_cutlass_moe = ( get_moe_runner_backend().is_flashinfer_cutlass() @@ -71,7 +96,47 @@ class StandardDispatcher(BaseDispatcher): def dispatch( self, hidden_states: torch.Tensor, topk_output: TopKOutput - ) -> DispatchOutput: + ) -> StandardDispatchOutput: + + if should_use_flashinfer_cutlass_moe_fp4_allgather(): + # all-gather fp4 hidden states + from flashinfer import nvfp4_block_scale_interleave + + global_scale = self.quant_config.get("input_global_scale", None) + assert global_scale is not None, "input_global_scale is not set" + topk_weights, topk_ids = topk_output.topk_weights, topk_output.topk_ids + + # Quantize before comm, swizzle after. + with use_symmetric_memory( + get_tp_group(), disabled=not is_allocation_symmetric() + ): + if hidden_states.shape[0] > 0: + x, x_sf = fp4_quantize_flashinfer( + hidden_states, global_scale, is_sf_swizzled_layout=False + ) + else: + x_col = hidden_states.shape[1] + x = torch.zeros( + 0, x_col // 2, dtype=torch.uint8, device=hidden_states.device + ) + x_sf = torch.zeros( + 0, x_col // 16, dtype=torch.uint8, device=hidden_states.device + ) + topk_weights, topk_ids, x, x_sf = get_tp_group().all_gatherv( + [topk_weights, topk_ids, x, x_sf], sizes=get_dp_global_num_tokens() + ) + x_sf = nvfp4_block_scale_interleave(x_sf) + + hidden_states = x + hidden_states_scale = x_sf + topk_output = StandardTopKOutput( + topk_weights=topk_weights, + topk_ids=topk_ids, + router_logits=topk_output.router_logits, # never tested + ) + else: + hidden_states = hidden_states + hidden_states_scale = None if ( self.moe_ep_size > 1 @@ -110,16 +175,18 @@ class StandardDispatcher(BaseDispatcher): raise NotImplementedError() return StandardDispatchOutput( - hidden_states=hidden_states, topk_output=topk_output + hidden_states=hidden_states, + hidden_states_scale=hidden_states_scale, + topk_output=topk_output, ) - def combine(self, combine_input: CombineInput) -> torch.Tensor: - if isinstance(combine_input, StandardCombineInput): - return combine_input.hidden_states - else: - # TODO: this branch should be removed in the future - assert isinstance(combine_input, torch.Tensor) - return combine_input - - def set_quant_config(self, quant_config: dict): - pass + def combine(self, combine_input: StandardCombineInput) -> torch.Tensor: + (hidden_states,) = combine_input + if should_use_flashinfer_cutlass_moe_fp4_allgather(): + hidden_states, global_hidden_states = get_local_dp_buffer(), hidden_states + get_tp_group().reduce_scatterv( + global_hidden_states, + output=hidden_states, + sizes=get_dp_global_num_tokens(), + ) + return hidden_states diff --git a/python/sglang/srt/layers/quantization/modelopt_quant.py b/python/sglang/srt/layers/quantization/modelopt_quant.py index 215b17a7e..e627a3151 100755 --- a/python/sglang/srt/layers/quantization/modelopt_quant.py +++ b/python/sglang/srt/layers/quantization/modelopt_quant.py @@ -12,17 +12,12 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) from sglang.srt.environ import envs -from sglang.srt.layers.dp_attention import ( - get_dp_global_num_tokens, - get_local_dp_buffer, - is_allocation_symmetric, -) +from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.moe import ( MoeRunner, MoeRunnerBackend, MoeRunnerConfig, get_moe_runner_backend, - should_use_flashinfer_cutlass_moe_fp4_allgather, ) from sglang.srt.layers.moe.cutlass_moe_params import CutlassMoEParams, CutlassMoEType from sglang.srt.layers.moe.moe_runner.triton import TritonMoeQuantInfo @@ -48,8 +43,12 @@ from sglang.srt.layers.quantization.utils import ( requantize_with_max_scale, ) from sglang.srt.layers.radix_attention import RadixAttention -from sglang.srt.utils import get_bool_env_var, is_cuda, next_power_of_2 -from sglang.srt.utils.common import is_sm120_supported +from sglang.srt.utils.common import ( + get_bool_env_var, + is_cuda, + is_sm120_supported, + next_power_of_2, +) if TYPE_CHECKING: from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE @@ -65,7 +64,6 @@ try: else: from sgl_kernel import scaled_fp4_quant as fp4_quantize - from flashinfer import fp4_quantize as fp4_quantize_flashinfer except ImportError: fp4_quantize = None @@ -1628,11 +1626,10 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): self, layer: FusedMoE, dispatch_output: StandardDispatchOutput, - forward_shared_experts=None, - alt_stream=None, ) -> CombineInput: x = dispatch_output.hidden_states + x_sf = dispatch_output.hidden_states_scale topk_output = dispatch_output.topk_output assert ( @@ -1656,39 +1653,13 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): # and fp4 quantized weights loaded from the checkpoint topk_weights, topk_ids = topk_output.topk_weights, topk_output.topk_ids - output_dtype = x.dtype - original_col = x.shape[1] - x_sf = None - - if should_use_flashinfer_cutlass_moe_fp4_allgather(): - from flashinfer import nvfp4_block_scale_interleave - - # Quantize before comm, swizzle after. - with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() - ): - if x.shape[0] > 0: - x, x_sf = fp4_quantize_flashinfer( - x, layer.w13_input_scale_quant, is_sf_swizzled_layout=False - ) - else: - x_col = x.shape[1] - x = torch.zeros( - 0, x_col // 2, dtype=torch.uint8, device=x.device - ) - x_sf = torch.zeros( - 0, x_col // 16, dtype=torch.uint8, device=x.device - ) - topk_weights, topk_ids, x, x_sf = get_tp_group().all_gatherv( - [topk_weights, topk_ids, x, x_sf], sizes=get_dp_global_num_tokens() - ) - x_sf = nvfp4_block_scale_interleave(x_sf) + output_dtype = torch.bfloat16 with use_symmetric_memory( get_tp_group(), disabled=not is_allocation_symmetric() ): symm_output = torch.empty( - x.shape[0], original_col, dtype=output_dtype, device=x.device + x.shape[0], x.shape[1] * 2, dtype=output_dtype, device=x.device ) output = flashinfer_cutlass_fused_moe( @@ -1714,20 +1685,6 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): tp_rank=layer.moe_tp_rank, tune_max_num_tokens=next_power_of_2(x.shape[0]), )[0] - if should_use_flashinfer_cutlass_moe_fp4_allgather(): - output, global_output = get_local_dp_buffer(), output - - if forward_shared_experts is not None: - alt_stream.wait_stream(torch.cuda.current_stream()) - with torch.cuda.stream(alt_stream): - forward_shared_experts() - - get_tp_group().reduce_scatterv( - global_output, output=output, sizes=get_dp_global_num_tokens() - ) - - if forward_shared_experts is not None: - torch.cuda.current_stream().wait_stream(alt_stream) from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput @@ -1762,7 +1719,6 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): x: tuple[torch.Tensor, Optional[torch.Tensor]], masked_m: torch.Tensor, moe_runner_config: MoeRunnerConfig, - down_gemm_overlap_args: Optional["DownGemmOverlapArgs"], ) -> torch.Tensor: assert ( moe_runner_config.activation == "silu" @@ -1777,6 +1733,10 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase): flashinfer_cutedsl_moe_masked, ) + down_gemm_overlap_args: Optional[DownGemmOverlapArgs] = getattr( + layer, "down_gemm_overlap_args", None + ) + out = flashinfer_cutedsl_moe_masked( hidden_states=x, input_global_scale=( diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index b7ed99d56..323ad864a 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -93,6 +93,11 @@ from sglang.srt.layers.moe import ( from sglang.srt.layers.moe.ep_moe.layer import DeepEPMoE, get_moe_impl_class from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE from sglang.srt.layers.moe.kt_ep_wrapper import KTEPWrapperMethod +from sglang.srt.layers.moe.token_dispatcher.base import ( + BaseDispatcher, + CombineInput, + DispatchOutput, +) from sglang.srt.layers.moe.topk import TopK, TopKOutputFormat from sglang.srt.layers.moe.utils import RoutingMethodType from sglang.srt.layers.quantization.base_config import QuantizationConfig @@ -131,7 +136,7 @@ from sglang.srt.model_loader.utils import ( ) from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.server_args import get_global_server_args -from sglang.srt.single_batch_overlap import SboFlags +from sglang.srt.single_batch_overlap import SboFlags, compute_overlap_args from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.two_batch_overlap import model_forward_maybe_tbo from sglang.srt.utils import ( @@ -847,7 +852,9 @@ class DeepseekV2MoE(nn.Module): return self.forward_cpu(hidden_states, should_allreduce_fusion) if hidden_states.shape[0] > 0: - if not self._fuse_shared_experts_inside_sbo: + if ( + not self._fuse_shared_experts_inside_sbo + ): # TODO: check if it supports mtp shared_output = self._forward_shared_experts( hidden_states, gemm_output_zero_allocator ) @@ -861,23 +868,36 @@ class DeepseekV2MoE(nn.Module): if self._fuse_shared_experts_inside_sbo: shared_output = None - def _forward_shared_experts_and_put_results(): + def _pre_combine_hook( + dispatcher: BaseDispatcher, combine_input: CombineInput + ): + nonlocal shared_output - shared_output = self._forward_shared_experts( - hidden_states, gemm_output_zero_allocator - ) + self.alt_stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(self.alt_stream): + shared_output = self._forward_shared_experts( + hidden_states, gemm_output_zero_allocator + ) + + pre_combine_hook_handle.remove() + + def _post_combine_hook( + dispatcher: BaseDispatcher, hidden_states: torch.Tensor + ): + nonlocal shared_output + torch.cuda.current_stream().wait_stream(self.alt_stream) + post_combine_hook_handle.remove() + + pre_combine_hook_handle = self.experts.dispatcher.register_pre_combine_hook( + _pre_combine_hook + ) + post_combine_hook_handle = ( + self.experts.dispatcher.register_post_combine_hook(_post_combine_hook) + ) final_hidden_states = self.experts( hidden_states, topk_output, - **( - dict( - forward_shared_experts=_forward_shared_experts_and_put_results, - alt_stream=self.alt_stream, - ) - if self._fuse_shared_experts_inside_sbo - else {} - ), ) if ( not _is_cuda @@ -961,10 +981,11 @@ class DeepseekV2MoE(nn.Module): forward_batch: ForwardBatch, ) -> torch.Tensor: shared_output = None + sbo_enabled_flag = self._fuse_shared_experts_inside_sbo and not self.is_nextn if hidden_states.shape[0] > 0: # router_logits: (num_tokens, n_experts) router_logits = self.gate(hidden_states, forward_batch=forward_batch) - if not self._fuse_shared_experts_inside_sbo: + if not sbo_enabled_flag: shared_output = self._forward_shared_experts(hidden_states) topk_output = self.topk( hidden_states, @@ -977,26 +998,67 @@ class DeepseekV2MoE(nn.Module): else: topk_output = self.topk.empty_topk_output(hidden_states.device) - if self._fuse_shared_experts_inside_sbo: + # SBO is not yet implemented for NextN + if sbo_enabled_flag: shared_output = None - def _forward_shared_experts_and_put_results(): + def _post_dispatch_hook( + dispatcher: BaseDispatcher, dispatch_output: DispatchOutput + ): + + combine_overlap_args, down_gemm_overlap_args, meta_overlap_args = ( + compute_overlap_args(dispatch_output, self.alt_stream) + ) + dispatcher.set_overlap_args( + combine_overlap_args=combine_overlap_args, + meta_overlap_args=meta_overlap_args, + ) + self.experts.set_overlap_args( + down_gemm_overlap_args=down_gemm_overlap_args, + meta_overlap_args=meta_overlap_args, + ) + + post_dispatch_hook_handle.remove() + + def _pre_combine_hook( + dispatcher: BaseDispatcher, combine_input: CombineInput + ): + nonlocal shared_output - shared_output = self._forward_shared_experts(hidden_states) + + if ( + e := dispatcher.meta_overlap_args.get("record_event_after_down") + ) is not None: + e.record() + + # TODO reduce sm for non-deepgemm + with deep_gemm_wrapper.configure_deep_gemm_num_sms( + dispatcher.meta_overlap_args["compute_num_sms"] + ): + shared_output = self._forward_shared_experts(hidden_states) + + pre_combine_hook_handle.remove() + + def _post_combine_hook( + dispatcher: BaseDispatcher, hidden_states: torch.Tensor + ): + dispatcher.clear_overlap_args() + self.experts.clear_overlap_args() + post_combine_hook_handle.remove() + + post_dispatch_hook_handle = ( + self.experts.dispatcher.register_post_dispatch_hook(_post_dispatch_hook) + ) + pre_combine_hook_handle = self.experts.dispatcher.register_pre_combine_hook( + _pre_combine_hook + ) + post_combine_hook_handle = ( + self.experts.dispatcher.register_post_combine_hook(_post_combine_hook) + ) final_hidden_states = self.experts( hidden_states=hidden_states, topk_output=topk_output, - **( - dict( - forward_shared_experts=_forward_shared_experts_and_put_results, - alt_stream=self.alt_stream, - # SBO is not yet implemented for NextN - disable_sbo=self.is_nextn, - ) - if self._fuse_shared_experts_inside_sbo - else {} - ), ) if shared_output is not None: diff --git a/python/sglang/srt/single_batch_overlap.py b/python/sglang/srt/single_batch_overlap.py index 93a35f5a6..7b324c37d 100644 --- a/python/sglang/srt/single_batch_overlap.py +++ b/python/sglang/srt/single_batch_overlap.py @@ -15,19 +15,14 @@ from __future__ import annotations from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, Callable, Optional +from typing import Optional import torch -from sglang.srt.layers import deep_gemm_wrapper from sglang.srt.layers.moe import get_moe_runner_backend -from sglang.srt.layers.moe.topk import TopKOutput from sglang.srt.layers.moe.utils import is_sbo_enabled from sglang.srt.utils import get_int_env_var -if TYPE_CHECKING: - from sglang.srt.layers.moe.fused_moe_triton import FusedMoE - class SboFlags: # TODO may have: "enable_dispatch_shared_one_stream_overlap", "enable_dispatch_gateup_gemm_two_stream_overlap", ... @@ -68,46 +63,8 @@ class DownGemmOverlapArgs: start_event: torch.cuda.Event -def execute_sbo( - forward_shared_experts: Callable[[], Any], - experts: FusedMoE, - hidden_states: torch.Tensor, - topk_output: TopKOutput, - alt_stream: Optional[torch.cuda.Stream] = None, - disable_sbo: bool = False, -): - - dispatch_output = experts.dispatcher.dispatch( - hidden_states=hidden_states, topk_output=topk_output - ) - - combine_overlap_args, down_gemm_overlap_args, meta_overlap_args = ( - _compute_overlap_args(dispatch_output, alt_stream, disable_sbo=disable_sbo) - ) - - combine_input = experts.run_moe_core( - dispatch_output, down_gemm_overlap_args=down_gemm_overlap_args - ) - if (e := meta_overlap_args.get("record_event_after_down")) is not None: - e.record() - - if (not disable_sbo) and SboFlags.enable_combine_shared_two_stream_overlap(): - # TODO reduce sm for non-deepgemm - with deep_gemm_wrapper.configure_deep_gemm_num_sms( - meta_overlap_args["compute_num_sms"] - ): - forward_shared_experts() - - hidden_states = experts.dispatcher.combine( - combine_input=combine_input, - overlap_args=combine_overlap_args, - ) - - return hidden_states - - -def _compute_overlap_args(dispatch_output, alt_stream, disable_sbo): - if disable_sbo or not ( +def compute_overlap_args(dispatch_output, alt_stream): + if not ( SboFlags.enable_combine_down_gemm_two_stream_overlap() or SboFlags.enable_combine_shared_two_stream_overlap() ): diff --git a/python/sglang/srt/two_batch_overlap.py b/python/sglang/srt/two_batch_overlap.py index 16fa3ef15..97cd7dff4 100644 --- a/python/sglang/srt/two_batch_overlap.py +++ b/python/sglang/srt/two_batch_overlap.py @@ -25,6 +25,7 @@ from sglang.srt.layers.moe.token_dispatcher import ( DeepEPDispatcher, MooncakeEPDispatcher, ) +from sglang.srt.layers.moe.token_dispatcher.base import BaseDispatcher from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.model_executor.forward_batch_info import ( ForwardBatch, @@ -39,6 +40,7 @@ from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_ if TYPE_CHECKING: from sglang.srt.layers.moe.token_dispatcher import DispatchOutput + from sglang.srt.single_batch_overlap import CombineOverlapArgs from sglang.srt.speculative.eagle_info import EagleVerifyInput _is_hip = is_hip() @@ -970,8 +972,9 @@ def _model_forward_tbo_merge_outputs(output_a, output_b): # -------------------------------- Utilities and wrappers --------------------------------------- -class MaybeTboDeepEPDispatcher: +class MaybeTboDeepEPDispatcher(BaseDispatcher): def __init__(self, **kwargs): + super().__init__() num_inner_dispatchers = 2 if is_tbo_enabled() else 1 if get_moe_a2a_backend().is_deepep(): self._inners = [ @@ -1004,5 +1007,18 @@ class MaybeTboDeepEPDispatcher: return self._execute("combine_b", **kwargs) def set_quant_config(self, quant_config: dict): + super().set_quant_config(quant_config) for inner in self._inners: inner.set_quant_config(quant_config) + + def set_overlap_args( + self, combine_overlap_args: CombineOverlapArgs, meta_overlap_args: dict + ): + super().set_overlap_args(combine_overlap_args, meta_overlap_args) + for inner in self._inners: + inner.set_overlap_args(combine_overlap_args, meta_overlap_args) + + def clear_overlap_args(self): + super().clear_overlap_args() + for inner in self._inners: + inner.clear_overlap_args() diff --git a/test/srt/run_suite.py b/test/srt/run_suite.py index 511142b0c..48510c959 100644 --- a/test/srt/run_suite.py +++ b/test/srt/run_suite.py @@ -180,7 +180,7 @@ suites = { "per-commit-8-gpu-b200": [], "per-commit-4-gpu-gb200": [ TestFile("test_cutedsl_moe.py", 300), - TestFile("test_deepseek_v3_cutedsl_4gpu.py", 3600), + TestFile("test_deepseek_v3_cutedsl_4gpu.py", 590), ], "per-commit-4-gpu-deepep": [ TestFile("ep/test_deepep_small.py", 531), diff --git a/test/srt/test_deepseek_v3_cutedsl_4gpu.py b/test/srt/test_deepseek_v3_cutedsl_4gpu.py index 260575ef0..63b5fc20b 100644 --- a/test/srt/test_deepseek_v3_cutedsl_4gpu.py +++ b/test/srt/test_deepseek_v3_cutedsl_4gpu.py @@ -14,6 +14,7 @@ from sglang.test.test_utils import ( ) +@unittest.skip("See https://github.com/sgl-project/sglang/issues/12533") class TestDeepseekR1Nvfp4CuteDSLDeepEP(CustomTestCase): @classmethod def setUpClass(cls): @@ -83,5 +84,79 @@ class TestDeepseekR1Nvfp4CuteDSLDeepEP(CustomTestCase): self.assertGreater(metrics["accuracy"], 0.92) +class TestDummyWithSBO(CustomTestCase): + @classmethod + def setUpClass(cls): + cls.model = try_cached_model(DEFAULT_DEEPSEEK_NVFP4_MODEL_FOR_TEST) + cls.base_url = DEFAULT_URL_FOR_TEST + other_args = [ + "--trust-remote-code", + "--disable-radix-cache", + "--mem-fraction-static", + "0.05", + "--max-prefill-tokens", + "16384", + "--max-running-requests", + "256", + "--chunked-prefill-size", + "1024", + "--cuda-graph-bs", + "64", + "--tp", + "4", + "--dp", + "4", + "--ep", + "4", + "--moe-dense-tp-size", + "1", + "--enable-dp-attention", + "--quantization", + "modelopt_fp4", + "--attention-backend", + "trtllm_mla", + "--moe-a2a-backend", + "deepep", + "--moe-runner-backend", + "flashinfer_cutedsl", + "--deepep-mode", + "low_latency", + "--json-model-override-args", + '{"num_hidden_layers": 1, "first_k_dense_replace": 0, "n_routed_experts": 24}', + "--enable-single-batch-overlap", + "--load-format", + "dummy", + ] + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=other_args, + env={ + **os.environ, + "SGLANG_DEEPEP_BF16_DISPATCH": "1", + "SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "256", + "SGLANG_CUTEDSL_MOE_NVFP4_DISPATCH": "0", + }, + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + def test_gsm8k(self): + args = SimpleNamespace( + num_shots=0, + data_path=None, + num_questions=512, + parallel=512, + max_new_tokens=16, + host="http://127.0.0.1", + port=int(self.base_url.split(":")[-1]), + ) + metrics = run_eval_few_shot_gsm8k(args) + print(f"Eval accuracy of GSM8K: {metrics=}") + + if __name__ == "__main__": unittest.main()