diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 6db0f8b0f..75ed09458 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -1474,6 +1474,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): split_forward_batch: ForwardBatch = None seq_lens_cpu_cache: torch.Tensor = None + # Forward-pass metrics + fpm_start_time: float = 0.0 + # Stream has_stream: bool = False @@ -2638,6 +2641,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): mamba_track_seqlens=self.mamba_track_seqlens, dp_cooperation_info=self.dp_cooperation_info, prefill_stats=self.prefill_stats, + fpm_start_time=self.fpm_start_time, forward_iter=self.forward_iter, ) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 2d33b3fc0..689cf9837 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2459,6 +2459,8 @@ class Scheduler( return batch def get_next_batch_to_run(self) -> Optional[ScheduleBatch]: + if self.enable_fpm: + self._fpm_batch_t0 = time.monotonic() self._abort_on_waiting_timeout() self._abort_on_running_timeout() if self.dllm_config is not None: @@ -2572,6 +2574,8 @@ class Scheduler( if ret: set_schedule_time_batch(ret) + if self.enable_fpm: + ret.fpm_start_time = self._fpm_batch_t0 return ret @@ -3153,6 +3157,11 @@ class Scheduler( self.process_batch_result_idle(batch, result) self.log_batch_result_stats(batch, result) + + # Emit forward pass metrics (every iteration when enabled) + if self.enable_fpm: + self._emit_forward_pass_metrics(batch, result) + self._maybe_clear_mm_inputs(batch) self.maybe_send_health_check_signal() self.update_device_timer() @@ -3981,6 +3990,7 @@ def run_scheduler_process( trace_set_thread_info(thread_label, tp_rank, dp_rank, pp_rank) # Create a scheduler and run the event loop + scheduler = None try: scheduler = Scheduler( server_args, @@ -4004,3 +4014,8 @@ def run_scheduler_process( traceback = get_exception_traceback() logger.error(f"Scheduler hit an exception: {traceback}") parent_process.send_signal(signal.SIGQUIT) + finally: + if scheduler is not None: + # FPM has a background ZMQ publisher thread that needs explicit + # teardown to flush queued metrics and close the socket cleanly. + scheduler._shutdown_fpm() diff --git a/python/sglang/srt/managers/utils.py b/python/sglang/srt/managers/utils.py index 6ac057572..a277f9b79 100644 --- a/python/sglang/srt/managers/utils.py +++ b/python/sglang/srt/managers/utils.py @@ -55,6 +55,10 @@ class GenerationBatchResult: # metrics expert_distribution_metrics: Optional[ExpertDistributionMetrics] = None + # Forward pass metrics (FPM) — GPU-accurate timing via CUDA events + fpm_start_event: Optional[torch.cuda.Event] = None + fpm_end_event: Optional[torch.cuda.Event] = None + def copy_to_cpu(self, return_logprob: bool): """Copy tensors to CPU in overlap scheduling. Only the tensors which are needed for processing results are copied, diff --git a/python/sglang/srt/observability/forward_pass_metrics.py b/python/sglang/srt/observability/forward_pass_metrics.py new file mode 100644 index 000000000..e271bd6de --- /dev/null +++ b/python/sglang/srt/observability/forward_pass_metrics.py @@ -0,0 +1,221 @@ +""" +Forward pass metrics for per-iteration scheduler telemetry. + +Emits per-iteration scheduling metrics over ZMQ PUB so that external +consumers can observe scheduler behavior in real time without polling +Prometheus. + +Uses msgspec.Struct for zero-copy serialization. + +Data flow:: + + Scheduler process: + SchedulerMetricsMixin._emit_forward_pass_metrics() + -> _FpmPublisherThread -> ZMQ PUB (localhost) + + External consumer: + ZMQ SUB -> deserialize ForwardPassMetrics + +""" + +from __future__ import annotations + +import logging +import queue +import threading +import time +from itertools import count + +import msgspec + +# Schema version. Must match the consumer (Dynamo's ForwardPassMetrics). +# Bump when the schema changes incompatibly. +FPM_VERSION: int = 1 + +logger = logging.getLogger(__name__) + + +class WelfordAccumulator: + """Welford's online algorithm for count / total / population-variance. + + Numerically stable single-pass computation. + """ + + __slots__ = ("count", "total", "_mean", "_m2") + + def __init__(self) -> None: + self.count = 0 + self.total = 0 + self._mean = 0.0 + self._m2 = 0.0 + + def add(self, v: int) -> None: + self.count += 1 + self.total += v + delta = v - self._mean + self._mean += delta / self.count + delta2 = v - self._mean + self._m2 += delta * delta2 + + def variance(self) -> float: + if self.count == 0: + return 0.0 + return self._m2 / self.count + + +class ScheduledRequestMetrics( + msgspec.Struct, + frozen=True, + gc=False, +): + """Metrics for requests scheduled in this iteration.""" + + num_prefill_requests: int = 0 + sum_prefill_tokens: int = 0 + var_prefill_length: float = 0.0 + sum_prefill_kv_tokens: int = 0 + num_decode_requests: int = 0 + sum_decode_kv_tokens: int = 0 + var_decode_kv_tokens: float = 0.0 + + +class QueuedRequestMetrics( + msgspec.Struct, + frozen=True, + gc=False, +): + """Metrics for requests waiting in the queue.""" + + num_prefill_requests: int = 0 + sum_prefill_tokens: int = 0 + var_prefill_length: float = 0.0 + num_decode_requests: int = 0 + sum_decode_kv_tokens: int = 0 + var_decode_kv_tokens: float = 0.0 + + +class ForwardPassMetrics( + msgspec.Struct, + frozen=True, + gc=False, +): + """Per-iteration metrics emitted by the scheduler. + + One message per scheduler iteration (one per forward pass). + ``wall_time`` is the iteration duration in seconds. + An idle heartbeat (all zeros, wall_time=0) is emitted when the + engine transitions from active to idle. + + Field order must match Dynamo's ``ForwardPassMetrics`` in + ``dynamo.common.forward_pass_metrics`` — msgspec uses positional + encoding so any mismatch silently corrupts data. + """ + + version: int = FPM_VERSION + worker_id: str = "" + dp_rank: int = 0 + counter_id: int = 0 + wall_time: float = 0.0 + scheduled_requests: ScheduledRequestMetrics = ScheduledRequestMetrics() + queued_requests: QueuedRequestMetrics = QueuedRequestMetrics() + + +_encoder = msgspec.msgpack.Encoder() +_decoder = msgspec.msgpack.Decoder(ForwardPassMetrics) + + +def encode(metrics: ForwardPassMetrics) -> bytes: + return _encoder.encode(metrics) + + +def decode(data: bytes) -> ForwardPassMetrics: + return _decoder.decode(data) + + +class _FpmPublisherThread: + """Background thread that serializes and sends ForwardPassMetrics over ZMQ. + + Also emits periodic heartbeats when idle. + """ + + SHUTDOWN_TIMEOUT: float = 1.0 + HEARTBEAT_INTERVAL: float = 1.0 + + def __init__( + self, + endpoint: str, + worker_id: str, + dp_rank: int, + max_queue_size: int = 10_000, + ) -> None: + import zmq + + self._queue: queue.Queue[ForwardPassMetrics | None] = queue.Queue( + maxsize=max_queue_size + ) + self._seq = count() + self._worker_id = worker_id + self._dp_rank = dp_rank + + self._ctx = zmq.Context() + self._pub = self._ctx.socket(zmq.PUB) + self._pub.bind(endpoint) + self._zmq = zmq + + self._running = True + self._thread = threading.Thread( + target=self._run, daemon=True, name="fpm-zmq-publisher" + ) + self._thread.start() + + def publish(self, metrics: ForwardPassMetrics) -> None: + if not self._running: + return + try: + self._queue.put_nowait(metrics) + except queue.Full: + pass + + def shutdown(self) -> None: + self._running = False + try: + self._queue.put_nowait(None) + except queue.Full: + pass + self._thread.join(timeout=self.SHUTDOWN_TIMEOUT) + try: + self._pub.close(linger=0) + self._ctx.term() + except Exception: + pass + + def _run(self) -> None: + zmq = self._zmq + topic = b"" + last_publish = time.monotonic() + + while self._running or not self._queue.empty(): + try: + metrics = self._queue.get(timeout=self.HEARTBEAT_INTERVAL) + if metrics is None: + break + except queue.Empty: + if time.monotonic() - last_publish >= self.HEARTBEAT_INTERVAL: + metrics = ForwardPassMetrics( + worker_id=self._worker_id, + dp_rank=self._dp_rank, + ) + else: + continue + + try: + seq = next(self._seq) + metrics = msgspec.structs.replace(metrics, counter_id=seq) + payload = encode(metrics) + seq_bytes = seq.to_bytes(8, "big") + self._pub.send_multipart((topic, seq_bytes, payload), flags=zmq.NOBLOCK) + last_publish = time.monotonic() + except zmq.Again: + pass + except Exception: + logger.warning("FPM publisher send failed", exc_info=True) diff --git a/python/sglang/srt/observability/scheduler_metrics_mixin.py b/python/sglang/srt/observability/scheduler_metrics_mixin.py index a1afa6c88..050895373 100644 --- a/python/sglang/srt/observability/scheduler_metrics_mixin.py +++ b/python/sglang/srt/observability/scheduler_metrics_mixin.py @@ -2,6 +2,7 @@ from __future__ import annotations import dataclasses import logging +import tempfile import time from collections import defaultdict from typing import TYPE_CHECKING, List, Optional, Tuple, Union @@ -18,7 +19,7 @@ from sglang.srt.managers.io_struct import ( QueueMetrics, SpeculativeMetrics, ) -from sglang.srt.managers.scheduler import ScheduleBatch +from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.utils import GenerationBatchResult from sglang.srt.observability.metrics_collector import ( DPCooperationInfo, @@ -86,6 +87,8 @@ class KvMetrics: class SchedulerMetricsMixin: + enable_fpm: bool = False + def init_metrics( self: Scheduler, tp_rank: int, pp_rank: int, dp_rank: Optional[int] ): @@ -175,6 +178,8 @@ class SchedulerMetricsMixin: self.init_kv_events(self.server_args.kv_events_config) + self._init_fpm() + self.scheduler_status_logger = SchedulerStatusLogger.maybe_create( enable_metrics=self.enable_metrics ) @@ -202,6 +207,128 @@ class SchedulerMetricsMixin: kv_events_config, self.attn_dp_rank ) + def _init_fpm(self: Scheduler): + """Initialize Forward Pass Metrics (FPM) publisher if configured.""" + self.enable_fpm = False + if ( + self.server_args.enable_forward_pass_metrics + and self.attn_tp_rank == 0 + and self.pp_rank == self.pp_size - 1 + ): + from sglang.srt.observability.forward_pass_metrics import ( + _FpmPublisherThread, + ) + + self._fpm_dp_rank = self.dp_rank if self.dp_rank is not None else 0 + self._fpm_worker_id = self.server_args.forward_pass_metrics_worker_id + base_endpoint = self.server_args.forward_pass_metrics_ipc_name + if base_endpoint is None: + ipc_path = tempfile.NamedTemporaryFile(delete=False).name + base_endpoint = f"ipc://{ipc_path}" + self.server_args.forward_pass_metrics_ipc_name = base_endpoint + endpoint = f"{base_endpoint}.{self._fpm_dp_rank}" + self._fpm_publisher = _FpmPublisherThread( + endpoint, + worker_id=self._fpm_worker_id, + dp_rank=self._fpm_dp_rank, + ) + self._fpm_gpu_time_acc = 0.0 + + def _fpm_device_timer_reporter(t, **_kwargs): + self._fpm_gpu_time_acc += t + + if hasattr(self, "forward_pass_device_timer"): + self.forward_pass_device_timer.add_reporter(_fpm_device_timer_reporter) + else: + self.forward_pass_device_timer = DeviceTimer( + reporter=_fpm_device_timer_reporter, + ) + self._fpm_uses_device_timer = True + self.enable_fpm = True + logger.info( + "FPM: ZMQ PUB bound on %s (dp_rank=%d, device_timer=%s)", + endpoint, + self._fpm_dp_rank, + self._fpm_uses_device_timer, + ) + + def _build_scheduled_request_metrics(self: Scheduler, batch: ScheduleBatch): + from sglang.srt.observability.forward_pass_metrics import ( + ScheduledRequestMetrics, + WelfordAccumulator, + ) + + num_prefill_requests = 0 + sum_prefill_tokens = 0 + sum_prefill_kv_tokens = 0 + prefill_lengths = WelfordAccumulator() + + if batch.forward_mode.is_mixed(): + decode_req_ids = {id(req) for req in batch.decoding_reqs or []} + prefill_reqs = [req for req in batch.reqs if id(req) not in decode_req_ids] + elif batch.forward_mode.is_extend(): + prefill_reqs = batch.reqs + else: + prefill_reqs = [] + + if prefill_reqs: + stats = batch.prefill_stats + for req in prefill_reqs: + prefill_lengths.add(len(req.origin_input_ids)) + num_prefill_requests = stats.num_new_seqs if stats else len(prefill_reqs) + sum_prefill_tokens = stats.log_input_tokens if stats else 0 + sum_prefill_kv_tokens = sum(len(req.prefix_indices) for req in prefill_reqs) + + decode_kv = WelfordAccumulator() + if batch.forward_mode.is_mixed(): + for req in batch.decoding_reqs or []: + decode_kv.add(req.seqlen) + elif batch.forward_mode.is_decode(): + for sl in batch.seq_lens_cpu: + decode_kv.add(int(sl)) + + return ScheduledRequestMetrics( + num_prefill_requests=num_prefill_requests, + sum_prefill_tokens=sum_prefill_tokens, + var_prefill_length=prefill_lengths.variance(), + sum_prefill_kv_tokens=sum_prefill_kv_tokens, + num_decode_requests=decode_kv.count, + sum_decode_kv_tokens=decode_kv.total, + var_decode_kv_tokens=decode_kv.variance(), + ) + + def _build_queued_request_metrics(self: Scheduler): + from sglang.srt.observability.forward_pass_metrics import ( + QueuedRequestMetrics, + WelfordAccumulator, + ) + + prefill_q = WelfordAccumulator() + decode_q = WelfordAccumulator() + if self.disaggregation_mode == DisaggregationMode.PREFILL: + for req in self.disagg_prefill_bootstrap_queue.queue: + prefill_q.add(len(req.origin_input_ids)) + elif self.disaggregation_mode == DisaggregationMode.DECODE: + for req in self.disagg_decode_prealloc_queue.queue: + decode_q.add(req.seqlen) + for req in self.disagg_decode_transfer_queue.queue: + decode_q.add(req.seqlen) + else: + for req in self.waiting_queue: + if len(req.output_ids) > 0: + decode_q.add(req.seqlen) + else: + prefill_q.add(len(req.origin_input_ids)) + + return QueuedRequestMetrics( + num_prefill_requests=prefill_q.count, + sum_prefill_tokens=prefill_q.total, + var_prefill_length=prefill_q.variance(), + num_decode_requests=decode_q.count, + sum_decode_kv_tokens=decode_q.total, + var_decode_kv_tokens=decode_q.variance(), + ) + def update_spec_metrics(self: Scheduler, bs: int, num_correct_drafts: int): self.spec_num_accept_tokens += num_correct_drafts + bs self.spec_num_forward_ct += bs @@ -719,6 +846,47 @@ class SchedulerMetricsMixin: batch = KVEventBatch(ts=time.time(), events=events) self.kv_event_publisher.publish(batch) + def _emit_forward_pass_metrics( + self: Scheduler, + batch: ScheduleBatch, + result=None, + ): + """Emit per-iteration ForwardPassMetrics over ZMQ PUB. + + Prefers GPU-accurate timing from DeviceTimer (which wraps + model_runner.forward / cuda_graph.replay via PR #24197). + Falls back to monotonic clock when DeviceTimer is not enabled. + """ + if not self.enable_fpm: + return + + from sglang.srt.observability.forward_pass_metrics import ( + ForwardPassMetrics, + ) + + if self._fpm_uses_device_timer: + self.forward_pass_device_timer._report() + wall_time = self._fpm_gpu_time_acc + self._fpm_gpu_time_acc = 0.0 + if wall_time == 0.0: + return + else: + wall_time = max(0.0, time.monotonic() - batch.fpm_start_time) + + fpm = ForwardPassMetrics( + worker_id=self._fpm_worker_id, + dp_rank=self._fpm_dp_rank, + wall_time=wall_time, + scheduled_requests=self._build_scheduled_request_metrics(batch), + queued_requests=self._build_queued_request_metrics(), + ) + self._fpm_publisher.publish(fpm) + + def _shutdown_fpm(self: Scheduler): + """Shut down the FPM publisher thread.""" + if self.enable_fpm: + self._fpm_publisher.shutdown() + def _log_hicache_stats(self: Scheduler): """Populate HiCache host-tier stats on self.stats. diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index f5f266fee..8a49db505 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -487,6 +487,9 @@ class ServerArgs: decode_log_interval: int = 40 enable_request_time_stats_logging: bool = False kv_events_config: Optional[str] = None + enable_forward_pass_metrics: bool = False + forward_pass_metrics_worker_id: str = "" + forward_pass_metrics_ipc_name: Optional[str] = None enable_trace: bool = False otlp_traces_endpoint: str = "localhost:4317" @@ -5236,6 +5239,25 @@ class ServerArgs: default=None, help="Config in json format for NVIDIA dynamo KV event publishing. Publishing will be enabled if this flag is used.", ) + parser.add_argument( + "--enable-forward-pass-metrics", + action="store_true", + help="Enable per-iteration forward pass metrics via ZMQ IPC. " + "External consumers (e.g. Dynamo planner) subscribe to the IPC " + "endpoint exposed in server_args.forward_pass_metrics_ipc_name.", + ) + parser.add_argument( + "--forward-pass-metrics-worker-id", + type=str, + default="", + help=argparse.SUPPRESS, + ) + parser.add_argument( + "--forward-pass-metrics-ipc-name", + type=str, + default=None, + help=argparse.SUPPRESS, + ) parser.add_argument( "--enable-trace", action="store_true", diff --git a/python/sglang/srt/utils/device_timer.py b/python/sglang/srt/utils/device_timer.py index eaa44b7f4..3562e4df3 100644 --- a/python/sglang/srt/utils/device_timer.py +++ b/python/sglang/srt/utils/device_timer.py @@ -1,7 +1,7 @@ from collections import deque from contextlib import contextmanager from dataclasses import dataclass -from typing import Callable, Deque, Dict, Optional +from typing import Callable, Deque, Dict, List, Optional import torch @@ -9,7 +9,10 @@ import torch class DeviceTimer: def __init__(self, reporter: Callable): self._intervals: Deque[_TimingInterval] = deque() - self._reporter = reporter + self._reporters: List[Callable] = [reporter] + + def add_reporter(self, reporter: Callable): + self._reporters.append(reporter) @contextmanager def wrap(self, metadata: Dict): @@ -27,8 +30,9 @@ class DeviceTimer: break self._intervals.popleft() - self._reporter(t=interval.elapsed_time() / 1000.0, **interval.metadata) - # print(f"{interval.elapsed_time()=:.6f}, {interval.metadata=}") + elapsed = interval.elapsed_time() / 1000.0 + for reporter in self._reporters: + reporter(t=elapsed, **interval.metadata) class GapTimer(DeviceTimer): diff --git a/test/manual/test_forward_pass_metrics.py b/test/manual/test_forward_pass_metrics.py new file mode 100644 index 000000000..7243e8545 --- /dev/null +++ b/test/manual/test_forward_pass_metrics.py @@ -0,0 +1,191 @@ +""" +Manual test for Forward Pass Metrics (FPM) ZMQ PUB/SUB path. + +Tests: +1. Schema encode/decode roundtrip +2. _FpmPublisherThread ZMQ PUB -> ZMQ SUB end-to-end +3. Heartbeat emission on idle +""" + +import sys +import time + +import zmq + + +def test_schema_roundtrip(): + from sglang.srt.observability.forward_pass_metrics import ( + ForwardPassMetrics, + QueuedRequestMetrics, + ScheduledRequestMetrics, + WelfordAccumulator, + decode, + encode, + ) + + # WelfordAccumulator + acc = WelfordAccumulator() + for v in [10, 20, 30]: + acc.add(v) + assert acc.count == 3 + assert acc.total == 60 + var = acc.variance() + assert abs(var - 66.667) < 0.01, f"Expected ~66.667, got {var}" + + # Encode/decode roundtrip + fpm = ForwardPassMetrics( + worker_id="test-worker", + dp_rank=1, + wall_time=0.042, + scheduled_requests=ScheduledRequestMetrics( + num_prefill_requests=5, + sum_prefill_tokens=1024, + var_prefill_length=33.3, + sum_prefill_kv_tokens=512, + num_decode_requests=32, + sum_decode_kv_tokens=8192, + var_decode_kv_tokens=100.0, + ), + queued_requests=QueuedRequestMetrics( + num_prefill_requests=3, + sum_prefill_tokens=768, + var_prefill_length=25.0, + num_decode_requests=1, + sum_decode_kv_tokens=128, + var_decode_kv_tokens=0.0, + ), + ) + + data = encode(fpm) + fpm2 = decode(data) + + assert fpm2.worker_id == "test-worker" + assert fpm2.dp_rank == 1 + assert fpm2.wall_time == 0.042 + assert fpm2.scheduled_requests.num_prefill_requests == 5 + assert fpm2.scheduled_requests.sum_prefill_tokens == 1024 + assert fpm2.scheduled_requests.num_decode_requests == 32 + assert fpm2.queued_requests.num_prefill_requests == 3 + + print("PASS: schema roundtrip") + + +def test_zmq_pub_sub(): + """Test _FpmPublisherThread -> ZMQ SUB end-to-end.""" + from sglang.srt.observability.forward_pass_metrics import ( + ForwardPassMetrics, + ScheduledRequestMetrics, + _FpmPublisherThread, + decode, + ) + + port = 29999 + endpoint = f"tcp://127.0.0.1:{port}" + + # Start publisher + pub = _FpmPublisherThread( + f"tcp://*:{port}", + worker_id="test-pub", + dp_rank=0, + ) + + # Connect subscriber + ctx = zmq.Context() + sub = ctx.socket(zmq.SUB) + sub.connect(endpoint) + sub.setsockopt(zmq.SUBSCRIBE, b"") + sub.setsockopt(zmq.RCVTIMEO, 5000) # 5s timeout + + # ZMQ PUB/SUB needs time to connect + time.sleep(0.5) + + # Publish a metric + fpm = ForwardPassMetrics( + worker_id="test-pub", + dp_rank=0, + wall_time=0.05, + scheduled_requests=ScheduledRequestMetrics( + num_prefill_requests=10, + sum_prefill_tokens=2048, + num_decode_requests=64, + sum_decode_kv_tokens=16384, + ), + ) + pub.publish(fpm) + + # Receive + frames = sub.recv_multipart() + assert len(frames) == 3, f"Expected 3 frames, got {len(frames)}" + + topic, seq_bytes, payload = frames + assert topic == b"" + seq = int.from_bytes(seq_bytes, "big") + assert seq == 0 + + received = decode(payload) + assert received.worker_id == "test-pub" + assert received.scheduled_requests.num_prefill_requests == 10 + assert received.scheduled_requests.sum_decode_kv_tokens == 16384 + print(f"PASS: ZMQ PUB/SUB (seq={seq}, {len(payload)} bytes)") + + # Publish a second message -- seq should increment + pub.publish(fpm) + frames2 = sub.recv_multipart() + seq2 = int.from_bytes(frames2[1], "big") + assert seq2 == 1, f"Expected seq=1, got {seq2}" + print(f"PASS: sequence incremented (seq={seq2})") + + # Cleanup + pub.shutdown() + sub.close() + ctx.term() + + +def test_heartbeat(): + """Test that heartbeat messages are emitted when idle.""" + from sglang.srt.observability.forward_pass_metrics import ( + _FpmPublisherThread, + decode, + ) + + port = 29998 + endpoint = f"tcp://127.0.0.1:{port}" + + pub = _FpmPublisherThread( + f"tcp://*:{port}", + worker_id="heartbeat-test", + dp_rank=0, + ) + # Override heartbeat interval for faster test + pub.HEARTBEAT_INTERVAL = 0.3 + + ctx = zmq.Context() + sub = ctx.socket(zmq.SUB) + sub.connect(endpoint) + sub.setsockopt(zmq.SUBSCRIBE, b"") + sub.setsockopt(zmq.RCVTIMEO, 3000) + + time.sleep(0.5) + + # Don't publish anything -- wait for heartbeat + try: + frames = sub.recv_multipart() + heartbeat = decode(frames[2]) + assert heartbeat.worker_id == "heartbeat-test" + assert heartbeat.wall_time == 0.0 # idle heartbeat + assert heartbeat.scheduled_requests.num_prefill_requests == 0 + print("PASS: heartbeat received") + except zmq.Again: + print("FAIL: no heartbeat received within timeout") + sys.exit(1) + + pub.shutdown() + sub.close() + ctx.term() + + +if __name__ == "__main__": + test_schema_roundtrip() + test_zmq_pub_sub() + test_heartbeat() + print("\nAll tests passed!") diff --git a/test/registered/unit/observability/test_forward_pass_metrics.py b/test/registered/unit/observability/test_forward_pass_metrics.py new file mode 100644 index 000000000..4eda3665c --- /dev/null +++ b/test/registered/unit/observability/test_forward_pass_metrics.py @@ -0,0 +1,271 @@ +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="stage-a-test-cpu") + +import types +import unittest +from unittest.mock import patch + +from sglang.srt.disaggregation.utils import DisaggregationMode +from sglang.srt.observability.scheduler_metrics_mixin import ( + PrefillStats, + SchedulerMetricsMixin, +) + + +class _FakeReq: + def __init__( + self, + prompt_len: int, + output_len: int = 0, + prefix_len: int = 0, + ): + self.origin_input_ids = list(range(prompt_len)) + self.output_ids = list(range(output_len)) + self.prefix_indices = list(range(prefix_len)) + self.seqlen = prompt_len + output_len + + +class _FakeForwardMode: + def __init__(self, *, is_mixed: bool = False, is_extend: bool = False): + self._is_mixed = is_mixed + self._is_extend = is_extend + + def is_mixed(self): + return self._is_mixed + + def is_extend(self, include_draft_extend_v2: bool = False): + return self._is_extend + + def is_decode(self): + return not self._is_mixed and not self._is_extend + + +class _CollectingPublisher: + def __init__(self): + self.metrics = [] + + def publish(self, metrics): + self.metrics.append(metrics) + + +class _DummyPublisherThread: + def __init__(self, endpoint: str, worker_id: str, dp_rank: int, **_: object): + self.endpoint = endpoint + self.worker_id = worker_id + self.dp_rank = dp_rank + + def shutdown(self): + pass + + +class _DummyScheduler(SchedulerMetricsMixin): + pass + + +class TestForwardPassMetrics(unittest.TestCase): + def setUp(self): + self.scheduler = _DummyScheduler() + self.scheduler.enable_fpm = True + self.scheduler._fpm_worker_id = "worker-7" + self.scheduler._fpm_dp_rank = 0 + self.scheduler._fpm_publisher = _CollectingPublisher() + self.scheduler._fpm_uses_device_timer = False + self.scheduler._fpm_gpu_time_acc = 0.0 + self.scheduler.waiting_queue = [] + self.scheduler.disaggregation_mode = DisaggregationMode.NULL + + def _make_batch(self, **overrides): + defaults = dict( + forward_mode=_FakeForwardMode(), + reqs=[], + decoding_reqs=[], + prefill_stats=None, + seq_lens_cpu=[], + fpm_start_time=100.0, + ) + defaults.update(overrides) + return types.SimpleNamespace(**defaults) + + def test_emit_mixed_batch_separates_prefill_and_decode(self): + self.scheduler._fpm_dp_rank = 3 + self.scheduler.waiting_queue = [_FakeReq(6), _FakeReq(4, output_len=2)] + + prefill_a = _FakeReq(10, prefix_len=2) + prefill_b = _FakeReq(14, prefix_len=3) + decode_req = _FakeReq(8, output_len=3) + batch = self._make_batch( + forward_mode=_FakeForwardMode(is_mixed=True, is_extend=True), + reqs=[prefill_a, prefill_b, decode_req], + decoding_reqs=[decode_req], + prefill_stats=PrefillStats( + log_input_tokens=12, + log_hit_tokens=5, + new_token_ratio=1.0, + num_running_reqs=types.SimpleNamespace(), + num_new_seqs=2, + ), + seq_lens_cpu=[decode_req.seqlen], + ) + + with patch( + "sglang.srt.observability.scheduler_metrics_mixin.time.monotonic", + return_value=104.5, + ): + self.scheduler._emit_forward_pass_metrics(batch) + + self.assertEqual(len(self.scheduler._fpm_publisher.metrics), 1) + metrics = self.scheduler._fpm_publisher.metrics[0] + self.assertEqual(metrics.worker_id, "worker-7") + self.assertEqual(metrics.dp_rank, 3) + self.assertEqual(metrics.wall_time, 4.5) + self.assertEqual(metrics.scheduled_requests.num_prefill_requests, 2) + self.assertEqual(metrics.scheduled_requests.sum_prefill_tokens, 12) + self.assertEqual(metrics.scheduled_requests.sum_prefill_kv_tokens, 5) + self.assertEqual(metrics.scheduled_requests.num_decode_requests, 1) + self.assertEqual( + metrics.scheduled_requests.sum_decode_kv_tokens, decode_req.seqlen + ) + self.assertEqual(metrics.queued_requests.num_prefill_requests, 1) + self.assertEqual(metrics.queued_requests.num_decode_requests, 1) + + 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( + _report=lambda: None, + ) + batch = self._make_batch() + + self.scheduler._emit_forward_pass_metrics(batch) + + self.assertEqual(len(self.scheduler._fpm_publisher.metrics), 1) + self.assertAlmostEqual( + self.scheduler._fpm_publisher.metrics[0].wall_time, 0.042, places=4 + ) + self.assertAlmostEqual(self.scheduler._fpm_gpu_time_acc, 0.0) + + 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( + _report=lambda: None, + ) + batch = self._make_batch() + + self.scheduler._emit_forward_pass_metrics(batch) + + self.assertEqual(len(self.scheduler._fpm_publisher.metrics), 0) + + def test_emit_uses_monotonic_without_device_timer(self): + batch = self._make_batch() + + with patch( + "sglang.srt.observability.scheduler_metrics_mixin.time.monotonic", + return_value=100.035, + ): + self.scheduler._emit_forward_pass_metrics(batch, result=None) + + self.assertEqual(len(self.scheduler._fpm_publisher.metrics), 1) + self.assertAlmostEqual( + self.scheduler._fpm_publisher.metrics[0].wall_time, 0.035, places=4 + ) + + def test_disagg_prefill_queued_metrics(self): + self.scheduler.disaggregation_mode = DisaggregationMode.PREFILL + self.scheduler.disagg_prefill_bootstrap_queue = types.SimpleNamespace( + queue=[_FakeReq(100), _FakeReq(200), _FakeReq(50)], + ) + batch = self._make_batch() + + with patch( + "sglang.srt.observability.scheduler_metrics_mixin.time.monotonic", + return_value=101.0, + ): + self.scheduler._emit_forward_pass_metrics(batch) + + metrics = self.scheduler._fpm_publisher.metrics[0] + self.assertEqual(metrics.queued_requests.num_prefill_requests, 3) + self.assertEqual(metrics.queued_requests.sum_prefill_tokens, 350) + self.assertEqual(metrics.queued_requests.num_decode_requests, 0) + + def test_disagg_decode_queued_metrics(self): + self.scheduler.disaggregation_mode = DisaggregationMode.DECODE + self.scheduler.disagg_decode_prealloc_queue = types.SimpleNamespace( + queue=[_FakeReq(10, output_len=5), _FakeReq(20, output_len=10)], + ) + self.scheduler.disagg_decode_transfer_queue = types.SimpleNamespace( + queue=[_FakeReq(30, output_len=15)], + ) + batch = self._make_batch() + + with patch( + "sglang.srt.observability.scheduler_metrics_mixin.time.monotonic", + return_value=101.0, + ): + self.scheduler._emit_forward_pass_metrics(batch) + + metrics = self.scheduler._fpm_publisher.metrics[0] + self.assertEqual(metrics.queued_requests.num_prefill_requests, 0) + self.assertEqual(metrics.queued_requests.num_decode_requests, 3) + self.assertEqual(metrics.queued_requests.sum_decode_kv_tokens, 15 + 30 + 45) + + def test_init_metrics_uses_server_worker_id(self): + scheduler = _DummyScheduler() + scheduler.server_args = types.SimpleNamespace( + enable_metrics=False, + enable_metrics_for_all_schedulers=False, + extra_metric_labels=None, + enable_forward_pass_metrics=True, + forward_pass_metrics_worker_id="endpoint-42", + forward_pass_metrics_ipc_name=None, + kv_events_config=None, + ) + scheduler.attn_tp_rank = 0 + scheduler.dp_rank = 2 + scheduler.pp_rank = 0 + scheduler.pp_size = 1 + scheduler.enable_kv_cache_events = False + + with patch( + "sglang.srt.observability.forward_pass_metrics._FpmPublisherThread", + _DummyPublisherThread, + ): + scheduler.init_metrics(tp_rank=0, pp_rank=0, dp_rank=2) + + self.assertTrue(scheduler.enable_fpm) + self.assertEqual(scheduler._fpm_worker_id, "endpoint-42") + self.assertEqual(scheduler._fpm_dp_rank, 2) + self.assertEqual(scheduler._fpm_publisher.worker_id, "endpoint-42") + self.assertEqual(scheduler._fpm_publisher.dp_rank, 2) + self.assertTrue(scheduler._fpm_publisher.endpoint.startswith("ipc://")) + self.assertIsNotNone(scheduler.server_args.forward_pass_metrics_ipc_name) + + def test_init_fpm_disabled_on_non_last_pp_rank(self): + scheduler = _DummyScheduler() + scheduler.server_args = types.SimpleNamespace( + enable_metrics=False, + enable_metrics_for_all_schedulers=False, + extra_metric_labels=None, + enable_forward_pass_metrics=True, + forward_pass_metrics_worker_id="endpoint-42", + forward_pass_metrics_ipc_name=None, + kv_events_config=None, + ) + scheduler.attn_tp_rank = 0 + scheduler.dp_rank = 0 + scheduler.pp_rank = 0 + scheduler.pp_size = 2 + scheduler.enable_kv_cache_events = False + + with patch( + "sglang.srt.observability.forward_pass_metrics._FpmPublisherThread", + _DummyPublisherThread, + ): + scheduler.init_metrics(tp_rank=0, pp_rank=0, dp_rank=0) + + self.assertFalse(scheduler.enable_fpm) + + +if __name__ == "__main__": + unittest.main()