[AMD][with CI Fix] support two batch overlapping for mori ep (#19216)
Co-authored-by: Duyi-Wang <duyi.wang@amd.com> Co-authored-by: kkHuang-amd <wunhuang@amd.com> Co-authored-by: Feiyue Zhai <feiyue.zhai@amd.com> Co-authored-by: HAI <hixiao@gmail.com>
This commit is contained in:
@@ -170,7 +170,7 @@ class _StateDict:
|
||||
def clear(self, expect_keys: Sequence[str]):
|
||||
if set(self._data.keys()) != set(expect_keys):
|
||||
raise Exception(
|
||||
f"Unexpected keys when clearning. This may indicate you do not release memory early enough but leave it to here. {list(self._data.keys())=} {expect_keys=}"
|
||||
f"Unexpected keys when clearing. This may indicate you do not release memory early enough but leave it until here. {list(self._data.keys())=} {expect_keys=}"
|
||||
)
|
||||
|
||||
self._data.clear()
|
||||
|
||||
@@ -7,6 +7,9 @@ from sglang.srt.batch_overlap import operations
|
||||
from sglang.srt.batch_overlap.operations import Operation
|
||||
from sglang.srt.layers.moe.token_dispatcher import DeepEPConfig
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.srt.utils import is_hip
|
||||
|
||||
_is_hip = is_hip()
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -91,7 +94,9 @@ def _compute_moe_deepseek_layer_operations_strategy_tbo(
|
||||
def _compute_moe_deepseek_blog_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
|
||||
deep_gemm_num_sms = None
|
||||
if not _is_hip:
|
||||
deep_gemm_num_sms = total_num_sms - DeepEPConfig.get_instance().num_sms
|
||||
|
||||
return OperationsStrategy(
|
||||
deep_gemm_num_sms=deep_gemm_num_sms,
|
||||
@@ -168,7 +173,9 @@ def _compute_moe_qwen3_layer_operations_strategy_tbo(
|
||||
def _compute_moe_qwen3_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
|
||||
deep_gemm_num_sms = None
|
||||
if not _is_hip:
|
||||
deep_gemm_num_sms = total_num_sms - DeepEPConfig.get_instance().num_sms
|
||||
|
||||
return OperationsStrategy(
|
||||
deep_gemm_num_sms=deep_gemm_num_sms,
|
||||
|
||||
@@ -30,6 +30,7 @@ from sglang.srt.layers.moe import (
|
||||
from sglang.srt.layers.moe.token_dispatcher import (
|
||||
DeepEPDispatcher,
|
||||
MooncakeEPDispatcher,
|
||||
MoriEPDispatcher,
|
||||
)
|
||||
from sglang.srt.layers.moe.token_dispatcher.base import BaseDispatcher
|
||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||
@@ -1027,6 +1028,10 @@ class MaybeTboDeepEPDispatcher(BaseDispatcher):
|
||||
self._inners = [
|
||||
MooncakeEPDispatcher(**kwargs) for _ in range(num_inner_dispatchers)
|
||||
]
|
||||
elif get_moe_a2a_backend().is_mori():
|
||||
self._inners = [
|
||||
MoriEPDispatcher(**kwargs) for _ in range(num_inner_dispatchers)
|
||||
]
|
||||
|
||||
def _execute(self, name, tbo_subbatch_index: Optional[int] = None, **kwargs):
|
||||
return getattr(self._inners[tbo_subbatch_index or 0], name)(**kwargs)
|
||||
|
||||
Reference in New Issue
Block a user