Move metrics reporting to SchedulerMetricsReporter and retire metrics mixin (#25630)

This commit is contained in:
fzyzcjy
2026-05-18 18:42:08 +08:00
committed by GitHub
parent 780d969699
commit fd97fbb096
8 changed files with 1001 additions and 1015 deletions
@@ -8,9 +8,9 @@ from unittest.mock import patch
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.observability.scheduler_metrics_mixin import (
from sglang.srt.managers.scheduler_components.metrics_reporter import (
PrefillStats,
SchedulerMetricsMixin,
SchedulerMetricsReporter,
)
@@ -85,14 +85,49 @@ class _DummyPublisherThread:
pass
class _DummyScheduler(SchedulerMetricsMixin):
pass
def _make_reporter(scheduler) -> SchedulerMetricsReporter:
if not hasattr(scheduler, "server_args"):
scheduler.server_args = types.SimpleNamespace(
enable_metrics=False,
enable_metrics_for_all_schedulers=False,
kv_events_config=None,
enable_mfu_metrics=False,
enable_forward_pass_metrics=False,
)
if not hasattr(scheduler, "ps"):
scheduler.ps = types.SimpleNamespace(attn_tp_rank=0, attn_cp_rank=0)
if not hasattr(scheduler, "kv_events_publisher"):
scheduler.kv_events_publisher = types.SimpleNamespace(
init_kv_events=lambda *a, **kw: None,
)
if not hasattr(scheduler, "tp_workers"):
scheduler.tp_workers = []
if not hasattr(scheduler, "tp_worker"):
scheduler.tp_worker = types.SimpleNamespace(
model_runner=types.SimpleNamespace(),
)
if not hasattr(scheduler, "draft_worker"):
scheduler.draft_worker = None
context = types.SimpleNamespace(
enable_metrics=False,
is_stats_logging_rank=True,
current_scheduler_metrics_enabled=False,
enable_kv_cache_events=False,
collector=None,
)
return SchedulerMetricsReporter(
scheduler=scheduler,
tp_rank=0,
pp_rank=0,
dp_rank=0,
metrics_collector_context=context,
metrics_collector=None,
)
class TestForwardPassMetrics(unittest.TestCase):
def setUp(self):
self.scheduler = _DummyScheduler()
self.scheduler.enable_fpm = True
self.scheduler = types.SimpleNamespace()
self.scheduler._fpm_worker_id = "worker-7"
self.scheduler._fpm_dp_rank = 0
self.scheduler._fpm_publisher = _CollectingPublisher()
@@ -100,6 +135,8 @@ class TestForwardPassMetrics(unittest.TestCase):
self.scheduler._fpm_gpu_time_acc = 0.0
self.scheduler.waiting_queue = []
self.scheduler.disaggregation_mode = DisaggregationMode.NULL
self.reporter = _make_reporter(self.scheduler)
self.scheduler.enable_fpm = True
def _make_batch(self, **overrides):
defaults = dict(
@@ -135,10 +172,10 @@ class TestForwardPassMetrics(unittest.TestCase):
)
with patch(
"sglang.srt.observability.scheduler_metrics_mixin.time.monotonic",
"sglang.srt.managers.scheduler_components.metrics_reporter.time.monotonic",
return_value=104.5,
):
self.scheduler._emit_forward_pass_metrics(batch)
self.reporter._emit_forward_pass_metrics(batch)
self.assertEqual(len(self.scheduler._fpm_publisher.metrics), 1)
metrics = self.scheduler._fpm_publisher.metrics[0]
@@ -158,12 +195,12 @@ class TestForwardPassMetrics(unittest.TestCase):
def test_emit_uses_device_timer_gpu_time(self):
self.scheduler._fpm_uses_device_timer = True
self.scheduler._fpm_gpu_time_acc = 0.042
self.scheduler.forward_pass_device_timer = types.SimpleNamespace(
self.reporter.forward_pass_device_timer = types.SimpleNamespace(
_report=lambda: None,
)
batch = self._make_batch()
self.scheduler._emit_forward_pass_metrics(batch)
self.reporter._emit_forward_pass_metrics(batch)
self.assertEqual(len(self.scheduler._fpm_publisher.metrics), 1)
self.assertAlmostEqual(
@@ -174,12 +211,12 @@ class TestForwardPassMetrics(unittest.TestCase):
def test_emit_skips_when_device_timer_zero(self):
self.scheduler._fpm_uses_device_timer = True
self.scheduler._fpm_gpu_time_acc = 0.0
self.scheduler.forward_pass_device_timer = types.SimpleNamespace(
self.reporter.forward_pass_device_timer = types.SimpleNamespace(
_report=lambda: None,
)
batch = self._make_batch()
self.scheduler._emit_forward_pass_metrics(batch)
self.reporter._emit_forward_pass_metrics(batch)
self.assertEqual(len(self.scheduler._fpm_publisher.metrics), 0)
@@ -187,10 +224,10 @@ class TestForwardPassMetrics(unittest.TestCase):
batch = self._make_batch()
with patch(
"sglang.srt.observability.scheduler_metrics_mixin.time.monotonic",
"sglang.srt.managers.scheduler_components.metrics_reporter.time.monotonic",
return_value=100.035,
):
self.scheduler._emit_forward_pass_metrics(batch, result=None)
self.reporter._emit_forward_pass_metrics(batch, result=None)
self.assertEqual(len(self.scheduler._fpm_publisher.metrics), 1)
self.assertAlmostEqual(
@@ -205,10 +242,10 @@ class TestForwardPassMetrics(unittest.TestCase):
batch = self._make_batch()
with patch(
"sglang.srt.observability.scheduler_metrics_mixin.time.monotonic",
"sglang.srt.managers.scheduler_components.metrics_reporter.time.monotonic",
return_value=101.0,
):
self.scheduler._emit_forward_pass_metrics(batch)
self.reporter._emit_forward_pass_metrics(batch)
metrics = self.scheduler._fpm_publisher.metrics[0]
self.assertEqual(metrics.queued_requests.num_prefill_requests, 3)
@@ -226,10 +263,10 @@ class TestForwardPassMetrics(unittest.TestCase):
batch = self._make_batch()
with patch(
"sglang.srt.observability.scheduler_metrics_mixin.time.monotonic",
"sglang.srt.managers.scheduler_components.metrics_reporter.time.monotonic",
return_value=101.0,
):
self.scheduler._emit_forward_pass_metrics(batch)
self.reporter._emit_forward_pass_metrics(batch)
metrics = self.scheduler._fpm_publisher.metrics[0]
self.assertEqual(metrics.queued_requests.num_prefill_requests, 0)
@@ -237,7 +274,7 @@ class TestForwardPassMetrics(unittest.TestCase):
self.assertEqual(metrics.queued_requests.sum_decode_kv_tokens, 15 + 30 + 45)
def test_init_metrics_uses_server_worker_id(self):
scheduler = _DummyScheduler()
scheduler = types.SimpleNamespace()
scheduler.server_args = types.SimpleNamespace(
enable_metrics=False,
enable_metrics_for_all_schedulers=False,
@@ -254,7 +291,7 @@ class TestForwardPassMetrics(unittest.TestCase):
"sglang.srt.observability.forward_pass_metrics._FpmPublisherThread",
_DummyPublisherThread,
):
scheduler.init_metrics(tp_rank=0, pp_rank=0, dp_rank=2)
reporter = _make_reporter(scheduler)
self.assertTrue(scheduler.enable_fpm)
self.assertEqual(scheduler._fpm_worker_id, "endpoint-42")
@@ -265,7 +302,7 @@ class TestForwardPassMetrics(unittest.TestCase):
self.assertIsNotNone(scheduler.server_args.forward_pass_metrics_ipc_name)
def test_init_fpm_disabled_on_non_last_pp_rank(self):
scheduler = _DummyScheduler()
scheduler = types.SimpleNamespace()
scheduler.server_args = types.SimpleNamespace(
enable_metrics=False,
enable_metrics_for_all_schedulers=False,
@@ -282,7 +319,7 @@ class TestForwardPassMetrics(unittest.TestCase):
"sglang.srt.observability.forward_pass_metrics._FpmPublisherThread",
_DummyPublisherThread,
):
scheduler.init_metrics(tp_rank=0, pp_rank=0, dp_rank=0)
reporter = _make_reporter(scheduler)
self.assertFalse(scheduler.enable_fpm)