[MiMoV2Flash] [feat]: support two batch overlap (#17634)
This commit is contained in:
@@ -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(),
|
||||
],
|
||||
)
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user