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
+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()