diff --git a/python/sglang/srt/models/minimax_m3.py b/python/sglang/srt/models/minimax_m3.py index 7ff7c0b36..5becb4e57 100644 --- a/python/sglang/srt/models/minimax_m3.py +++ b/python/sglang/srt/models/minimax_m3.py @@ -79,13 +79,14 @@ from sglang.srt.model_executor.forward_context import ( get_forward_context, has_forward_context, ) +from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_loader.weight_utils import ( default_weight_loader, maybe_remap_kv_scale_name, ) from sglang.srt.models.minimax_m2 import MiniMaxM2RMSNormTP from sglang.srt.models.utils import WeightsMapper -from sglang.srt.runtime_context import get_exec, get_parallel +from sglang.srt.runtime_context import get_exec, get_parallel, get_stream from sglang.srt.utils import ( add_prefix, get_device_sm, @@ -312,10 +313,12 @@ class MiniMaxM3MoE(nn.Module): config: PretrainedConfig, layer_id: int, quant_config: Optional[QuantizationConfig] = None, + alt_stream: Optional[torch.cuda.Stream] = None, prefix: str = "", ): super().__init__() self.tp_size = get_parallel().tp_size + self.alt_stream = alt_stream self.n_shared_experts = getattr(config, "n_shared_experts", None) self.num_fused_shared_experts = ( 0 if is_shared_experts_fusion_disabled() else config.n_shared_experts @@ -425,14 +428,24 @@ class MiniMaxM3MoE(nn.Module): use_reduce_scatter: bool = False, ) -> torch.Tensor: if hidden_states.shape[0] > 0: - shared_output = self._forward_shared_experts(hidden_states) - router_logits = self._compute_router_logits(hidden_states) - topk_output = self.topk(hidden_states, router_logits) + if ( + self.alt_stream is not None + and self.shared_experts is not None + and get_is_capture_mode() + ): + current_stream = torch.cuda.current_stream() + self.alt_stream.wait_stream(current_stream) + shared_output = self._forward_shared_experts(hidden_states) + with torch.cuda.stream(self.alt_stream): + final_hidden_states = self._forward_router_experts(hidden_states) + current_stream.wait_stream(self.alt_stream) + else: + shared_output = self._forward_shared_experts(hidden_states) + final_hidden_states = self._forward_router_experts(hidden_states) else: shared_output = None topk_output = self.topk.empty_topk_output(hidden_states.device) - - final_hidden_states = self.experts(hidden_states, topk_output) + final_hidden_states = self.experts(hidden_states, topk_output) if shared_output is not None: final_hidden_states = final_hidden_states + shared_output @@ -441,6 +454,11 @@ class MiniMaxM3MoE(nn.Module): return final_hidden_states + def _forward_router_experts(self, hidden_states: torch.Tensor) -> torch.Tensor: + router_logits = self._compute_router_logits(hidden_states) + topk_output = self.topk(hidden_states, router_logits) + return self.experts(hidden_states, topk_output) + def forward_deepep( self, hidden_states: torch.Tensor, forward_batch: ForwardBatch ) -> torch.Tensor: @@ -1233,6 +1251,7 @@ class MiniMaxM3DecoderLayer(nn.Module): config: PretrainedConfig, layer_id: int, quant_config: Optional[QuantizationConfig] = None, + alt_stream: Optional[torch.cuda.Stream] = None, prefix: str = "", ) -> None: super().__init__() @@ -1272,6 +1291,7 @@ class MiniMaxM3DecoderLayer(nn.Module): config=config, layer_id=layer_id, quant_config=quant_config, + alt_stream=alt_stream, prefix=add_prefix("mlp", prefix), ) else: @@ -1416,11 +1436,14 @@ class MiniMaxM3Model(nn.Module): else: self.embed_tokens = PPMissingLayer() + alt_stream = get_stream("alt") if _is_cuda else None + def layer_fn(idx, prefix: str) -> nn.Module: return MiniMaxM3DecoderLayer( config=config, layer_id=idx, quant_config=quant_config, + alt_stream=alt_stream, prefix=prefix, )