Support GPU execution time breakdown by forward mode metrics (#15396)
This commit is contained in:
@@ -365,6 +365,9 @@ class Envs:
|
|||||||
# Numa
|
# Numa
|
||||||
SGLANG_NUMA_BIND_V2 = EnvBool(True)
|
SGLANG_NUMA_BIND_V2 = EnvBool(True)
|
||||||
|
|
||||||
|
# Metrics
|
||||||
|
SGLANG_ENABLE_METRICS_DEVICE_TIMER = EnvBool(False)
|
||||||
|
|
||||||
# fmt: on
|
# fmt: on
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -2105,10 +2105,11 @@ class Scheduler(
|
|||||||
with self.forward_stream_ctx:
|
with self.forward_stream_ctx:
|
||||||
self.forward_stream.wait_stream(self.default_stream)
|
self.forward_stream.wait_stream(self.default_stream)
|
||||||
self.future_map.resolve_future(model_worker_batch)
|
self.future_map.resolve_future(model_worker_batch)
|
||||||
batch_result = self.model_worker.forward_batch_generation(
|
with self.record_forward_metrics(batch):
|
||||||
model_worker_batch
|
batch_result = self.model_worker.forward_batch_generation(
|
||||||
# here pp is not compatible with overlap
|
model_worker_batch
|
||||||
)
|
# here pp is not compatible with overlap
|
||||||
|
)
|
||||||
# FIXME(lsyin): maybe move this to forward_batch_generation
|
# FIXME(lsyin): maybe move this to forward_batch_generation
|
||||||
batch_result.copy_done = self.device_module.Event()
|
batch_result.copy_done = self.device_module.Event()
|
||||||
if batch_result.delay_sample_func is None:
|
if batch_result.delay_sample_func is None:
|
||||||
@@ -2144,9 +2145,10 @@ class Scheduler(
|
|||||||
if self.spec_algorithm.is_none()
|
if self.spec_algorithm.is_none()
|
||||||
else {}
|
else {}
|
||||||
)
|
)
|
||||||
batch_result = self.model_worker.forward_batch_generation(
|
with self.record_forward_metrics(batch):
|
||||||
worker_batch_or_batch, **kwargs
|
batch_result = self.model_worker.forward_batch_generation(
|
||||||
)
|
worker_batch_or_batch, **kwargs
|
||||||
|
)
|
||||||
future_indices_or_next_token_ids = batch_result.next_token_ids
|
future_indices_or_next_token_ids = batch_result.next_token_ids
|
||||||
self.update_cache_from_scheduler(batch, batch_result)
|
self.update_cache_from_scheduler(batch, batch_result)
|
||||||
|
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ from __future__ import annotations
|
|||||||
import logging
|
import logging
|
||||||
import time
|
import time
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
|
from contextlib import contextmanager
|
||||||
from typing import TYPE_CHECKING, List, Optional
|
from typing import TYPE_CHECKING, List, Optional
|
||||||
|
|
||||||
from sglang.srt.disaggregation.kv_events import EventPublisherFactory, KVEventBatch
|
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.managers.scheduler import Req, ScheduleBatch
|
||||||
from sglang.srt.metrics.collector import SchedulerMetricsCollector, SchedulerStats
|
from sglang.srt.metrics.collector import SchedulerMetricsCollector, SchedulerStats
|
||||||
from sglang.srt.utils import get_bool_env_var
|
from sglang.srt.utils import get_bool_env_var
|
||||||
|
from sglang.srt.utils.device_timer import DeviceTimer
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.scheduler import Scheduler
|
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")
|
RECORD_STEP_TIME = get_bool_env_var("SGLANG_RECORD_STEP_TIME")
|
||||||
LOG_FORWARD_ITERS = envs.SGLANG_LOG_FORWARD_ITERS.get()
|
LOG_FORWARD_ITERS = envs.SGLANG_LOG_FORWARD_ITERS.get()
|
||||||
|
ENABLE_METRICS_DEVICE_TIMER = envs.SGLANG_ENABLE_METRICS_DEVICE_TIMER.get()
|
||||||
|
|
||||||
|
|
||||||
class KvMetrics:
|
class KvMetrics:
|
||||||
@@ -80,6 +83,11 @@ class SchedulerMetricsMixin:
|
|||||||
labels["dp_rank"] = dp_rank
|
labels["dp_rank"] = dp_rank
|
||||||
self.metrics_collector = SchedulerMetricsCollector(labels=labels)
|
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:
|
if self.enable_kv_cache_events:
|
||||||
self.init_kv_events(self.server_args.kv_events_config)
|
self.init_kv_events(self.server_args.kv_events_config)
|
||||||
|
|
||||||
@@ -455,3 +463,13 @@ class SchedulerMetricsMixin:
|
|||||||
num_waiting_reqs=num_waiting_reqs,
|
num_waiting_reqs=num_waiting_reqs,
|
||||||
num_tokens=num_tokens,
|
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
|
||||||
|
|||||||
@@ -12,6 +12,7 @@
|
|||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
# ==============================================================================
|
# ==============================================================================
|
||||||
"""Utilities for Prometheus Metrics Collection."""
|
"""Utilities for Prometheus Metrics Collection."""
|
||||||
|
import logging
|
||||||
import os
|
import os
|
||||||
import time
|
import time
|
||||||
from dataclasses import dataclass, field
|
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")
|
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]]:
|
def get_histogram_conf_from_env(env_var_name: str) -> Optional[List[float]]:
|
||||||
"""
|
"""
|
||||||
Get the histogram configuration from the environment variable.
|
Get the histogram configuration from the environment variable.
|
||||||
@@ -660,6 +664,12 @@ class SchedulerMetricsCollector:
|
|||||||
labelnames=labels.keys(),
|
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:
|
def _log_gauge(self, gauge, data: Union[int, float]) -> None:
|
||||||
# Convenience function for logging to gauge.
|
# Convenience function for logging to gauge.
|
||||||
gauge.labels(**self.labels).set(data)
|
gauge.labels(**self.labels).set(data)
|
||||||
@@ -699,6 +709,10 @@ class SchedulerMetricsCollector:
|
|||||||
)
|
)
|
||||||
self.realtime_decode_tokens_total.labels(**self.labels).inc(decode_tokens)
|
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:
|
def log_stats(self, stats: SchedulerStats) -> None:
|
||||||
self._log_gauge(self.num_running_reqs, stats.num_running_reqs)
|
self._log_gauge(self.num_running_reqs, stats.num_running_reqs)
|
||||||
self._log_gauge(self.num_used_tokens, stats.num_used_tokens)
|
self._log_gauge(self.num_used_tokens, stats.num_used_tokens)
|
||||||
|
|||||||
@@ -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)
|
||||||
Reference in New Issue
Block a user