From a3fd5c24becc465a01d28a8199c08a36e693bf08 Mon Sep 17 00:00:00 2001 From: JoyFuture <35593546+JoyFuture@users.noreply.github.com> Date: Sun, 14 Jun 2026 08:16:53 +0800 Subject: [PATCH] feat: add NVTX markers for the scheduler main loop (#27901) --- python/sglang/srt/disaggregation/decode.py | 2 + python/sglang/srt/disaggregation/prefill.py | 2 + python/sglang/srt/environ.py | 1 + python/sglang/srt/managers/scheduler.py | 5 + .../scheduler_components/request_receiver.py | 2 + python/sglang/srt/utils/nvtx_utils.py | 98 +++++++++++++++++++ 6 files changed, 110 insertions(+) create mode 100644 python/sglang/srt/utils/nvtx_utils.py diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 247ba6673..208d0976e 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -85,6 +85,7 @@ from sglang.srt.observability.req_time_stats import ( ) from sglang.srt.utils import get_num_new_pages from sglang.srt.utils.network import NetworkAddress +from sglang.srt.utils.nvtx_utils import nvtx_annotated_method from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter logger = logging.getLogger(__name__) @@ -1824,6 +1825,7 @@ class SchedulerDisaggregationDecodeMixin: return GenerationBatchResult() + @nvtx_annotated_method("scheduler.get_next_batch_to_run") def get_next_disagg_decode_batch_to_run( self: Scheduler, ) -> Optional[ScheduleBatch]: diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 89ee2e2d8..eee6fd133 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -61,6 +61,7 @@ from sglang.srt.mem_cache.common import ( ) from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.observability.req_time_stats import set_schedule_time_batch +from sglang.srt.utils.nvtx_utils import nvtx_annotated_method if TYPE_CHECKING: from torch.distributed import ProcessGroup @@ -405,6 +406,7 @@ class SchedulerDisaggregationPrefillMixin: if room is not None and room in kv_mgr.transfer_infos: prefetch(room) + @nvtx_annotated_method("scheduler.get_next_batch_to_run") def get_next_disagg_prefill_batch_to_run( self: Scheduler, ) -> Optional[ScheduleBatch]: diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 2f3a15aae..f7bf43b19 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -243,6 +243,7 @@ class Envs: SGLANG_PROFILE_WITH_STACK = EnvBool(True) SGLANG_PROFILE_RECORD_SHAPES = EnvBool(True) SGLANG_PROFILE_V2 = EnvBool(False) + SGLANG_ENABLE_NVTX = EnvBool(False) SGLANG_RECORD_STEP_TIME = EnvBool(False) SGLANG_FORCE_SHUTDOWN = EnvBool(False) SGLANG_DEBUG_MEMORY_POOL = EnvBool(False) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 24f2ead54..98ba64f2d 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -263,6 +263,7 @@ from sglang.srt.utils.hf_transformers_utils import ( get_tokenizer_from_processor, ) from sglang.srt.utils.numa_utils import get_numa_node_if_available, numa_bind_to_node +from sglang.srt.utils.nvtx_utils import nvtx_annotated_method from sglang.srt.utils.tensor_bridge import use_mlx from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter from sglang.utils import TypeBasedDispatcher, get_exception_traceback @@ -1585,6 +1586,7 @@ class Scheduler( return disable_overlap_for_batch or need_grammar_sync + @nvtx_annotated_method("scheduler.process_input_requests") def process_input_requests(self, recv_reqs: List): now = time.monotonic() self.session_controller.maybe_reap(now) @@ -2461,6 +2463,7 @@ class Scheduler( # todo hisparse, maybe other info to contain for the new batch return batch + @nvtx_annotated_method("scheduler.get_next_batch_to_run") def get_next_batch_to_run(self) -> Optional[ScheduleBatch]: if self.enable_fpm: self._fpm_batch_t0 = time.monotonic() @@ -3048,6 +3051,7 @@ class Scheduler( else: batch.sampling_info = sched_sampling_info + @nvtx_annotated_method("scheduler.run_batch") def run_batch( self, batch: ScheduleBatch, @@ -3269,6 +3273,7 @@ class Scheduler( if batch_result.logits_output is not None: batch_result.logits_output.next_token_logits = None + @nvtx_annotated_method("scheduler.process_batch_result") def process_batch_result( self, batch: ScheduleBatch, diff --git a/python/sglang/srt/managers/scheduler_components/request_receiver.py b/python/sglang/srt/managers/scheduler_components/request_receiver.py index 3f10a2790..1d06d0750 100644 --- a/python/sglang/srt/managers/scheduler_components/request_receiver.py +++ b/python/sglang/srt/managers/scheduler_components/request_receiver.py @@ -29,6 +29,7 @@ from sglang.srt.utils import ( broadcast_pyobj, point_to_point_pyobj, ) +from sglang.srt.utils.nvtx_utils import nvtx_annotated_method if TYPE_CHECKING: from sglang.srt.configs.model_config import ModelConfig @@ -67,6 +68,7 @@ class SchedulerRequestReceiver: return False return num_recv_reqs >= self.max_recv_per_poll + @nvtx_annotated_method("scheduler.recv_requests") def recv_requests( self, ) -> List[Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput, Any]]: diff --git a/python/sglang/srt/utils/nvtx_utils.py b/python/sglang/srt/utils/nvtx_utils.py new file mode 100644 index 000000000..834180387 --- /dev/null +++ b/python/sglang/srt/utils/nvtx_utils.py @@ -0,0 +1,98 @@ +# Copyright 2023-2024 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Lightweight NVTX annotations for the scheduler main loop. + +Enabled via the ``SGLANG_ENABLE_NVTX`` environment variable (off by default). +When disabled, the decorator/context manager add zero runtime overhead so they +are safe to leave on hot scheduler paths. +""" + +import logging +from contextlib import contextmanager, nullcontext +from functools import wraps + +import torch + +from sglang.srt.environ import envs + +logger = logging.getLogger(__name__) + +_NVTX_ENV_ENABLED = envs.SGLANG_ENABLE_NVTX.get() +_nvtx_module = None +if _NVTX_ENV_ENABLED: + try: + import nvtx as _nvtx_module # type: ignore + except ImportError: + logger.warning( + "SGLANG_ENABLE_NVTX is set, but the `nvtx` package is missing. " + "NVTX annotations are disabled." + ) + +NVTX_ENABLED = _nvtx_module is not None + +# Colors are assigned per scheduler main-loop stage so the markers are easy to +# distinguish in Nsight Systems. +_NVTX_COLOR_MAP = { + # === Scheduler main loop (pipeline order) === + "scheduler.recv_requests": "blue", + "scheduler.process_input_requests": "purple", + "scheduler.get_next_batch_to_run": "green", + "scheduler.run_batch": "red", + "scheduler.process_batch_result": "cyan", +} + + +@contextmanager +def _nvtx_range_enabled(debug_name: str): + color = _NVTX_COLOR_MAP.get(debug_name) + # record_function carries a non-trivial (~microseconds) cost per call even + # when no PyTorch profiler is collecting, so only pay it when one is active + # (e.g. Chrome-trace export). For Nsight-only runs the nvtx.annotate marker + # alone is enough. + if torch.autograd._profiler_enabled(): + with torch.autograd.profiler.record_function(debug_name): + with _nvtx_module.annotate(debug_name, color=color): + yield + else: + with _nvtx_module.annotate(debug_name, color=color): + yield + + +if NVTX_ENABLED: + nvtx_range = _nvtx_range_enabled +else: + # When NVTX is disabled, hand back a shared no-op context manager so hot + # paths using `with nvtx_range(...)` pay no per-call generator overhead. + _NULL_CONTEXT = nullcontext() + + def nvtx_range(debug_name: str): + return _NULL_CONTEXT + + +def nvtx_annotated_method(debug_name: str): + # Decide at decoration time. When NVTX is disabled this returns the + # original function untouched, so decorated methods on hot paths have zero + # runtime cost. + if not NVTX_ENABLED: + return lambda func: func + + def decorator(func): + @wraps(func) + def wrapper(*args, **kwargs): + with nvtx_range(debug_name): + return func(*args, **kwargs) + + return wrapper + + return decorator