From b1e13e7cea7b3a3a5687e70ba4a7839696db9df5 Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Tue, 28 Oct 2025 01:23:06 -0700 Subject: [PATCH] [hotfix] Incorrect CombineOverlapArgs in SBO (#12230) --- python/sglang/srt/layers/moe/ep_moe/layer.py | 1 - .../srt/layers/moe/token_dispatcher/deepep.py | 21 +++++++++++-------- python/sglang/srt/single_batch_overlap.py | 5 ++++- 3 files changed, 16 insertions(+), 11 deletions(-) diff --git a/python/sglang/srt/layers/moe/ep_moe/layer.py b/python/sglang/srt/layers/moe/ep_moe/layer.py index 8d00a4010..e9f736a18 100644 --- a/python/sglang/srt/layers/moe/ep_moe/layer.py +++ b/python/sglang/srt/layers/moe/ep_moe/layer.py @@ -235,7 +235,6 @@ class DeepEPMoE(FusedMoE): hidden_states=output, topk_ids=dispatch_output.topk_ids, topk_weights=dispatch_output.topk_weights, - overlap_args=down_gemm_overlap_args, ) def combine( diff --git a/python/sglang/srt/layers/moe/token_dispatcher/deepep.py b/python/sglang/srt/layers/moe/token_dispatcher/deepep.py index 33fc824fb..12fccaacb 100644 --- a/python/sglang/srt/layers/moe/token_dispatcher/deepep.py +++ b/python/sglang/srt/layers/moe/token_dispatcher/deepep.py @@ -97,7 +97,6 @@ class DeepEPNormalCombineInput(NamedTuple): hidden_states: torch.Tensor topk_ids: torch.Tensor topk_weights: torch.Tensor - overlap_args: Optional[CombineOverlapArgs] = None @property def format(self) -> CombineInputFormat: @@ -110,7 +109,6 @@ class DeepEPLLCombineInput(NamedTuple): hidden_states: torch.Tensor topk_ids: torch.Tensor topk_weights: torch.Tensor - overlap_args: Optional[CombineOverlapArgs] = None @property def format(self) -> CombineInputFormat: @@ -333,7 +331,7 @@ class _DeepEPDispatcherImplBase: hidden_states: torch.Tensor, topk_ids: torch.Tensor, topk_weights: torch.Tensor, - overlap_args: Optional["CombineOverlapArgs"], + overlap_args: Optional[CombineOverlapArgs] = None, ): raise NotImplementedError @@ -463,7 +461,7 @@ class _DeepEPDispatcherImplNormal(_DeepEPDispatcherImplBase): hidden_states: torch.Tensor, topk_ids: torch.Tensor, topk_weights: torch.Tensor, - overlap_args: Optional["CombineOverlapArgs"], + overlap_args: Optional[CombineOverlapArgs] = None, ): if deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM or _use_aiter or _is_npu: @@ -619,7 +617,7 @@ class _DeepEPDispatcherImplLowLatency(_DeepEPDispatcherImplBase): hidden_states: torch.Tensor, topk_ids: torch.Tensor, topk_weights: torch.Tensor, - overlap_args: Optional["CombineOverlapArgs"], + overlap_args: Optional[CombineOverlapArgs] = None, ): hidden_states, event, hook = self._combine_core( hidden_states, @@ -645,7 +643,7 @@ class _DeepEPDispatcherImplLowLatency(_DeepEPDispatcherImplBase): hidden_states: torch.Tensor, topk_ids: torch.Tensor, topk_weights: torch.Tensor, - overlap_args: Optional["CombineOverlapArgs"], + overlap_args: Optional[CombineOverlapArgs] = None, ): buffer = self._get_buffer() @@ -762,16 +760,21 @@ class DeepEPDispatcher(BaseDispatcher): del self._dispatch_intermediate_state return self._get_impl().dispatch_b(*inner_state) - def combine(self, combine_input: CombineInput) -> Tuple: - self.combine_a(combine_input) + def combine( + self, + combine_input: CombineInput, + overlap_args: Optional[CombineOverlapArgs] = None, + ) -> Tuple: + self.combine_a(combine_input, overlap_args) ret = self.combine_b() return ret def combine_a( self, combine_input: CombineInput, + overlap_args: Optional[CombineOverlapArgs] = None, ): - hidden_states, topk_ids, topk_weights, overlap_args = combine_input + hidden_states, topk_ids, topk_weights = combine_input self._update_stage(_Stage.AFTER_DISPATCH_B, _Stage.AFTER_COMBINE_A) inner_state = self._get_impl().combine_a( hidden_states=hidden_states, diff --git a/python/sglang/srt/single_batch_overlap.py b/python/sglang/srt/single_batch_overlap.py index ad7630cf8..93a35f5a6 100644 --- a/python/sglang/srt/single_batch_overlap.py +++ b/python/sglang/srt/single_batch_overlap.py @@ -98,7 +98,10 @@ def execute_sbo( ): forward_shared_experts() - hidden_states = experts.dispatcher.combine(combine_input=combine_input) + hidden_states = experts.dispatcher.combine( + combine_input=combine_input, + overlap_args=combine_overlap_args, + ) return hidden_states