Fix metrics and request tracing (TimeStats) (#11123)
This commit is contained in:
@@ -14,9 +14,9 @@
|
||||
"""Utilities for Prometheus Metrics Collection."""
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Dict, List, Optional, Union
|
||||
|
||||
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||
from sglang.srt.metrics.utils import exponential_buckets, generate_buckets
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils import get_bool_env_var
|
||||
@@ -34,6 +34,7 @@ class TimeStats:
|
||||
Decode: prealloc_queue -> transfer_queue -> wait_queue -> forward -> completion
|
||||
"""
|
||||
|
||||
disagg_mode: DisaggregationMode = DisaggregationMode.NULL
|
||||
lb_entry_time: float = 0.0
|
||||
wait_queue_entry_time: float = 0.0
|
||||
forward_entry_time: float = 0.0
|
||||
@@ -43,20 +44,11 @@ class TimeStats:
|
||||
decode_prealloc_queue_entry_time: float = 0.0
|
||||
decode_transfer_queue_entry_time: float = 0.0
|
||||
|
||||
class RequestType(Enum):
|
||||
UNIFIED = "unified"
|
||||
PREFILL = "prefill"
|
||||
DECODE = "decode"
|
||||
INVALID = "invalid"
|
||||
|
||||
def get_queueing_time(self) -> float:
|
||||
return self.forward_entry_time - self.wait_queue_entry_time
|
||||
|
||||
def __str__(self) -> str:
|
||||
# if unified
|
||||
_type = self.get_type()
|
||||
|
||||
if _type == self.RequestType.UNIFIED:
|
||||
def convert_to_duration(self) -> str:
|
||||
if self.disagg_mode == DisaggregationMode.NULL:
|
||||
queue_duration = self.forward_entry_time - self.wait_queue_entry_time
|
||||
forward_duration = self.completion_time - self.forward_entry_time
|
||||
|
||||
@@ -65,30 +57,28 @@ class TimeStats:
|
||||
queue_duration >= 0 and forward_duration >= 0
|
||||
), f"queue_duration={queue_duration} < 0 or forward_duration={forward_duration} < 0"
|
||||
|
||||
return f"queue_duration={self.format_duration(queue_duration)}, forward_duration={self.format_duration(forward_duration)}, start_time={self.wait_queue_entry_time}"
|
||||
elif _type == self.RequestType.PREFILL:
|
||||
return f"queue_duration={self.format_duration(queue_duration)}, forward_duration={self.format_duration(forward_duration)}, start_time={self.wait_queue_entry_time:.3f}"
|
||||
elif self.disagg_mode == DisaggregationMode.PREFILL:
|
||||
bootstrap_duration = (
|
||||
self.wait_queue_entry_time - self.prefill_bootstrap_queue_entry_time
|
||||
)
|
||||
|
||||
queue_duration = self.forward_entry_time - self.wait_queue_entry_time
|
||||
|
||||
forward_duration = self.completion_time - self.forward_entry_time
|
||||
|
||||
if SGLANG_TEST_REQUEST_TIME_STATS:
|
||||
assert (
|
||||
bootstrap_duration >= 0
|
||||
and queue_duration >= 0
|
||||
and forward_duration >= 0
|
||||
), f"bootstrap_duration={bootstrap_duration} < 0 or queue_duration={queue_duration} < 0 or forward_duration={forward_duration} < 0"
|
||||
return f"bootstrap_duration={self.format_duration(bootstrap_duration)}, queue_duration={self.format_duration(queue_duration)}, forward_duration={self.format_duration(forward_duration)}, start_time={self.prefill_bootstrap_queue_entry_time}"
|
||||
# if decode
|
||||
elif _type == self.RequestType.DECODE:
|
||||
if self.wait_queue_entry_time > 0:
|
||||
assert (
|
||||
bootstrap_duration >= 0
|
||||
and queue_duration >= 0
|
||||
and forward_duration >= 0
|
||||
), f"bootstrap_duration={bootstrap_duration} < 0 or queue_duration={queue_duration} < 0 or forward_duration={forward_duration} < 0"
|
||||
|
||||
return f"bootstrap_duration={self.format_duration(bootstrap_duration)}, queue_duration={self.format_duration(queue_duration)}, forward_duration={self.format_duration(forward_duration)}, start_time={self.prefill_bootstrap_queue_entry_time:.3f}"
|
||||
elif self.disagg_mode == DisaggregationMode.DECODE:
|
||||
prealloc_duration = (
|
||||
self.decode_transfer_queue_entry_time
|
||||
- self.decode_prealloc_queue_entry_time
|
||||
)
|
||||
|
||||
transfer_duration = (
|
||||
self.wait_queue_entry_time - self.decode_transfer_queue_entry_time
|
||||
)
|
||||
@@ -96,42 +86,30 @@ class TimeStats:
|
||||
forward_duration = self.completion_time - self.forward_entry_time
|
||||
|
||||
if SGLANG_TEST_REQUEST_TIME_STATS:
|
||||
assert (
|
||||
prealloc_duration >= 0
|
||||
and transfer_duration >= 0
|
||||
and queue_duration >= 0
|
||||
and forward_duration >= 0
|
||||
), f"prealloc_duration={prealloc_duration} < 0 or transfer_duration={transfer_duration} < 0 or queue_duration={queue_duration} < 0 or forward_duration={forward_duration} < 0"
|
||||
if self.wait_queue_entry_time > 0:
|
||||
assert (
|
||||
prealloc_duration >= 0
|
||||
and transfer_duration >= 0
|
||||
and queue_duration >= 0
|
||||
and forward_duration >= 0
|
||||
), f"prealloc_duration={prealloc_duration} < 0 or transfer_duration={transfer_duration} < 0 or queue_duration={queue_duration} < 0 or forward_duration={forward_duration} < 0. {self=}"
|
||||
|
||||
return f"prealloc_duration={self.format_duration(prealloc_duration)}, transfer_duration={self.format_duration(transfer_duration)}, queue_duration={self.format_duration(queue_duration)}, forward_duration={self.format_duration(forward_duration)}, start_time={self.decode_prealloc_queue_entry_time}"
|
||||
return f"prealloc_duration={self.format_duration(prealloc_duration)}, transfer_duration={self.format_duration(transfer_duration)}, queue_duration={self.format_duration(queue_duration)}, forward_duration={self.format_duration(forward_duration)}, start_time={self.decode_prealloc_queue_entry_time:.3f}"
|
||||
else:
|
||||
return "Invalid Time Stats"
|
||||
return "Unknown Time Stats"
|
||||
|
||||
def format_duration(self, duration: float) -> str:
|
||||
return f"{duration * 1e3:.2f}ms"
|
||||
|
||||
def get_type(self) -> RequestType:
|
||||
"""Determine the type of request based on timestamp values."""
|
||||
if (
|
||||
self.prefill_bootstrap_queue_entry_time == 0.0
|
||||
and self.prefill_transfer_queue_entry_time == 0.0
|
||||
and self.decode_prealloc_queue_entry_time == 0.0
|
||||
and self.decode_transfer_queue_entry_time == 0.0
|
||||
):
|
||||
return self.RequestType.UNIFIED
|
||||
elif (
|
||||
self.prefill_bootstrap_queue_entry_time > 0.0
|
||||
and self.prefill_transfer_queue_entry_time > 0.0
|
||||
):
|
||||
return self.RequestType.PREFILL
|
||||
elif (
|
||||
self.decode_prealloc_queue_entry_time > 0.0
|
||||
and self.decode_transfer_queue_entry_time > 0.0
|
||||
and self.wait_queue_entry_time > 0.0
|
||||
):
|
||||
return self.RequestType.DECODE
|
||||
def disagg_mode_str(self) -> str:
|
||||
if self.disagg_mode == DisaggregationMode.NULL:
|
||||
return "unified"
|
||||
elif self.disagg_mode == DisaggregationMode.DECODE:
|
||||
return "decode"
|
||||
elif self.disagg_mode == DisaggregationMode.PREFILL:
|
||||
return "prefill"
|
||||
else:
|
||||
return self.RequestType.INVALID
|
||||
return "unknown"
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -145,12 +123,15 @@ class SchedulerStats:
|
||||
num_queue_reqs: int = 0
|
||||
num_grammar_queue_reqs: int = 0
|
||||
num_running_reqs_offline_batch: int = 0
|
||||
avg_request_queue_latency: float = 0.0
|
||||
cache_hit_rate: float = 0.0
|
||||
|
||||
# Speculative decoding
|
||||
spec_accept_length: float = 0.0
|
||||
|
||||
# Retract
|
||||
num_retracted_reqs: int = 0
|
||||
num_paused_reqs: int = 0
|
||||
|
||||
# PD disaggregation
|
||||
num_prefill_prealloc_queue_reqs: int = 0
|
||||
num_prefill_inflight_queue_reqs: int = 0
|
||||
@@ -159,11 +140,6 @@ class SchedulerStats:
|
||||
kv_transfer_speed_gb_s: float = 0.0
|
||||
kv_transfer_latency_ms: float = 0.0
|
||||
|
||||
# Retract
|
||||
total_retracted_reqs: int = 0
|
||||
num_retracted_reqs: int = 0
|
||||
num_paused_reqs: int = 0
|
||||
|
||||
# Utilization
|
||||
utilization: float = 0.0
|
||||
max_running_requests_under_SLO: Optional[int] = None
|
||||
@@ -230,12 +206,6 @@ class SchedulerMetricsCollector:
|
||||
labelnames=labels.keys(),
|
||||
multiprocess_mode="mostrecent",
|
||||
)
|
||||
self.avg_request_queue_latency = Gauge(
|
||||
name="sglang:avg_request_queue_latency",
|
||||
documentation="The average request queue latency for the last batch of requests in seconds.",
|
||||
labelnames=labels.keys(),
|
||||
multiprocess_mode="mostrecent",
|
||||
)
|
||||
self.cache_hit_rate = Gauge(
|
||||
name="sglang:cache_hit_rate",
|
||||
documentation="The prefix cache hit rate.",
|
||||
@@ -251,6 +221,18 @@ class SchedulerMetricsCollector:
|
||||
multiprocess_mode="mostrecent",
|
||||
)
|
||||
|
||||
# Retract
|
||||
self.num_retracted_reqs = Gauge(
|
||||
name="sglang:num_retracted_reqs",
|
||||
documentation="The number of retracted requests.",
|
||||
labelnames=labels.keys(),
|
||||
)
|
||||
self.num_paused_reqs = Gauge(
|
||||
name="sglang:num_paused_reqs",
|
||||
documentation="The number of paused requests by async weight sync.",
|
||||
labelnames=labels.keys(),
|
||||
)
|
||||
|
||||
# PD disaggregation
|
||||
self.num_prefill_prealloc_queue_reqs = Gauge(
|
||||
name="sglang:num_prefill_prealloc_queue_reqs",
|
||||
@@ -299,24 +281,6 @@ class SchedulerMetricsCollector:
|
||||
multiprocess_mode="mostrecent",
|
||||
)
|
||||
|
||||
# Retract
|
||||
self.total_retracted_reqs = Gauge(
|
||||
name="sglang:total_retracted_reqs",
|
||||
documentation="The total number of retracted requests due to kvcache full.",
|
||||
labelnames=labels.keys(),
|
||||
multiprocess_mode="mostrecent",
|
||||
)
|
||||
self.num_retracted_reqs = Gauge(
|
||||
name="sglang:num_retracted_reqs",
|
||||
documentation="The number of retracted requests.",
|
||||
labelnames=labels.keys(),
|
||||
)
|
||||
self.num_paused_reqs = Gauge(
|
||||
name="sglang:num_paused_reqs",
|
||||
documentation="The number of paused requests by async weight sync.",
|
||||
labelnames=labels.keys(),
|
||||
)
|
||||
|
||||
# Utilization
|
||||
self.utilization = Gauge(
|
||||
name="sglang:utilization",
|
||||
@@ -347,7 +311,7 @@ class SchedulerMetricsCollector:
|
||||
|
||||
# Additional queueing time histogram
|
||||
self.queue_time = Histogram(
|
||||
name="sglang:queue_time_s",
|
||||
name="sglang:queue_time_seconds",
|
||||
documentation="Histogram of queueing time in seconds.",
|
||||
labelnames=labels.keys(),
|
||||
buckets=[
|
||||
@@ -513,8 +477,8 @@ class SchedulerMetricsCollector:
|
||||
buckets=tree_traversal_time_buckets,
|
||||
)
|
||||
|
||||
self.request_latency_seconds = Histogram(
|
||||
name="sglang:request_latency_seconds",
|
||||
self.per_stage_req_latency_seconds = Histogram(
|
||||
name="sglang:per_stage_req_latency_seconds",
|
||||
documentation="The latency of each stage of requests.",
|
||||
# captures latency in range [1ms - ~1191s]
|
||||
buckets=exponential_buckets(start=0.001, width=1.62, length=30),
|
||||
@@ -525,7 +489,7 @@ class SchedulerMetricsCollector:
|
||||
# Convenience function for logging to gauge.
|
||||
gauge.labels(**self.labels).set(data)
|
||||
|
||||
def log_histogram(self, histogram, data: Union[int, float]) -> None:
|
||||
def _log_histogram(self, histogram, data: Union[int, float]) -> None:
|
||||
histogram.labels(**self.labels).observe(data)
|
||||
|
||||
def increment_bootstrap_failed_reqs(self) -> None:
|
||||
@@ -534,9 +498,12 @@ class SchedulerMetricsCollector:
|
||||
def increment_transfer_failed_reqs(self) -> None:
|
||||
self.num_transfer_failed_reqs.labels(**self.labels).inc(1)
|
||||
|
||||
def observe_request_latency_seconds(self, stage: str, latency: float) -> None:
|
||||
def observe_per_stage_req_latency(self, stage: str, latency: float) -> None:
|
||||
labels_with_stage = {**self.labels, "stage": stage}
|
||||
self.request_latency_seconds.labels(**labels_with_stage).observe(latency)
|
||||
self.per_stage_req_latency_seconds.labels(**labels_with_stage).observe(latency)
|
||||
|
||||
def observe_queue_time(self, latency: float) -> None:
|
||||
self._log_histogram(self.queue_time, latency)
|
||||
|
||||
def log_stats(self, stats: SchedulerStats) -> None:
|
||||
self._log_gauge(self.num_running_reqs, stats.num_running_reqs)
|
||||
@@ -550,7 +517,6 @@ class SchedulerMetricsCollector:
|
||||
self.num_running_reqs_offline_batch, stats.num_running_reqs_offline_batch
|
||||
)
|
||||
self._log_gauge(self.cache_hit_rate, stats.cache_hit_rate)
|
||||
self._log_gauge(self.avg_request_queue_latency, stats.avg_request_queue_latency)
|
||||
|
||||
# Speculative decoding
|
||||
self._log_gauge(self.spec_accept_length, stats.spec_accept_length)
|
||||
@@ -572,7 +538,6 @@ class SchedulerMetricsCollector:
|
||||
self._log_gauge(self.kv_transfer_latency_ms, stats.kv_transfer_latency_ms)
|
||||
|
||||
# Retract
|
||||
self._log_gauge(self.total_retracted_reqs, stats.total_retracted_reqs)
|
||||
self._log_gauge(self.num_retracted_reqs, stats.num_retracted_reqs)
|
||||
self._log_gauge(self.num_paused_reqs, stats.num_paused_reqs)
|
||||
|
||||
@@ -596,19 +561,19 @@ class SchedulerMetricsCollector:
|
||||
def log_grammar_stats(self, grammar_stats) -> None:
|
||||
# Duck-typed GrammarStats to avoid cross-package dependency
|
||||
if getattr(grammar_stats, "compilation_time", None) is not None:
|
||||
self.log_histogram(
|
||||
self._log_histogram(
|
||||
self.grammar_compilation_time, grammar_stats.compilation_time
|
||||
)
|
||||
if getattr(grammar_stats, "schema_count", None) is not None:
|
||||
self.log_histogram(self.grammar_schema_count, grammar_stats.schema_count)
|
||||
self._log_histogram(self.grammar_schema_count, grammar_stats.schema_count)
|
||||
if getattr(grammar_stats, "ebnf_size", None) is not None:
|
||||
self.log_histogram(self.grammar_ebnf_size, grammar_stats.ebnf_size)
|
||||
self._log_histogram(self.grammar_ebnf_size, grammar_stats.ebnf_size)
|
||||
tree_times = getattr(grammar_stats, "tree_traversal_time", None)
|
||||
if tree_times:
|
||||
max_time = max(tree_times)
|
||||
avg_time = sum(tree_times) / len(tree_times)
|
||||
self.log_histogram(self.grammar_tree_traversal_time_max, max_time)
|
||||
self.log_histogram(self.grammar_tree_traversal_time_avg, avg_time)
|
||||
self._log_histogram(self.grammar_tree_traversal_time_max, max_time)
|
||||
self._log_histogram(self.grammar_tree_traversal_time_avg, avg_time)
|
||||
if getattr(grammar_stats, "is_cache_hit", False):
|
||||
self.num_grammar_cache_hit.labels(**self.labels).inc(1)
|
||||
if getattr(grammar_stats, "is_grammar_aborted", False):
|
||||
@@ -714,7 +679,7 @@ class TokenizerMetricsCollector:
|
||||
)
|
||||
|
||||
self.num_aborted_requests_total = Counter(
|
||||
name="sglang:num_aborted_requests",
|
||||
name="sglang:num_aborted_requests_total",
|
||||
documentation="Number of requests aborted.",
|
||||
labelnames=labels.keys(),
|
||||
)
|
||||
@@ -801,7 +766,7 @@ class TokenizerMetricsCollector:
|
||||
buckets=bucket_time_to_first_token,
|
||||
)
|
||||
|
||||
self.histogram_inter_token_latency_seconds = Histogram(
|
||||
self.histogram_inter_token_latency = Histogram(
|
||||
name="sglang:inter_token_latency_seconds",
|
||||
documentation="Histogram of inter-token latency in seconds.",
|
||||
labelnames=labels.keys(),
|
||||
@@ -815,14 +780,6 @@ class TokenizerMetricsCollector:
|
||||
buckets=bucket_e2e_request_latency,
|
||||
)
|
||||
|
||||
# Offline batch specific TTFB histogram
|
||||
self.histogram_time_to_first_token_offline_batch = Histogram(
|
||||
name="sglang:time_to_first_token_seconds_offline_batch",
|
||||
documentation="Histogram of time to first token in seconds for offline batch requests.",
|
||||
labelnames=labels.keys(),
|
||||
buckets=bucket_time_to_first_token,
|
||||
)
|
||||
|
||||
def observe_one_finished_request(
|
||||
self,
|
||||
labels: Dict[str, str],
|
||||
@@ -846,15 +803,8 @@ class TokenizerMetricsCollector:
|
||||
float(generation_tokens)
|
||||
)
|
||||
|
||||
def observe_time_to_first_token(
|
||||
self, labels: Dict[str, str], value: float, type: str = ""
|
||||
):
|
||||
if type == "batch":
|
||||
self.histogram_time_to_first_token_offline_batch.labels(**labels).observe(
|
||||
value
|
||||
)
|
||||
else:
|
||||
self.histogram_time_to_first_token.labels(**labels).observe(value)
|
||||
def observe_time_to_first_token(self, labels: Dict[str, str], value: float):
|
||||
self.histogram_time_to_first_token.labels(**labels).observe(value)
|
||||
|
||||
def check_time_to_first_token_straggler(self, value: float) -> bool:
|
||||
his = self.histogram_time_to_first_token.labels(**self.labels)
|
||||
@@ -876,7 +826,7 @@ class TokenizerMetricsCollector:
|
||||
|
||||
# A faster version of the Histogram::observe which observes multiple values at the same time.
|
||||
# reference: https://github.com/prometheus/client_python/blob/v0.21.1/prometheus_client/metrics.py#L639
|
||||
his = self.histogram_inter_token_latency_seconds.labels(**labels)
|
||||
his = self.histogram_inter_token_latency.labels(**labels)
|
||||
his._sum.inc(internval)
|
||||
|
||||
for i, bound in enumerate(his._upper_bounds):
|
||||
@@ -884,8 +834,8 @@ class TokenizerMetricsCollector:
|
||||
his._buckets[i].inc(num_new_tokens)
|
||||
break
|
||||
|
||||
def observe_one_aborted_request(self):
|
||||
self.num_aborted_requests_total.labels(**self.labels).inc(1)
|
||||
def observe_one_aborted_request(self, labels: Dict[str, str]):
|
||||
self.num_aborted_requests_total.labels(**labels).inc(1)
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
Reference in New Issue
Block a user