From b0a511560cbd2f020387fad95ea97a82a21416cd Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Mon, 18 May 2026 18:32:46 +0800 Subject: [PATCH] Stand up SchedulerProfilerManager; migrate profiler state to it (#25613) --- python/sglang/srt/managers/scheduler.py | 16 ++- .../scheduler_components/profiler_manager.py | 43 ++++++++ .../srt/managers/scheduler_profiler_mixin.py | 99 +++++++++---------- .../unit/utils/test_profile_merger.py | 2 +- 4 files changed, 101 insertions(+), 59 deletions(-) create mode 100644 python/sglang/srt/managers/scheduler_components/profiler_manager.py diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 7c429020c..0b1ab7b99 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -167,6 +167,9 @@ from sglang.srt.managers.schedule_policy import ( from sglang.srt.managers.scheduler_components.dp_attn import ( SchedulerDPAttnAdapter, ) +from sglang.srt.managers.scheduler_components.profiler_manager import ( + SchedulerProfilerManager, +) from sglang.srt.managers.scheduler_components.request_receiver import ( SchedulerRequestReceiver, ) @@ -526,7 +529,11 @@ class Scheduler( self.init_watch_dog_memory_saver_input_blocker() # 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 self.init_disaggregation() @@ -1316,7 +1323,10 @@ class Scheduler( (ResumeMemoryOccupationReqInput, self.resume_memory_occupation), (CheckWeightsReqInput, self.check_weights), (SlowDownReqInput, self.slow_down), - (ProfileReq, self.profile), + ( + ProfileReq, + lambda req: self._profile(self.profiler_manager, req), + ), (FreezeGCReq, self.handle_freeze_gc), (GetInternalStateReq, self.get_internal_state), (SetInternalStateReq, self.set_internal_state), @@ -2715,7 +2725,7 @@ class Scheduler( batch.forward_iter = self.forward_ct # 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: logger.info(f"Scheduler.run_batch sleep {self.forward_sleep_time}s") time.sleep(self.forward_sleep_time) diff --git a/python/sglang/srt/managers/scheduler_components/profiler_manager.py b/python/sglang/srt/managers/scheduler_components/profiler_manager.py new file mode 100644 index 000000000..e3d21d3f7 --- /dev/null +++ b/python/sglang/srt/managers/scheduler_components/profiler_manager.py @@ -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 diff --git a/python/sglang/srt/managers/scheduler_profiler_mixin.py b/python/sglang/srt/managers/scheduler_profiler_mixin.py index 66011de44..c1c714ab2 100644 --- a/python/sglang/srt/managers/scheduler_profiler_mixin.py +++ b/python/sglang/srt/managers/scheduler_profiler_mixin.py @@ -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.utils import is_npu from sglang.srt.utils.profile_merger import ProfileMerger -from sglang.srt.utils.profile_utils import ProfileManager if TYPE_CHECKING: 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() if _is_npu: @@ -35,36 +36,9 @@ logger = logging.getLogger(__name__) class SchedulerProfilerMixin: - def init_profiler(self: Scheduler): - 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 - - def init_profile( - self: Scheduler, + @staticmethod + def _init_profile( + self: "SchedulerProfilerManager", output_dir: Optional[str], start_step: Optional[int], num_steps: Optional[int], @@ -114,7 +88,7 @@ class SchedulerProfilerMixin: self.profile_prefix = profile_prefix 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 self.profile_by_stage: @@ -127,15 +101,16 @@ class SchedulerProfilerMixin: self.profiler_start_forward_ct + num_steps ) 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 else: self.profiler_target_forward_ct = None return ProfileReqOutput(success=True, message="Succeeded") - def start_profile( - self: Scheduler, stage: Optional[ForwardMode] = None + @staticmethod + def _start_profile( + self: "SchedulerProfilerManager", stage: Optional[ForwardMode] = None ) -> ProfileReqOutput | None: if envs.SGLANG_PROFILE_V2.get(): return self._profile_manager.manual_start() @@ -215,7 +190,8 @@ class SchedulerProfilerMixin: 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: return "" @@ -247,8 +223,9 @@ class SchedulerProfilerMixin: else: return merge_message - def stop_profile( - self: Scheduler, stage: Optional[ForwardMode] = None + @staticmethod + def _stop_profile( + self: "SchedulerProfilerManager", stage: Optional[ForwardMode] = None ) -> ProfileReqOutput | None: if envs.SGLANG_PROFILE_V2.get(): return self._profile_manager.manual_stop() @@ -322,7 +299,7 @@ class SchedulerProfilerMixin: if self.ps.gpu_id == get_global_server_args().base_gpu_id: torch.cuda.cudart().cudaProfilerStop() - merge_message = self._merge_profile_traces() + merge_message = SchedulerProfilerMixin._merge_profile_traces(self) logger.info( "Profiling done. Traces are saved to: %s%s", @@ -335,7 +312,10 @@ class SchedulerProfilerMixin: 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(): self._profile_manager.step(forward_mode=batch.forward_mode) return @@ -343,21 +323,27 @@ class SchedulerProfilerMixin: if self.profile_by_stage: if batch.forward_mode.is_prefill(): if self.profiler_prefill_ct == 0: - self.start_profile(batch.forward_mode) + SchedulerProfilerMixin._start_profile(self, batch.forward_mode) self.profiler_prefill_ct += 1 if self.profiler_prefill_ct > self.profiler_target_prefill_ct: if self.profile_in_progress: - self.stop_profile(stage=ForwardMode.EXTEND) + SchedulerProfilerMixin._stop_profile( + self, stage=ForwardMode.EXTEND + ) elif batch.forward_mode.is_decode(): if self.profiler_decode_ct == 0: if self.profile_in_progress: # force trace flush - self.stop_profile(stage=ForwardMode.EXTEND) - self.start_profile(batch.forward_mode) + SchedulerProfilerMixin._stop_profile( + self, stage=ForwardMode.EXTEND + ) + SchedulerProfilerMixin._start_profile(self, batch.forward_mode) self.profiler_decode_ct += 1 if self.profiler_decode_ct > self.profiler_target_decode_ct: if self.profile_in_progress: - self.stop_profile(stage=ForwardMode.DECODE) + SchedulerProfilerMixin._stop_profile( + self, stage=ForwardMode.DECODE + ) elif batch.forward_mode.is_idle(): pass else: @@ -366,19 +352,21 @@ class SchedulerProfilerMixin: # Check profiler if ( 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 ( 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.profile_by_stage or recv_req.start_step: - return self.init_profile( + return SchedulerProfilerMixin._init_profile( + self, recv_req.output_dir, recv_req.start_step, recv_req.num_steps, @@ -392,7 +380,8 @@ class SchedulerProfilerMixin: recv_req.profile_stages, ) else: - self.init_profile( + SchedulerProfilerMixin._init_profile( + self, recv_req.output_dir, recv_req.start_step, recv_req.num_steps, @@ -404,6 +393,6 @@ class SchedulerProfilerMixin: recv_req.merge_profiles, recv_req.profile_prefix, ) - return self.start_profile() + return SchedulerProfilerMixin._start_profile(self) else: - return self.stop_profile() + return SchedulerProfilerMixin._stop_profile(self) diff --git a/test/registered/unit/utils/test_profile_merger.py b/test/registered/unit/utils/test_profile_merger.py index 879afb514..9f211dbc0 100644 --- a/test/registered/unit/utils/test_profile_merger.py +++ b/test/registered/unit/utils/test_profile_merger.py @@ -231,7 +231,7 @@ class TestProfileMergerIntegration(unittest.TestCase): # Test 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) # Test CLI profiler