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:
co-authored by
Ishan Dhanani
Claude Opus 4.6
ishandhanani
parent
fd3eb77d45
commit
e86fb42736
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 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.
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user