[MiMoV2Flash] [feat]: support two batch overlap (#17634)

This commit is contained in:
TZHelloWorld
2026-02-02 14:03:57 -08:00
committed by GitHub
parent b0a6d5244c
commit cbf1500390
2 changed files with 292 additions and 8 deletions
@@ -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(),
],
)