diff --git a/python/sglang/srt/batch_overlap/operations_strategy.py b/python/sglang/srt/batch_overlap/operations_strategy.py index 152e4874d..41f40275e 100644 --- a/python/sglang/srt/batch_overlap/operations_strategy.py +++ b/python/sglang/srt/batch_overlap/operations_strategy.py @@ -51,6 +51,15 @@ class OperationsStrategy: for layer in layers ] ) + elif layer_name == "MiMoV2DecoderLayer": + return OperationsStrategy.concat( + [ + _compute_moe_mimov2_layer_operations_strategy_tbo( + layer, forward_mode + ) + for layer in layers + ] + ) else: raise NotImplementedError @@ -209,3 +218,78 @@ def _compute_moe_qwen3_decode(layer): operations.YieldOperation(), ], ) + + +# -------------------------------- Strategy for MiMoV2DecoderLayer --------------------------------------- + + +# TODO: unstable; current strategy matches DeepSeek for the common operations (MiMoV2 has no op_shared_experts), +# so we keep this redundant code here for convenience when adjusting the strategy +def _compute_moe_mimov2_layer_operations_strategy_tbo( + layer: torch.nn.Module, + forward_mode: ForwardMode, +) -> OperationsStrategy: + assert layer.is_layer_sparse, "MiMoV2DecoderLayer moe only support sparse layers" + if forward_mode == ForwardMode.EXTEND: + return _compute_moe_mimov2_prefill(layer) + elif ( + forward_mode == ForwardMode.DECODE or forward_mode == ForwardMode.TARGET_VERIFY + ): + return _compute_moe_mimov2_decode(layer) + else: + raise NotImplementedError(f"Unsupported {forward_mode=}") + + +def _compute_moe_mimov2_prefill(layer): + device_properties = torch.cuda.get_device_properties(device="cuda") + total_num_sms = device_properties.multi_processor_count + deep_gemm_num_sms = total_num_sms - DeepEPConfig.get_instance().num_sms + + return OperationsStrategy( + deep_gemm_num_sms=deep_gemm_num_sms, + tbo_delta_stages=0, + operations=[ + layer.op_comm_prepare_attn, + layer.self_attn.op_prepare, + layer.self_attn.op_core, + layer.op_comm_prepare_mlp, + layer.mlp.op_gate, + layer.mlp.op_select_experts, + layer.mlp.op_dispatch_a, + operations.YieldOperation(), + layer.mlp.op_dispatch_b, + layer.mlp.op_experts, + layer.mlp.op_combine_a, + operations.YieldOperation(), + layer.mlp.op_combine_b, + layer.mlp.op_output, + layer.op_comm_postprocess_layer, + ], + ) + + +def _compute_moe_mimov2_decode(layer): + return OperationsStrategy( + deep_gemm_num_sms=None, + tbo_delta_stages=2, + operations=[ + layer.op_comm_prepare_attn, + layer.self_attn.op_prepare, + operations.YieldOperation(), + layer.self_attn.op_core, + layer.op_comm_prepare_mlp, + layer.mlp.op_gate, + layer.mlp.op_select_experts, + operations.YieldOperation(), + layer.mlp.op_dispatch_a, + operations.YieldOperation(), + layer.mlp.op_dispatch_b, + layer.mlp.op_experts, + layer.mlp.op_combine_a, + operations.YieldOperation(), + layer.mlp.op_combine_b, + layer.mlp.op_output, + layer.op_comm_postprocess_layer, + operations.YieldOperation(), + ], + ) diff --git a/python/sglang/srt/models/mimo_v2_flash.py b/python/sglang/srt/models/mimo_v2_flash.py index ca3c903a1..d6f5eb07f 100644 --- a/python/sglang/srt/models/mimo_v2_flash.py +++ b/python/sglang/srt/models/mimo_v2_flash.py @@ -19,18 +19,21 @@ import torch import torch.nn.functional as F from torch import nn +from sglang.srt.batch_overlap.two_batch_overlap import model_forward_maybe_tbo from sglang.srt.distributed import ( get_moe_expert_parallel_world_size, get_pp_group, get_tensor_model_parallel_world_size, tensor_model_parallel_all_reduce, ) +from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.communicator import ( LayerCommunicator, LayerScatterModes, + ScatterMode, enable_moe_dense_fully_dp, ) from sglang.srt.layers.dp_attention import ( @@ -66,7 +69,12 @@ from sglang.srt.model_loader.weight_utils import ( kv_cache_scales_loader, ) from sglang.srt.server_args import get_global_server_args -from sglang.srt.utils import LazyValue, add_prefix, make_layers +from sglang.srt.utils import ( + LazyValue, + add_prefix, + is_non_idle_and_non_empty, + make_layers, +) MiMoV2FlashConfig = None @@ -322,6 +330,72 @@ class MiMoV2MoE(nn.Module): return final_hidden_states + def op_gate(self, state): + if is_non_idle_and_non_empty( + state.forward_batch.forward_mode, state.hidden_states_mlp_input + ): + # router_logits: (num_tokens, n_experts) + state.router_logits = self.gate(state.hidden_states_mlp_input) + else: + state.router_logits = None + + def op_select_experts(self, state): + router_logits = state.pop("router_logits") + hidden_states = state.hidden_states_mlp_input + if router_logits is not None: + with get_global_expert_distribution_recorder().with_current_layer( + self.layer_id + ): + state.topk_output = self.topk( + hidden_states=hidden_states, + router_logits=router_logits, + num_token_non_padded=state.forward_batch.num_token_non_padded, + expert_location_dispatch_info=ExpertLocationDispatchInfo.init_new( + layer_id=self.layer_id, + ), + ) + else: + state.topk_output = self.topk.empty_topk_output(hidden_states.device) + + def op_dispatch_a(self, state): + if self.ep_size > 1: + self.experts.dispatcher.dispatch_a( + hidden_states=state.pop("hidden_states_mlp_input"), + topk_output=state.pop("topk_output"), + tbo_subbatch_index=state.get("tbo_subbatch_index"), + ) + + def op_dispatch_b(self, state): + if self.ep_size > 1: + with get_global_expert_distribution_recorder().with_current_layer( + self.layer_id + ): + state.dispatch_output = self.experts.dispatcher.dispatch_b( + tbo_subbatch_index=state.get("tbo_subbatch_index"), + ) + + def op_experts(self, state): + state.combine_input = self.experts.run_moe_core( + dispatch_output=state.dispatch_output, + ) + + def op_combine_a(self, state): + if self.ep_size > 1: + self.experts.dispatcher.combine_a( + combine_input=state.pop("combine_input"), + tbo_subbatch_index=state.get("tbo_subbatch_index"), + ) + state.pop("dispatch_output") + + def op_combine_b(self, state): + if self.ep_size > 1: + state.hidden_states_after_combine = self.experts.dispatcher.combine_b( + tbo_subbatch_index=state.get("tbo_subbatch_index"), + ) + + def op_output(self, state): + state.hidden_states_mlp_output = state.pop("hidden_states_after_combine") + class MiMoV2Attention(nn.Module): def __init__( @@ -425,6 +499,47 @@ class MiMoV2Attention(nn.Module): else None ) + def op_prepare(self, state): + state.attn_intermediate_state = self.forward_prepare( + positions=state.positions, + hidden_states=state.pop("hidden_states_after_comm_pre_attn"), + forward_batch=state.forward_batch, + ) + + def op_core(self, state): + state.hidden_states_after_attn = self.forward_core( + state.pop("attn_intermediate_state") + ) + + def forward_prepare( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + ): + if hidden_states.shape[0] == 0: + return hidden_states, forward_batch, None + qkv, _ = self.qkv_proj(hidden_states) + q, k, v = qkv.split([self.q_size, self.k_size, self.v_size], dim=-1) + + q, k = self.rotary_emb(positions, q, k) + if self.v_scale is not None: + v = v * self.v_scale + + inner_state = q, k, v, forward_batch + return None, forward_batch, inner_state + + def forward_core(self, intermediate_state): + hidden_states, forward_batch, inner_state = intermediate_state + if inner_state is None: + return hidden_states + attn_output = self.attn( + *inner_state, + sinks=self.attention_sink_bias, + ) + output, _ = self.o_proj(attn_output) + return output + def forward( self, positions: torch.Tensor, @@ -609,6 +724,63 @@ class MiMoV2DecoderLayer(nn.Module): def is_swa_layer(self) -> bool: return self.config.hybrid_layer_pattern[self.layer_id] == 1 + def op_comm_prepare_attn( + self, + state, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + residual: Optional[torch.Tensor], + tbo_subbatch_index: Optional[int] = None, + ): + state.hidden_states_after_comm_pre_attn, state.residual_after_input_ln = ( + self.layer_communicator.prepare_attn(hidden_states, residual, forward_batch) + ) + state.update( + dict( + forward_batch=forward_batch, + positions=positions, + tbo_subbatch_index=tbo_subbatch_index, + ) + ) + + def op_comm_prepare_mlp(self, state): + state.hidden_states_mlp_input, state.residual_after_comm_pre_mlp = ( + self.layer_communicator.prepare_mlp( + state.pop("hidden_states_after_attn"), + state.pop("residual_after_input_ln"), + state.forward_batch, + ) + ) + + def op_mlp(self, state): + hidden_states = state.pop("hidden_states_mlp_input") + state.hidden_states_mlp_output = self.mlp(hidden_states, state.forward_batch) + + def op_comm_postprocess_layer(self, state): + hidden_states, residual = self.layer_communicator.postprocess_layer( + state.pop("hidden_states_mlp_output"), + state.pop("residual_after_comm_pre_mlp"), + state.forward_batch, + ) + + output = dict( + positions=state.positions, + hidden_states=hidden_states, + residual=residual, + forward_batch=state.forward_batch, + tbo_subbatch_index=state.tbo_subbatch_index, + ) + + state.clear( + expect_keys={ + "positions", + "forward_batch", + "tbo_subbatch_index", + } + ) + return output + class MiMoV2Model(nn.Module): def __init__( @@ -682,14 +854,42 @@ class MiMoV2Model(nn.Module): hidden_states = pp_proxy_tensors["hidden_states"] residual = pp_proxy_tensors["residual"] - for i in range(self.start_layer, self.end_layer): - layer = self.layers[i] - hidden_states, residual = layer( - positions, - hidden_states, - forward_batch, - residual, + if forward_batch.can_run_tbo: + tbo_start_layer = self.start_layer + tbo_end_layer = self.end_layer + + # skip first layer for TBO when starting from layer 0 + if self.start_layer == 0: + layer = self.layers[0] + hidden_states, residual = layer( + positions, hidden_states, forward_batch, residual + ) + tbo_start_layer = tbo_start_layer + 1 + + hidden_states, residual = model_forward_maybe_tbo( + layers=self.layers[tbo_start_layer:tbo_end_layer], + enable_tbo=True, + input_data_scatter_mode=( + ScatterMode.model_input_output() + if tbo_start_layer == self.start_layer + else self.layers[ + tbo_start_layer - 1 + ].layer_scatter_modes.layer_output_mode + ), + positions=positions, + forward_batch=forward_batch, + hidden_states=hidden_states, + residual=residual, ) + else: + for i in range(self.start_layer, self.end_layer): + layer = self.layers[i] + hidden_states, residual = layer( + positions, + hidden_states, + forward_batch, + residual, + ) hidden_states_before_norm = None if not self.pp_group.is_last_rank: