feat: emit per-iteration forward pass metrics via ZMQ PUB (#22789)

Co-authored-by: Ishan Dhanani <ishandhanani@gmail.com>
Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
Co-authored-by: ishandhanani <82981111+ishandhanani@users.noreply.github.com>
This commit is contained in:
Krishnan Prashanth
2026-05-12 10:28:17 -07:00
committed by GitHub
co-authored by Ishan Dhanani Claude Opus 4.6 ishandhanani
parent fd3eb77d45
commit e86fb42736
9 changed files with 905 additions and 5 deletions
@@ -1474,6 +1474,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
split_forward_batch: ForwardBatch = None split_forward_batch: ForwardBatch = None
seq_lens_cpu_cache: torch.Tensor = None seq_lens_cpu_cache: torch.Tensor = None
# Forward-pass metrics
fpm_start_time: float = 0.0
# Stream # Stream
has_stream: bool = False has_stream: bool = False
@@ -2638,6 +2641,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
mamba_track_seqlens=self.mamba_track_seqlens, mamba_track_seqlens=self.mamba_track_seqlens,
dp_cooperation_info=self.dp_cooperation_info, dp_cooperation_info=self.dp_cooperation_info,
prefill_stats=self.prefill_stats, prefill_stats=self.prefill_stats,
fpm_start_time=self.fpm_start_time,
forward_iter=self.forward_iter, forward_iter=self.forward_iter,
) )
+15
View File
@@ -2459,6 +2459,8 @@ class Scheduler(
return batch return batch
def get_next_batch_to_run(self) -> Optional[ScheduleBatch]: 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_waiting_timeout()
self._abort_on_running_timeout() self._abort_on_running_timeout()
if self.dllm_config is not None: if self.dllm_config is not None:
@@ -2572,6 +2574,8 @@ class Scheduler(
if ret: if ret:
set_schedule_time_batch(ret) set_schedule_time_batch(ret)
if self.enable_fpm:
ret.fpm_start_time = self._fpm_batch_t0
return ret return ret
@@ -3153,6 +3157,11 @@ class Scheduler(
self.process_batch_result_idle(batch, result) self.process_batch_result_idle(batch, result)
self.log_batch_result_stats(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_clear_mm_inputs(batch)
self.maybe_send_health_check_signal() self.maybe_send_health_check_signal()
self.update_device_timer() self.update_device_timer()
@@ -3981,6 +3990,7 @@ def run_scheduler_process(
trace_set_thread_info(thread_label, tp_rank, dp_rank, pp_rank) trace_set_thread_info(thread_label, tp_rank, dp_rank, pp_rank)
# Create a scheduler and run the event loop # Create a scheduler and run the event loop
scheduler = None
try: try:
scheduler = Scheduler( scheduler = Scheduler(
server_args, server_args,
@@ -4004,3 +4014,8 @@ def run_scheduler_process(
traceback = get_exception_traceback() traceback = get_exception_traceback()
logger.error(f"Scheduler hit an exception: {traceback}") logger.error(f"Scheduler hit an exception: {traceback}")
parent_process.send_signal(signal.SIGQUIT) 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()
+4
View File
@@ -55,6 +55,10 @@ class GenerationBatchResult:
# metrics # metrics
expert_distribution_metrics: Optional[ExpertDistributionMetrics] = None 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): def copy_to_cpu(self, return_logprob: bool):
"""Copy tensors to CPU in overlap scheduling. """Copy tensors to CPU in overlap scheduling.
Only the tensors which are needed for processing results are copied, Only the tensors which are needed for processing results are copied,
@@ -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)
@@ -2,6 +2,7 @@ from __future__ import annotations
import dataclasses import dataclasses
import logging import logging
import tempfile
import time import time
from collections import defaultdict from collections import defaultdict
from typing import TYPE_CHECKING, List, Optional, Tuple, Union from typing import TYPE_CHECKING, List, Optional, Tuple, Union
@@ -18,7 +19,7 @@ from sglang.srt.managers.io_struct import (
QueueMetrics, QueueMetrics,
SpeculativeMetrics, 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.managers.utils import GenerationBatchResult
from sglang.srt.observability.metrics_collector import ( from sglang.srt.observability.metrics_collector import (
DPCooperationInfo, DPCooperationInfo,
@@ -86,6 +87,8 @@ class KvMetrics:
class SchedulerMetricsMixin: class SchedulerMetricsMixin:
enable_fpm: bool = False
def init_metrics( def init_metrics(
self: Scheduler, tp_rank: int, pp_rank: int, dp_rank: Optional[int] 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_kv_events(self.server_args.kv_events_config)
self._init_fpm()
self.scheduler_status_logger = SchedulerStatusLogger.maybe_create( self.scheduler_status_logger = SchedulerStatusLogger.maybe_create(
enable_metrics=self.enable_metrics enable_metrics=self.enable_metrics
) )
@@ -202,6 +207,128 @@ class SchedulerMetricsMixin:
kv_events_config, self.attn_dp_rank 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): def update_spec_metrics(self: Scheduler, bs: int, num_correct_drafts: int):
self.spec_num_accept_tokens += num_correct_drafts + bs self.spec_num_accept_tokens += num_correct_drafts + bs
self.spec_num_forward_ct += bs self.spec_num_forward_ct += bs
@@ -719,6 +846,47 @@ class SchedulerMetricsMixin:
batch = KVEventBatch(ts=time.time(), events=events) batch = KVEventBatch(ts=time.time(), events=events)
self.kv_event_publisher.publish(batch) 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): def _log_hicache_stats(self: Scheduler):
"""Populate HiCache host-tier stats on self.stats. """Populate HiCache host-tier stats on self.stats.
+22
View File
@@ -487,6 +487,9 @@ class ServerArgs:
decode_log_interval: int = 40 decode_log_interval: int = 40
enable_request_time_stats_logging: bool = False enable_request_time_stats_logging: bool = False
kv_events_config: Optional[str] = None 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 enable_trace: bool = False
otlp_traces_endpoint: str = "localhost:4317" otlp_traces_endpoint: str = "localhost:4317"
@@ -5236,6 +5239,25 @@ class ServerArgs:
default=None, default=None,
help="Config in json format for NVIDIA dynamo KV event publishing. Publishing will be enabled if this flag is used.", 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( parser.add_argument(
"--enable-trace", "--enable-trace",
action="store_true", action="store_true",
+8 -4
View File
@@ -1,7 +1,7 @@
from collections import deque from collections import deque
from contextlib import contextmanager from contextlib import contextmanager
from dataclasses import dataclass from dataclasses import dataclass
from typing import Callable, Deque, Dict, Optional from typing import Callable, Deque, Dict, List, Optional
import torch import torch
@@ -9,7 +9,10 @@ import torch
class DeviceTimer: class DeviceTimer:
def __init__(self, reporter: Callable): def __init__(self, reporter: Callable):
self._intervals: Deque[_TimingInterval] = deque() self._intervals: Deque[_TimingInterval] = deque()
self._reporter = reporter self._reporters: List[Callable] = [reporter]
def add_reporter(self, reporter: Callable):
self._reporters.append(reporter)
@contextmanager @contextmanager
def wrap(self, metadata: Dict): def wrap(self, metadata: Dict):
@@ -27,8 +30,9 @@ class DeviceTimer:
break break
self._intervals.popleft() self._intervals.popleft()
self._reporter(t=interval.elapsed_time() / 1000.0, **interval.metadata) elapsed = interval.elapsed_time() / 1000.0
# print(f"{interval.elapsed_time()=:.6f}, {interval.metadata=}") for reporter in self._reporters:
reporter(t=elapsed, **interval.metadata)
class GapTimer(DeviceTimer): class GapTimer(DeviceTimer):
+191
View File
@@ -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!")
@@ -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()