[AMD] [MORI-EP] Skip LocalExpertCount kernel in decode graph when not recording (#30302)
This commit is contained in:
@@ -5,7 +5,10 @@ import os
|
|||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import TYPE_CHECKING, List, NamedTuple, Optional, Tuple
|
from typing import TYPE_CHECKING, List, NamedTuple, Optional, Tuple
|
||||||
|
|
||||||
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
|
from sglang.srt.eplb.expert_distribution import (
|
||||||
|
_ExpertDistributionRecorderNoop,
|
||||||
|
get_global_expert_distribution_recorder,
|
||||||
|
)
|
||||||
from sglang.srt.layers.dp_attention import get_is_extend_in_batch
|
from sglang.srt.layers.dp_attention import get_is_extend_in_batch
|
||||||
from sglang.srt.layers.moe.token_dispatcher.base import (
|
from sglang.srt.layers.moe.token_dispatcher.base import (
|
||||||
BaseDispatcher,
|
BaseDispatcher,
|
||||||
@@ -53,7 +56,15 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
def _should_record_expert_distribution() -> bool:
|
def _should_record_expert_distribution() -> bool:
|
||||||
recorder = get_global_expert_distribution_recorder()
|
recorder = get_global_expert_distribution_recorder()
|
||||||
return recorder.recording or torch.get_device_module().is_current_stream_capturing()
|
if recorder.recording:
|
||||||
|
return True
|
||||||
|
# While capturing, only bake in the count kernel if a recorder is actually
|
||||||
|
# configured (non-Noop); otherwise it would replay as dead work every decode
|
||||||
|
# step. Configured recorders still bake it in, so start_record() works after
|
||||||
|
# capture.
|
||||||
|
if torch.get_device_module().is_current_stream_capturing():
|
||||||
|
return not isinstance(recorder, _ExpertDistributionRecorderNoop)
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
class MoriEPPDispatchHooks(DeepEPPDispatchHooks):
|
class MoriEPPDispatchHooks(DeepEPPDispatchHooks):
|
||||||
|
|||||||
Reference in New Issue
Block a user