Stand up SchedulerProfilerManager; migrate profiler state to it (#25613)
This commit is contained in:
@@ -167,6 +167,9 @@ from sglang.srt.managers.schedule_policy import (
|
|||||||
from sglang.srt.managers.scheduler_components.dp_attn import (
|
from sglang.srt.managers.scheduler_components.dp_attn import (
|
||||||
SchedulerDPAttnAdapter,
|
SchedulerDPAttnAdapter,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.managers.scheduler_components.profiler_manager import (
|
||||||
|
SchedulerProfilerManager,
|
||||||
|
)
|
||||||
from sglang.srt.managers.scheduler_components.request_receiver import (
|
from sglang.srt.managers.scheduler_components.request_receiver import (
|
||||||
SchedulerRequestReceiver,
|
SchedulerRequestReceiver,
|
||||||
)
|
)
|
||||||
@@ -526,7 +529,11 @@ class Scheduler(
|
|||||||
self.init_watch_dog_memory_saver_input_blocker()
|
self.init_watch_dog_memory_saver_input_blocker()
|
||||||
|
|
||||||
# Init profiler
|
# Init profiler
|
||||||
self.init_profiler()
|
self.profiler_manager = SchedulerProfilerManager(
|
||||||
|
ps=self.ps,
|
||||||
|
dp_tp_cpu_group=self.dp_tp_cpu_group,
|
||||||
|
get_forward_ct=lambda: self.forward_ct,
|
||||||
|
)
|
||||||
|
|
||||||
# Init prefill-decodedisaggregation
|
# Init prefill-decodedisaggregation
|
||||||
self.init_disaggregation()
|
self.init_disaggregation()
|
||||||
@@ -1316,7 +1323,10 @@ class Scheduler(
|
|||||||
(ResumeMemoryOccupationReqInput, self.resume_memory_occupation),
|
(ResumeMemoryOccupationReqInput, self.resume_memory_occupation),
|
||||||
(CheckWeightsReqInput, self.check_weights),
|
(CheckWeightsReqInput, self.check_weights),
|
||||||
(SlowDownReqInput, self.slow_down),
|
(SlowDownReqInput, self.slow_down),
|
||||||
(ProfileReq, self.profile),
|
(
|
||||||
|
ProfileReq,
|
||||||
|
lambda req: self._profile(self.profiler_manager, req),
|
||||||
|
),
|
||||||
(FreezeGCReq, self.handle_freeze_gc),
|
(FreezeGCReq, self.handle_freeze_gc),
|
||||||
(GetInternalStateReq, self.get_internal_state),
|
(GetInternalStateReq, self.get_internal_state),
|
||||||
(SetInternalStateReq, self.set_internal_state),
|
(SetInternalStateReq, self.set_internal_state),
|
||||||
@@ -2715,7 +2725,7 @@ class Scheduler(
|
|||||||
batch.forward_iter = self.forward_ct
|
batch.forward_iter = self.forward_ct
|
||||||
|
|
||||||
# Whether to run the profiler
|
# Whether to run the profiler
|
||||||
self._profile_batch_predicate(batch)
|
self._profile_batch_predicate(self.profiler_manager, batch)
|
||||||
if self.forward_sleep_time is not None:
|
if self.forward_sleep_time is not None:
|
||||||
logger.info(f"Scheduler.run_batch sleep {self.forward_sleep_time}s")
|
logger.info(f"Scheduler.run_batch sleep {self.forward_sleep_time}s")
|
||||||
time.sleep(self.forward_sleep_time)
|
time.sleep(self.forward_sleep_time)
|
||||||
|
|||||||
@@ -0,0 +1,43 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Callable, List, Optional
|
||||||
|
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.srt.utils.profile_utils import ProfileManager
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(kw_only=True)
|
||||||
|
class SchedulerProfilerManager:
|
||||||
|
ps: Any
|
||||||
|
dp_tp_cpu_group: Any
|
||||||
|
get_forward_ct: Callable[[], int]
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
if envs.SGLANG_PROFILE_V2.get():
|
||||||
|
self._profile_manager = ProfileManager(
|
||||||
|
ps=self.ps,
|
||||||
|
cpu_group=self.dp_tp_cpu_group,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
self.torch_profiler = None
|
||||||
|
self.torch_profiler_output_dir: Optional[Path] = None
|
||||||
|
self.profiler_activities: Optional[List[str]] = None
|
||||||
|
self.profile_id: Optional[str] = None
|
||||||
|
|
||||||
|
self.profiler_start_forward_ct: Optional[int] = None
|
||||||
|
self.profiler_target_forward_ct: Optional[int] = None
|
||||||
|
|
||||||
|
self.profiler_prefill_ct: Optional[int] = None
|
||||||
|
self.profiler_decode_ct: Optional[int] = None
|
||||||
|
self.profiler_target_prefill_ct: Optional[int] = None
|
||||||
|
self.profiler_target_decode_ct: Optional[int] = None
|
||||||
|
|
||||||
|
self.profile_by_stage: bool = False
|
||||||
|
self.profile_in_progress: bool = False
|
||||||
|
self.merge_profiles = False
|
||||||
|
|
||||||
|
# For ROCM
|
||||||
|
self.rpd_profiler = None
|
||||||
@@ -14,11 +14,12 @@ from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
|||||||
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 is_npu
|
from sglang.srt.utils import is_npu
|
||||||
from sglang.srt.utils.profile_merger import ProfileMerger
|
from sglang.srt.utils.profile_merger import ProfileMerger
|
||||||
from sglang.srt.utils.profile_utils import ProfileManager
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||||
from sglang.srt.managers.scheduler import Scheduler
|
from sglang.srt.managers.scheduler_components.profiler_manager import (
|
||||||
|
SchedulerProfilerManager,
|
||||||
|
)
|
||||||
|
|
||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
if _is_npu:
|
if _is_npu:
|
||||||
@@ -35,36 +36,9 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
|
|
||||||
class SchedulerProfilerMixin:
|
class SchedulerProfilerMixin:
|
||||||
def init_profiler(self: Scheduler):
|
@staticmethod
|
||||||
if envs.SGLANG_PROFILE_V2.get():
|
def _init_profile(
|
||||||
self._profile_manager = ProfileManager(
|
self: "SchedulerProfilerManager",
|
||||||
ps=self.ps,
|
|
||||||
cpu_group=self.dp_tp_cpu_group,
|
|
||||||
)
|
|
||||||
return
|
|
||||||
|
|
||||||
self.torch_profiler = None
|
|
||||||
self.torch_profiler_output_dir: Optional[Path] = None
|
|
||||||
self.profiler_activities: Optional[List[str]] = None
|
|
||||||
self.profile_id: Optional[str] = None
|
|
||||||
|
|
||||||
self.profiler_start_forward_ct: Optional[int] = None
|
|
||||||
self.profiler_target_forward_ct: Optional[int] = None
|
|
||||||
|
|
||||||
self.profiler_prefill_ct: Optional[int] = None
|
|
||||||
self.profiler_decode_ct: Optional[int] = None
|
|
||||||
self.profiler_target_prefill_ct: Optional[int] = None
|
|
||||||
self.profiler_target_decode_ct: Optional[int] = None
|
|
||||||
|
|
||||||
self.profile_by_stage: bool = False
|
|
||||||
self.profile_in_progress: bool = False
|
|
||||||
self.merge_profiles = False
|
|
||||||
|
|
||||||
# For ROCM
|
|
||||||
self.rpd_profiler = None
|
|
||||||
|
|
||||||
def init_profile(
|
|
||||||
self: Scheduler,
|
|
||||||
output_dir: Optional[str],
|
output_dir: Optional[str],
|
||||||
start_step: Optional[int],
|
start_step: Optional[int],
|
||||||
num_steps: Optional[int],
|
num_steps: Optional[int],
|
||||||
@@ -114,7 +88,7 @@ class SchedulerProfilerMixin:
|
|||||||
self.profile_prefix = profile_prefix
|
self.profile_prefix = profile_prefix
|
||||||
|
|
||||||
if start_step:
|
if start_step:
|
||||||
self.profiler_start_forward_ct = max(start_step, self.forward_ct + 1)
|
self.profiler_start_forward_ct = max(start_step, self.get_forward_ct() + 1)
|
||||||
|
|
||||||
if num_steps:
|
if num_steps:
|
||||||
if self.profile_by_stage:
|
if self.profile_by_stage:
|
||||||
@@ -127,15 +101,16 @@ class SchedulerProfilerMixin:
|
|||||||
self.profiler_start_forward_ct + num_steps
|
self.profiler_start_forward_ct + num_steps
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.profiler_target_forward_ct = self.forward_ct + num_steps
|
self.profiler_target_forward_ct = self.get_forward_ct() + num_steps
|
||||||
# The caller will be notified when reaching profiler_target_forward_ct
|
# The caller will be notified when reaching profiler_target_forward_ct
|
||||||
else:
|
else:
|
||||||
self.profiler_target_forward_ct = None
|
self.profiler_target_forward_ct = None
|
||||||
|
|
||||||
return ProfileReqOutput(success=True, message="Succeeded")
|
return ProfileReqOutput(success=True, message="Succeeded")
|
||||||
|
|
||||||
def start_profile(
|
@staticmethod
|
||||||
self: Scheduler, stage: Optional[ForwardMode] = None
|
def _start_profile(
|
||||||
|
self: "SchedulerProfilerManager", stage: Optional[ForwardMode] = None
|
||||||
) -> ProfileReqOutput | None:
|
) -> ProfileReqOutput | None:
|
||||||
if envs.SGLANG_PROFILE_V2.get():
|
if envs.SGLANG_PROFILE_V2.get():
|
||||||
return self._profile_manager.manual_start()
|
return self._profile_manager.manual_start()
|
||||||
@@ -215,7 +190,8 @@ class SchedulerProfilerMixin:
|
|||||||
|
|
||||||
return ProfileReqOutput(success=True, message="Succeeded")
|
return ProfileReqOutput(success=True, message="Succeeded")
|
||||||
|
|
||||||
def _merge_profile_traces(self: Scheduler) -> str:
|
@staticmethod
|
||||||
|
def _merge_profile_traces(self: "SchedulerProfilerManager") -> str:
|
||||||
if not self.merge_profiles:
|
if not self.merge_profiles:
|
||||||
return ""
|
return ""
|
||||||
|
|
||||||
@@ -247,8 +223,9 @@ class SchedulerProfilerMixin:
|
|||||||
else:
|
else:
|
||||||
return merge_message
|
return merge_message
|
||||||
|
|
||||||
def stop_profile(
|
@staticmethod
|
||||||
self: Scheduler, stage: Optional[ForwardMode] = None
|
def _stop_profile(
|
||||||
|
self: "SchedulerProfilerManager", stage: Optional[ForwardMode] = None
|
||||||
) -> ProfileReqOutput | None:
|
) -> ProfileReqOutput | None:
|
||||||
if envs.SGLANG_PROFILE_V2.get():
|
if envs.SGLANG_PROFILE_V2.get():
|
||||||
return self._profile_manager.manual_stop()
|
return self._profile_manager.manual_stop()
|
||||||
@@ -322,7 +299,7 @@ class SchedulerProfilerMixin:
|
|||||||
if self.ps.gpu_id == get_global_server_args().base_gpu_id:
|
if self.ps.gpu_id == get_global_server_args().base_gpu_id:
|
||||||
torch.cuda.cudart().cudaProfilerStop()
|
torch.cuda.cudart().cudaProfilerStop()
|
||||||
|
|
||||||
merge_message = self._merge_profile_traces()
|
merge_message = SchedulerProfilerMixin._merge_profile_traces(self)
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"Profiling done. Traces are saved to: %s%s",
|
"Profiling done. Traces are saved to: %s%s",
|
||||||
@@ -335,7 +312,10 @@ class SchedulerProfilerMixin:
|
|||||||
|
|
||||||
return ProfileReqOutput(success=True, message=f"Succeeded.{merge_message}")
|
return ProfileReqOutput(success=True, message=f"Succeeded.{merge_message}")
|
||||||
|
|
||||||
def _profile_batch_predicate(self: Scheduler, batch: ScheduleBatch):
|
@staticmethod
|
||||||
|
def _profile_batch_predicate(
|
||||||
|
self: "SchedulerProfilerManager", batch: ScheduleBatch
|
||||||
|
):
|
||||||
if envs.SGLANG_PROFILE_V2.get():
|
if envs.SGLANG_PROFILE_V2.get():
|
||||||
self._profile_manager.step(forward_mode=batch.forward_mode)
|
self._profile_manager.step(forward_mode=batch.forward_mode)
|
||||||
return
|
return
|
||||||
@@ -343,21 +323,27 @@ class SchedulerProfilerMixin:
|
|||||||
if self.profile_by_stage:
|
if self.profile_by_stage:
|
||||||
if batch.forward_mode.is_prefill():
|
if batch.forward_mode.is_prefill():
|
||||||
if self.profiler_prefill_ct == 0:
|
if self.profiler_prefill_ct == 0:
|
||||||
self.start_profile(batch.forward_mode)
|
SchedulerProfilerMixin._start_profile(self, batch.forward_mode)
|
||||||
self.profiler_prefill_ct += 1
|
self.profiler_prefill_ct += 1
|
||||||
if self.profiler_prefill_ct > self.profiler_target_prefill_ct:
|
if self.profiler_prefill_ct > self.profiler_target_prefill_ct:
|
||||||
if self.profile_in_progress:
|
if self.profile_in_progress:
|
||||||
self.stop_profile(stage=ForwardMode.EXTEND)
|
SchedulerProfilerMixin._stop_profile(
|
||||||
|
self, stage=ForwardMode.EXTEND
|
||||||
|
)
|
||||||
elif batch.forward_mode.is_decode():
|
elif batch.forward_mode.is_decode():
|
||||||
if self.profiler_decode_ct == 0:
|
if self.profiler_decode_ct == 0:
|
||||||
if self.profile_in_progress:
|
if self.profile_in_progress:
|
||||||
# force trace flush
|
# force trace flush
|
||||||
self.stop_profile(stage=ForwardMode.EXTEND)
|
SchedulerProfilerMixin._stop_profile(
|
||||||
self.start_profile(batch.forward_mode)
|
self, stage=ForwardMode.EXTEND
|
||||||
|
)
|
||||||
|
SchedulerProfilerMixin._start_profile(self, batch.forward_mode)
|
||||||
self.profiler_decode_ct += 1
|
self.profiler_decode_ct += 1
|
||||||
if self.profiler_decode_ct > self.profiler_target_decode_ct:
|
if self.profiler_decode_ct > self.profiler_target_decode_ct:
|
||||||
if self.profile_in_progress:
|
if self.profile_in_progress:
|
||||||
self.stop_profile(stage=ForwardMode.DECODE)
|
SchedulerProfilerMixin._stop_profile(
|
||||||
|
self, stage=ForwardMode.DECODE
|
||||||
|
)
|
||||||
elif batch.forward_mode.is_idle():
|
elif batch.forward_mode.is_idle():
|
||||||
pass
|
pass
|
||||||
else:
|
else:
|
||||||
@@ -366,19 +352,21 @@ class SchedulerProfilerMixin:
|
|||||||
# Check profiler
|
# Check profiler
|
||||||
if (
|
if (
|
||||||
self.profiler_target_forward_ct
|
self.profiler_target_forward_ct
|
||||||
and self.profiler_target_forward_ct <= self.forward_ct
|
and self.profiler_target_forward_ct <= self.get_forward_ct()
|
||||||
):
|
):
|
||||||
self.stop_profile()
|
SchedulerProfilerMixin._stop_profile(self)
|
||||||
if (
|
if (
|
||||||
self.profiler_start_forward_ct
|
self.profiler_start_forward_ct
|
||||||
and self.profiler_start_forward_ct == self.forward_ct
|
and self.profiler_start_forward_ct == self.get_forward_ct()
|
||||||
):
|
):
|
||||||
self.start_profile()
|
SchedulerProfilerMixin._start_profile(self)
|
||||||
|
|
||||||
def profile(self: Scheduler, recv_req: ProfileReq):
|
@staticmethod
|
||||||
|
def _profile(self: "SchedulerProfilerManager", recv_req: ProfileReq):
|
||||||
if recv_req.type == ProfileReqType.START_PROFILE:
|
if recv_req.type == ProfileReqType.START_PROFILE:
|
||||||
if recv_req.profile_by_stage or recv_req.start_step:
|
if recv_req.profile_by_stage or recv_req.start_step:
|
||||||
return self.init_profile(
|
return SchedulerProfilerMixin._init_profile(
|
||||||
|
self,
|
||||||
recv_req.output_dir,
|
recv_req.output_dir,
|
||||||
recv_req.start_step,
|
recv_req.start_step,
|
||||||
recv_req.num_steps,
|
recv_req.num_steps,
|
||||||
@@ -392,7 +380,8 @@ class SchedulerProfilerMixin:
|
|||||||
recv_req.profile_stages,
|
recv_req.profile_stages,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.init_profile(
|
SchedulerProfilerMixin._init_profile(
|
||||||
|
self,
|
||||||
recv_req.output_dir,
|
recv_req.output_dir,
|
||||||
recv_req.start_step,
|
recv_req.start_step,
|
||||||
recv_req.num_steps,
|
recv_req.num_steps,
|
||||||
@@ -404,6 +393,6 @@ class SchedulerProfilerMixin:
|
|||||||
recv_req.merge_profiles,
|
recv_req.merge_profiles,
|
||||||
recv_req.profile_prefix,
|
recv_req.profile_prefix,
|
||||||
)
|
)
|
||||||
return self.start_profile()
|
return SchedulerProfilerMixin._start_profile(self)
|
||||||
else:
|
else:
|
||||||
return self.stop_profile()
|
return SchedulerProfilerMixin._stop_profile(self)
|
||||||
|
|||||||
@@ -231,7 +231,7 @@ class TestProfileMergerIntegration(unittest.TestCase):
|
|||||||
# Test SchedulerProfilerMixin
|
# Test SchedulerProfilerMixin
|
||||||
from sglang.srt.managers.scheduler_profiler_mixin import SchedulerProfilerMixin
|
from sglang.srt.managers.scheduler_profiler_mixin import SchedulerProfilerMixin
|
||||||
|
|
||||||
sig = inspect.signature(SchedulerProfilerMixin.init_profile)
|
sig = inspect.signature(SchedulerProfilerMixin._init_profile)
|
||||||
self.assertIn("merge_profiles", sig.parameters)
|
self.assertIn("merge_profiles", sig.parameters)
|
||||||
|
|
||||||
# Test CLI profiler
|
# Test CLI profiler
|
||||||
|
|||||||
Reference in New Issue
Block a user