Support GPU execution time breakdown by forward mode metrics (#15396)
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user