[MiMoV2Flash] [feat]: support two batch overlap (#17634)
This commit is contained in:
@@ -51,6 +51,15 @@ class OperationsStrategy:
|
|||||||
for layer in layers
|
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:
|
else:
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
@@ -209,3 +218,78 @@ def _compute_moe_qwen3_decode(layer):
|
|||||||
operations.YieldOperation(),
|
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
|
import torch.nn.functional as F
|
||||||
from torch import nn
|
from torch import nn
|
||||||
|
|
||||||
|
from sglang.srt.batch_overlap.two_batch_overlap import model_forward_maybe_tbo
|
||||||
from sglang.srt.distributed import (
|
from sglang.srt.distributed import (
|
||||||
get_moe_expert_parallel_world_size,
|
get_moe_expert_parallel_world_size,
|
||||||
get_pp_group,
|
get_pp_group,
|
||||||
get_tensor_model_parallel_world_size,
|
get_tensor_model_parallel_world_size,
|
||||||
tensor_model_parallel_all_reduce,
|
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 import ModelConfigForExpertLocation
|
||||||
from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo
|
from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo
|
||||||
from sglang.srt.layers.activation import SiluAndMul
|
from sglang.srt.layers.activation import SiluAndMul
|
||||||
from sglang.srt.layers.communicator import (
|
from sglang.srt.layers.communicator import (
|
||||||
LayerCommunicator,
|
LayerCommunicator,
|
||||||
LayerScatterModes,
|
LayerScatterModes,
|
||||||
|
ScatterMode,
|
||||||
enable_moe_dense_fully_dp,
|
enable_moe_dense_fully_dp,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
@@ -66,7 +69,12 @@ from sglang.srt.model_loader.weight_utils import (
|
|||||||
kv_cache_scales_loader,
|
kv_cache_scales_loader,
|
||||||
)
|
)
|
||||||
from sglang.srt.server_args import get_global_server_args
|
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
|
MiMoV2FlashConfig = None
|
||||||
|
|
||||||
@@ -322,6 +330,72 @@ class MiMoV2MoE(nn.Module):
|
|||||||
|
|
||||||
return final_hidden_states
|
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):
|
class MiMoV2Attention(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -425,6 +499,47 @@ class MiMoV2Attention(nn.Module):
|
|||||||
else None
|
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(
|
def forward(
|
||||||
self,
|
self,
|
||||||
positions: torch.Tensor,
|
positions: torch.Tensor,
|
||||||
@@ -609,6 +724,63 @@ class MiMoV2DecoderLayer(nn.Module):
|
|||||||
def is_swa_layer(self) -> bool:
|
def is_swa_layer(self) -> bool:
|
||||||
return self.config.hybrid_layer_pattern[self.layer_id] == 1
|
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):
|
class MiMoV2Model(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -682,14 +854,42 @@ class MiMoV2Model(nn.Module):
|
|||||||
hidden_states = pp_proxy_tensors["hidden_states"]
|
hidden_states = pp_proxy_tensors["hidden_states"]
|
||||||
residual = pp_proxy_tensors["residual"]
|
residual = pp_proxy_tensors["residual"]
|
||||||
|
|
||||||
for i in range(self.start_layer, self.end_layer):
|
if forward_batch.can_run_tbo:
|
||||||
layer = self.layers[i]
|
tbo_start_layer = self.start_layer
|
||||||
hidden_states, residual = layer(
|
tbo_end_layer = self.end_layer
|
||||||
positions,
|
|
||||||
hidden_states,
|
# skip first layer for TBO when starting from layer 0
|
||||||
forward_batch,
|
if self.start_layer == 0:
|
||||||
residual,
|
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
|
hidden_states_before_norm = None
|
||||||
if not self.pp_group.is_last_rank:
|
if not self.pp_group.is_last_rank:
|
||||||
|
|||||||
Reference in New Issue
Block a user