From c5f4e20f2f5c3bc7f798f0dda3072c34ad192b1f Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Thu, 18 Dec 2025 22:18:17 +0800 Subject: [PATCH] Support GPU execution time breakdown by forward mode metrics (#15396) --- python/sglang/srt/environ.py | 3 ++ python/sglang/srt/managers/scheduler.py | 16 +++--- .../srt/managers/scheduler_metrics_mixin.py | 18 +++++++ python/sglang/srt/metrics/collector.py | 14 +++++ python/sglang/srt/utils/device_timer.py | 54 +++++++++++++++++++ 5 files changed, 98 insertions(+), 7 deletions(-) create mode 100644 python/sglang/srt/utils/device_timer.py diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 417b9d378..a843aa87d 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -365,6 +365,9 @@ class Envs: # Numa SGLANG_NUMA_BIND_V2 = EnvBool(True) + # Metrics + SGLANG_ENABLE_METRICS_DEVICE_TIMER = EnvBool(False) + # fmt: on diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 033c37e3f..983093d00 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2105,10 +2105,11 @@ class Scheduler( with self.forward_stream_ctx: self.forward_stream.wait_stream(self.default_stream) self.future_map.resolve_future(model_worker_batch) - batch_result = self.model_worker.forward_batch_generation( - model_worker_batch - # here pp is not compatible with overlap - ) + with self.record_forward_metrics(batch): + batch_result = self.model_worker.forward_batch_generation( + model_worker_batch + # here pp is not compatible with overlap + ) # FIXME(lsyin): maybe move this to forward_batch_generation batch_result.copy_done = self.device_module.Event() if batch_result.delay_sample_func is None: @@ -2144,9 +2145,10 @@ class Scheduler( if self.spec_algorithm.is_none() else {} ) - batch_result = self.model_worker.forward_batch_generation( - worker_batch_or_batch, **kwargs - ) + with self.record_forward_metrics(batch): + batch_result = self.model_worker.forward_batch_generation( + worker_batch_or_batch, **kwargs + ) future_indices_or_next_token_ids = batch_result.next_token_ids self.update_cache_from_scheduler(batch, batch_result) diff --git a/python/sglang/srt/managers/scheduler_metrics_mixin.py b/python/sglang/srt/managers/scheduler_metrics_mixin.py index 13a6c7159..1a0d51daa 100644 --- a/python/sglang/srt/managers/scheduler_metrics_mixin.py +++ b/python/sglang/srt/managers/scheduler_metrics_mixin.py @@ -3,6 +3,7 @@ from __future__ import annotations import logging import time from collections import defaultdict +from contextlib import contextmanager from typing import TYPE_CHECKING, List, Optional from sglang.srt.disaggregation.kv_events import EventPublisherFactory, KVEventBatch @@ -13,6 +14,7 @@ from sglang.srt.managers.schedule_policy import PrefillAdder from sglang.srt.managers.scheduler import Req, ScheduleBatch from sglang.srt.metrics.collector import SchedulerMetricsCollector, SchedulerStats from sglang.srt.utils import get_bool_env_var +from sglang.srt.utils.device_timer import DeviceTimer if TYPE_CHECKING: from sglang.srt.managers.scheduler import Scheduler @@ -21,6 +23,7 @@ logger = logging.getLogger(__name__) RECORD_STEP_TIME = get_bool_env_var("SGLANG_RECORD_STEP_TIME") LOG_FORWARD_ITERS = envs.SGLANG_LOG_FORWARD_ITERS.get() +ENABLE_METRICS_DEVICE_TIMER = envs.SGLANG_ENABLE_METRICS_DEVICE_TIMER.get() class KvMetrics: @@ -80,6 +83,11 @@ class SchedulerMetricsMixin: labels["dp_rank"] = dp_rank self.metrics_collector = SchedulerMetricsCollector(labels=labels) + if ENABLE_METRICS_DEVICE_TIMER: + self.forward_pass_device_timer = DeviceTimer( + reporter=self.metrics_collector.increment_gpu_execution_seconds + ) + if self.enable_kv_cache_events: self.init_kv_events(self.server_args.kv_events_config) @@ -455,3 +463,13 @@ class SchedulerMetricsMixin: num_waiting_reqs=num_waiting_reqs, num_tokens=num_tokens, ) + + @contextmanager + def record_forward_metrics(self: Scheduler, batch): + if not (self.enable_metrics and ENABLE_METRICS_DEVICE_TIMER): + yield + return + + category = "forward_" + batch.forward_mode.name.lower() + with self.forward_pass_device_timer.wrap(category=category): + yield diff --git a/python/sglang/srt/metrics/collector.py b/python/sglang/srt/metrics/collector.py index 765b256ed..65c856041 100644 --- a/python/sglang/srt/metrics/collector.py +++ b/python/sglang/srt/metrics/collector.py @@ -12,6 +12,7 @@ # limitations under the License. # ============================================================================== """Utilities for Prometheus Metrics Collection.""" +import logging import os import time from dataclasses import dataclass, field @@ -25,6 +26,9 @@ from sglang.srt.utils import get_bool_env_var SGLANG_TEST_REQUEST_TIME_STATS = get_bool_env_var("SGLANG_TEST_REQUEST_TIME_STATS") +logger = logging.getLogger(__name__) + + def get_histogram_conf_from_env(env_var_name: str) -> Optional[List[float]]: """ Get the histogram configuration from the environment variable. @@ -660,6 +664,12 @@ class SchedulerMetricsCollector: labelnames=labels.keys(), ) + self.gpu_execution_seconds_total = Counter( + name="sglang:gpu_execution_seconds_total", + documentation="Total time that GPU is busy executing a workload.", + labelnames=list(labels.keys()) + ["category"], + ) + def _log_gauge(self, gauge, data: Union[int, float]) -> None: # Convenience function for logging to gauge. gauge.labels(**self.labels).set(data) @@ -699,6 +709,10 @@ class SchedulerMetricsCollector: ) self.realtime_decode_tokens_total.labels(**self.labels).inc(decode_tokens) + def increment_gpu_execution_seconds(self, category: str, t: float): + logger.debug(f"GPU execution seconds: {category=} {t=:.3f}") + self.gpu_execution_seconds_total.labels(**self.labels, category=category).inc(t) + def log_stats(self, stats: SchedulerStats) -> None: self._log_gauge(self.num_running_reqs, stats.num_running_reqs) self._log_gauge(self.num_used_tokens, stats.num_used_tokens) diff --git a/python/sglang/srt/utils/device_timer.py b/python/sglang/srt/utils/device_timer.py new file mode 100644 index 000000000..e426918bb --- /dev/null +++ b/python/sglang/srt/utils/device_timer.py @@ -0,0 +1,54 @@ +from collections import deque +from contextlib import contextmanager +from dataclasses import dataclass +from typing import Callable, Deque, Optional + +import torch + + +class DeviceTimer: + def __init__(self, reporter: Callable[[str, float], None]): + self._intervals: Deque[_TimingInterval] = deque() + self._reporter = reporter + + @contextmanager + def wrap(self, category: str): + self._intervals.append(_TimingInterval.create()) + try: + yield + finally: + self._intervals[-1].end(category=category) + self._report() + + def _report(self): + while len(self._intervals) > 0: + interval = self._intervals[0] + if not interval.end_event.query(): + break + + self._intervals.popleft() + self._reporter(interval.category, interval.elapsed_time() / 1000.0) + + +@dataclass +class _TimingInterval: + start_event: torch.cuda.Event + end_event: Optional[torch.cuda.Event] = None + category: Optional[str] = None + + @staticmethod + def create(): + start_event = torch.cuda.Event(enable_timing=True) + start_event.record() + return _TimingInterval(start_event=start_event) + + def end(self, category: str): + end_event = torch.cuda.Event(enable_timing=True) + end_event.record() + + assert self.end_event is None + self.end_event = end_event + self.category = category + + def elapsed_time(self) -> float: + return self.start_event.elapsed_time(self.end_event)