From fd97fbb096bd70b925abc12e59e2fc3346e49563 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Mon, 18 May 2026 18:42:08 +0800 Subject: [PATCH] Move metrics reporting to SchedulerMetricsReporter and retire metrics mixin (#25630) --- python/sglang/srt/disaggregation/prefill.py | 3 +- python/sglang/srt/dllm/mixin/scheduler.py | 7 +- python/sglang/srt/managers/schedule_batch.py | 2 +- python/sglang/srt/managers/scheduler.py | 20 +- .../scheduler_components/metrics_reporter.py | 932 ++++++++++++++++- .../scheduler_output_processor_mixin.py | 10 +- .../observability/scheduler_metrics_mixin.py | 961 ------------------ .../test_forward_pass_metrics.py | 81 +- 8 files changed, 1001 insertions(+), 1015 deletions(-) delete mode 100644 python/sglang/srt/observability/scheduler_metrics_mixin.py diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 72f0450d8..732e78453 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -587,8 +587,7 @@ class SchedulerDisaggregationPrefillMixin: req.time_stats.set_last_chunked_prefill_finish_time() can_run_cuda_graph = getattr(result, "can_run_cuda_graph", False) - self.report_prefill_stats( - self.metrics_reporter, + self.metrics_reporter.report_prefill_stats( batch=batch, prefill_stats=batch.prefill_stats, can_run_cuda_graph=can_run_cuda_graph, diff --git a/python/sglang/srt/dllm/mixin/scheduler.py b/python/sglang/srt/dllm/mixin/scheduler.py index 8179674bb..39ea06a99 100644 --- a/python/sglang/srt/dllm/mixin/scheduler.py +++ b/python/sglang/srt/dllm/mixin/scheduler.py @@ -93,8 +93,7 @@ class SchedulerDllmMixin: self.token_to_kv_pool_allocator.free_group_end() can_run_cuda_graph = getattr(result, "can_run_cuda_graph", False) - self.report_prefill_stats( - self.metrics_reporter, + self.metrics_reporter.report_prefill_stats( batch=batch, prefill_stats=batch.prefill_stats, can_run_cuda_graph=can_run_cuda_graph, @@ -226,7 +225,9 @@ class SchedulerDllmMixin: new_batch.decoding_reqs = None # Record prefill stats for logging after forward - from sglang.srt.observability.scheduler_metrics_mixin import PrefillStats + from sglang.srt.managers.scheduler_components.metrics_reporter import ( + PrefillStats, + ) new_batch.prefill_stats = PrefillStats.from_adder( self.adder, self.running_batch.reqs, self.enable_priority_scheduling diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 0b6d7b3e9..33da2b1e3 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -96,7 +96,7 @@ if TYPE_CHECKING: from sglang.srt.configs.model_config import ModelConfig from sglang.srt.managers.hisparse_coordinator import HiSparseCoordinator - from sglang.srt.observability.scheduler_metrics_mixin import PrefillStats + from sglang.srt.managers.scheduler_components.metrics_reporter import PrefillStats from sglang.srt.session.session_controller import Session from sglang.srt.speculative.eagle_info import EagleDraftInput from sglang.srt.speculative.spec_info import SpecInput, SpeculativeAlgorithm diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 420885ae0..5712ef1c9 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -177,6 +177,8 @@ from sglang.srt.managers.scheduler_components.load_inquirer import ( SchedulerLoadInquirer, ) from sglang.srt.managers.scheduler_components.metrics_reporter import ( + RECORD_STEP_TIME, + PrefillStats, SchedulerMetricsReporter, ) from sglang.srt.managers.scheduler_components.pool_stats_observer import ( @@ -212,11 +214,6 @@ from sglang.srt.observability.req_time_stats import ( set_schedule_time_batch, set_time_batch, ) -from sglang.srt.observability.scheduler_metrics_mixin import ( - RECORD_STEP_TIME, - PrefillStats, - SchedulerMetricsMixin, -) from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info from sglang.srt.parser.reasoning_parser import ReasoningParser from sglang.srt.platforms import current_platform @@ -361,7 +358,6 @@ def create_scheduler_watchdog( class Scheduler( SchedulerOutputProcessorMixin, - SchedulerMetricsMixin, SchedulerDisaggregationDecodeMixin, SchedulerDisaggregationPrefillMixin, SchedulerMultiplexMixin, @@ -3067,15 +3063,15 @@ class Scheduler( elif batch.forward_mode.is_idle(): self.process_batch_result_idle(batch, result) - self.log_batch_result_stats(self.metrics_reporter, batch, result) + self.metrics_reporter.log_batch_result_stats(batch, result) # Emit forward pass metrics (every iteration when enabled) if self.enable_fpm: - self._emit_forward_pass_metrics(self.metrics_reporter, batch, result) + self.metrics_reporter._emit_forward_pass_metrics(batch, result) self._maybe_clear_mm_inputs(batch) self.maybe_send_health_check_signal() - self.update_device_timer(self.metrics_reporter) + self.metrics_reporter.update_device_timer() def maybe_send_health_check_signal(self): if self.return_health_check_ipcs: @@ -3201,7 +3197,7 @@ class Scheduler( self.new_token_ratio = self.init_new_token_ratio # reset device timer window so idle time isn't counted - self.reset_device_timer_window(self.metrics_reporter) + self.metrics_reporter.reset_device_timer_window() # sleep until next event self.maybe_sleep_on_idle() @@ -3362,7 +3358,7 @@ class Scheduler( self.req_to_token_pool.clear() self.token_to_kv_pool_allocator.clear() self.grammar_manager.clear() - self.reset_metrics(self.metrics_reporter) + self.metrics_reporter.reset_metrics() if self.draft_worker: self.draft_worker.clear_cache_pool() @@ -3981,4 +3977,4 @@ def run_scheduler_process( 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(scheduler.metrics_reporter) + scheduler.metrics_reporter._shutdown_fpm() diff --git a/python/sglang/srt/managers/scheduler_components/metrics_reporter.py b/python/sglang/srt/managers/scheduler_components/metrics_reporter.py index bbe17d101..3f05bac45 100644 --- a/python/sglang/srt/managers/scheduler_components/metrics_reporter.py +++ b/python/sglang/srt/managers/scheduler_components/metrics_reporter.py @@ -1,24 +1,82 @@ from __future__ import annotations +import dataclasses import logging +import tempfile +import time +from collections import defaultdict from dataclasses import dataclass -from typing import TYPE_CHECKING, Optional +from typing import ( + TYPE_CHECKING, + List, + Optional, + Tuple, + Union, +) +from sglang.srt.disaggregation.utils import DisaggregationMode +from sglang.srt.environ import envs +from sglang.srt.managers.schedule_batch import ScheduleBatch +from sglang.srt.managers.utils import GenerationBatchResult from sglang.srt.observability.metrics_collector import ( + DPCooperationInfo, + QueueCount, SchedulerMetricsCollector, SchedulerMetricsCollectorContext, + SchedulerStats, + compute_routing_key_stats, ) -from sglang.srt.observability.scheduler_metrics_mixin import ( - SchedulerMetricsMixin, -) +from sglang.srt.utils.device_timer import DeviceTimer +from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger if TYPE_CHECKING: - from sglang.srt.managers.scheduler import Scheduler + from sglang.srt.managers.schedule_batch import Req + from sglang.srt.managers.schedule_policy import PrefillAdder + from sglang.srt.managers.scheduler import ( + EmbeddingBatchResult, + Scheduler, + ) logger = logging.getLogger(__name__) +RECORD_STEP_TIME = envs.SGLANG_RECORD_STEP_TIME.get() +LOG_FORWARD_ITERS = envs.SGLANG_LOG_FORWARD_ITERS.get() +ENABLE_METRICS_DEVICE_TIMER = envs.SGLANG_ENABLE_METRICS_DEVICE_TIMER.get() + + +@dataclasses.dataclass +class PrefillStats: + """Stats for logging prefill batch metrics.""" + + log_input_tokens: int + log_hit_tokens: int + new_token_ratio: float + num_running_reqs: QueueCount + num_new_seqs: int # len(can_run_list) + num_pending_tokens: int = 0 + + @classmethod + def from_adder( + cls, + adder: PrefillAdder, + running_reqs: List[Req], + enable_priority_scheduling: bool = False, + num_pending_tokens: int = 0, + ): + return cls( + log_input_tokens=adder.log_input_tokens, + log_hit_tokens=adder.log_hit_tokens, + new_token_ratio=adder.new_token_ratio, + num_running_reqs=QueueCount.from_reqs( + running_reqs, enable_priority_scheduling + ), + num_new_seqs=len(adder.can_run_list), + num_pending_tokens=num_pending_tokens, + ) + + @dataclass(kw_only=True) class SchedulerMetricsReporter: scheduler: "Scheduler" @@ -41,7 +99,865 @@ class SchedulerMetricsReporter: self.enable_kv_cache_events = ( self.metrics_collector_context.enable_kv_cache_events ) - SchedulerMetricsMixin._init_metrics( - self, self.tp_rank, self.pp_rank, self.dp_rank + self._init_metrics(self.tp_rank, self.pp_rank, self.dp_rank) + self._install_device_timer_on_runners() + + def _init_metrics( + self, + tp_rank: int, + pp_rank: int, + dp_rank: Optional[int], + ): + # Basic stats + self.forward_ct_decode = 0 + self.num_generated_tokens = 0 + self.last_decode_stats_tic = time.perf_counter() + self.last_prefill_stats_tic = time.perf_counter() + self.last_gen_throughput: float = 0.0 + self.last_input_throughput: float = 0.0 + self.step_time_dict = defaultdict(list) # Dict[batch size -> step time] + self.stats = SchedulerStats() + self._graph_backend_label = { + "cpu": "cpu graph", + "npu": "npu graph", + "musa": "musa graph", + }.get(getattr(self.scheduler, "device", ""), "cuda graph") + + # Cumulative spec-decoding counters (reset every decode_log_interval). + # Each update adds (num_correct_drafts + bs, bs). + # `*_accept_tokens` = drafts + bonus; `*_correct_drafts` = drafts-only. + self.spec_num_accept_tokens = 0 # per-log-interval + self.spec_num_forward_ct = 0 + self.spec_total_num_accept_tokens = 0 # lifetime + self.spec_total_num_forward_ct = 0 + + # For PD disaggregation + self.kv_transfer_speed_gb_s: float = 0.0 + self.kv_transfer_latency_ms: float = 0.0 + + self.enable_mfu_metrics = False + + if self.enable_metrics: + self.enable_mfu_metrics = self.scheduler.server_args.enable_mfu_metrics + if self.enable_mfu_metrics: + self._init_estimated_perf_constants() + self._mfu_log_flops = 0.0 + self._mfu_log_read_bytes = 0.0 + self._mfu_log_write_bytes = 0.0 + + self.fwd_occupancy = float("nan") + + self.forward_pass_device_timer: Optional[DeviceTimer] = None + + if ENABLE_METRICS_DEVICE_TIMER: + self._device_timer_window_batch_count = 0 + self._device_timer_window_gpu_time = 0.0 + self._device_timer_window_start = None + + def _wrap_execution_reporter(**kwargs): + self._device_timer_window_gpu_time += kwargs["t"] + if self.enable_metrics: + self.metrics_collector.increment_forward_execution_seconds(**kwargs) + + self.forward_pass_device_timer = DeviceTimer( + reporter=_wrap_execution_reporter, + ) + + self._init_fpm() + + self.scheduler_status_logger = SchedulerStatusLogger.maybe_create( + enable_metrics=self.enable_metrics ) - SchedulerMetricsMixin._install_device_timer_on_runners(self) + + def _install_device_timer_on_runners(self): + if self.forward_pass_device_timer is None: + return + timer = self.forward_pass_device_timer + self.scheduler.tp_worker.model_runner.device_timer = timer + if self.scheduler.draft_worker is not None: + dw = getattr(self.scheduler.draft_worker, "draft_worker", None) + if dw is not None: + if hasattr(dw, "draft_runner"): + dw.draft_runner.device_timer = timer + for r in getattr(dw, "draft_runner_list", []): + r.device_timer = timer + + def _init_fpm(self): + """Initialize Forward Pass Metrics (FPM) publisher if configured.""" + self.scheduler.enable_fpm = False + if ( + self.scheduler.server_args.enable_forward_pass_metrics + and self.scheduler.ps.attn_tp_rank == 0 + and self.scheduler.ps.pp_rank == self.scheduler.ps.pp_size - 1 + ): + from sglang.srt.observability.forward_pass_metrics import ( + _FpmPublisherThread, + ) + + self.scheduler._fpm_dp_rank = ( + self.scheduler.ps.dp_rank + if self.scheduler.ps.dp_rank is not None + else 0 + ) + self.scheduler._fpm_worker_id = ( + self.scheduler.server_args.forward_pass_metrics_worker_id + ) + base_endpoint = self.scheduler.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.scheduler.server_args.forward_pass_metrics_ipc_name = base_endpoint + endpoint = f"{base_endpoint}.{self.scheduler._fpm_dp_rank}" + self.scheduler._fpm_publisher = _FpmPublisherThread( + endpoint, + worker_id=self.scheduler._fpm_worker_id, + dp_rank=self.scheduler._fpm_dp_rank, + ) + self.scheduler._fpm_gpu_time_acc = 0.0 + + def _fpm_device_timer_reporter(t, **_kwargs): + self.scheduler._fpm_gpu_time_acc += t + + if self.forward_pass_device_timer is not None: + self.forward_pass_device_timer.add_reporter(_fpm_device_timer_reporter) + else: + self.forward_pass_device_timer = DeviceTimer( + reporter=_fpm_device_timer_reporter, + ) + self.scheduler._fpm_uses_device_timer = True + self.scheduler.enable_fpm = True + logger.info( + "FPM: ZMQ PUB bound on %s (dp_rank=%d, device_timer=%s)", + endpoint, + self.scheduler._fpm_dp_rank, + self.scheduler._fpm_uses_device_timer, + ) + + def _build_scheduled_request_metrics(self, 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): + from sglang.srt.observability.forward_pass_metrics import ( + QueuedRequestMetrics, + WelfordAccumulator, + ) + + prefill_q = WelfordAccumulator() + decode_q = WelfordAccumulator() + if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL: + for req in self.scheduler.disagg_prefill_bootstrap_queue.queue: + prefill_q.add(len(req.origin_input_ids)) + elif self.scheduler.disaggregation_mode == DisaggregationMode.DECODE: + for req in self.scheduler.disagg_decode_prealloc_queue.queue: + decode_q.add(req.seqlen) + for req in self.scheduler.disagg_decode_transfer_queue.queue: + decode_q.add(req.seqlen) + else: + for req in self.scheduler.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, bs: int, num_correct_drafts: int): + self.spec_num_accept_tokens += num_correct_drafts + bs + self.spec_num_forward_ct += bs + + # Bonus tokens updated elsewhere + self.num_generated_tokens += num_correct_drafts + + def _init_estimated_perf_constants(self) -> None: + model_config = self.scheduler.model_config + hf_text_config = model_config.hf_text_config + + hidden_size = float(model_config.hidden_size) + num_layers = float(getattr(model_config, "num_attention_layers", 0)) + head_dim = float(getattr(model_config, "head_dim", 0)) + num_attn_heads = float( + model_config.get_num_attention_heads(self.scheduler.ps.tp_size) + ) + num_kv_heads = float(model_config.get_num_kv_heads(self.scheduler.ps.tp_size)) + intermediate_size = getattr(hf_text_config, "intermediate_size", None) + if intermediate_size is None: + intermediate_size = getattr(hf_text_config, "ffn_hidden_size", 0) + intermediate_size = float(intermediate_size) + + dtype_num_bytes = getattr(model_config.dtype, "itemsize", None) + if dtype_num_bytes is None: + dtype_num_bytes = 2 + # Keep this estimator lightweight and consistent with current server dtype. + # KV cache quantization-aware bytes can be added in a follow-up. + act_bytes = float(dtype_num_bytes) + w_bytes = float(dtype_num_bytes) + cache_bytes = float(dtype_num_bytes) + + # Linear-layer FLOPs per token on one GPU. + attn_linear_flops = ( + 2.0 * hidden_size * head_dim * (num_attn_heads + 2.0 * num_kv_heads) + + 2.0 * hidden_size * head_dim * num_attn_heads + ) + mlp_flops = ( + 6.0 * hidden_size * intermediate_size if intermediate_size > 0 else 0.0 + ) + self._linear_flops_per_token = max( + 0.0, (attn_linear_flops + mlp_flops) * num_layers + ) + + # Attention dot-product FLOPs coefficient to multiply token-context product. + # attn_qk + attn_av = 4 * q * TC * d * L + self._attn_dot_flops_coeff = 4.0 * num_attn_heads * head_dim * num_layers + + # KV cache bytes (write one K and one V vector per generated token). + self._kv_cache_bytes_per_token = ( + 2.0 * num_layers * num_kv_heads * head_dim * cache_bytes + ) + + # Weight read bytes per token. + self._weight_read_bytes_per_token = ( + hidden_size + * head_dim + * (num_attn_heads + 2.0 * num_kv_heads) + * w_bytes + * num_layers + + hidden_size * head_dim * num_attn_heads * w_bytes * num_layers + + ( + 3.0 * hidden_size * intermediate_size * w_bytes * num_layers + if intermediate_size > 0 + else 0.0 + ) + ) + + # Activation movement bytes per token (coarse approximation). + self._qkv_act_bytes_per_token = ( + hidden_size * act_bytes * num_layers + + (num_attn_heads + 2.0 * num_kv_heads) * head_dim * act_bytes * num_layers + + head_dim * num_attn_heads * act_bytes * num_layers + + hidden_size * act_bytes * num_layers + ) + self._ffn_act_bytes_per_token = ( + 3.0 * intermediate_size * act_bytes * num_layers + if intermediate_size > 0 + else 0.0 + ) + + # Prefill reads Q/K/V activations from on-device memory. + self._prefill_attn_act_read_per_token = ( + (num_attn_heads + 2.0 * num_kv_heads) * head_dim * act_bytes * num_layers + ) + + # Decode reads Q from activation memory; K/V reads are from KV cache. + self._decode_q_read_bytes_per_token = ( + num_attn_heads * head_dim * act_bytes * num_layers + ) + + def _estimate_prefill_perf(self, num_tokens: int) -> Tuple[float, float, float]: + tokens = max(0, int(num_tokens)) + if tokens == 0: + return 0.0, 0.0, 0.0 + + # Causal prefill token-context product. + context_product = tokens * (tokens + 1) / 2.0 + flops = ( + tokens * self._linear_flops_per_token + + self._attn_dot_flops_coeff * context_product + ) + + read_bytes = ( + tokens * self._weight_read_bytes_per_token + + tokens * self._qkv_act_bytes_per_token + + tokens * self._prefill_attn_act_read_per_token + ) + write_bytes = ( + tokens * self._kv_cache_bytes_per_token + + tokens * self._qkv_act_bytes_per_token + + tokens * self._ffn_act_bytes_per_token + ) + return flops, read_bytes, write_bytes + + def _estimate_decode_perf( + self, batch: ScheduleBatch, num_tokens: int + ) -> Tuple[float, float, float]: + tokens = max(0, int(num_tokens)) + if tokens == 0: + return 0.0, 0.0, 0.0 + + total_context = float(batch.seq_lens_cpu.sum().item()) + flops = ( + tokens * self._linear_flops_per_token + + self._attn_dot_flops_coeff * total_context + ) + read_bytes = ( + tokens * self._weight_read_bytes_per_token + + tokens * self._qkv_act_bytes_per_token + + tokens * self._decode_q_read_bytes_per_token + + total_context * self._kv_cache_bytes_per_token + ) + write_bytes = ( + tokens * self._kv_cache_bytes_per_token + + tokens * self._qkv_act_bytes_per_token + + tokens * self._ffn_act_bytes_per_token + ) + return flops, read_bytes, write_bytes + + def reset_metrics(self): + self.forward_ct_decode = 0 + self.num_generated_tokens = 0 + self.spec_num_accept_tokens = 0 + self.spec_num_forward_ct = 0 + self.spec_total_num_accept_tokens = 0 + self.spec_total_num_forward_ct = 0 + + def report_prefill_stats( + self, + batch: Optional[ScheduleBatch], + prefill_stats: PrefillStats, + can_run_cuda_graph: bool, + dp_cooperation_info: Optional[DPCooperationInfo] = None, + ): + if ( + not self.is_stats_logging_rank + and not self.current_scheduler_metrics_enabled + ): + return + + now = time.perf_counter() + gap_latency = now - self.last_prefill_stats_tic + self.last_prefill_stats_tic = now + self.last_input_throughput = ( + prefill_stats.log_input_tokens / gap_latency if gap_latency > 0 else 0.0 + ) + + pool_stats = self.scheduler.pool_stats_observer.get_pool_stats() + token_usage_msg = ", ".join(pool_stats.get_prefill_usage_msg_parts()) + ", " + + self.stats.new_token_ratio = prefill_stats.new_token_ratio + batch_iter = ( + batch.forward_iter + if batch is not None and batch.forward_iter is not None + else self.scheduler.forward_ct + ) + iter_msg = f" [{batch_iter}]" if LOG_FORWARD_ITERS else "" + + msg = ( + f"Prefill batch{iter_msg}, " + f"#new-seq: {prefill_stats.num_new_seqs}, " + f"#new-token: {prefill_stats.log_input_tokens}, " + f"#cached-token: {prefill_stats.log_hit_tokens}, " + f"{token_usage_msg}" + f"#running-req: {prefill_stats.num_running_reqs.total}, " + f"#queue-req: {len(self.scheduler.waiting_queue)}, " + f"#pending-token: {prefill_stats.num_pending_tokens}, " + ) + + if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL: + msg += f"#bootstrap-req: {len(self.scheduler.disagg_prefill_bootstrap_queue.queue)}, " + msg += ( + f"#inflight-req: {len(self.scheduler.disagg_prefill_inflight_queue)}, " + ) + + if ( + self.scheduler.server_args.language_only + and self.scheduler.server_args.encoder_transfer_backend + == "zmq_to_scheduler" + ): + msg += ( + f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, " + ) + + msg += f"{self._graph_backend_label}: {can_run_cuda_graph}, " + msg += f"input throughput (token/s): {self.last_input_throughput:.2f}" + + if self.enable_mfu_metrics and gap_latency > 0: + flops, _, _ = self._estimate_prefill_perf(prefill_stats.log_input_tokens) + tflops_per_s = flops / gap_latency / 1e12 + msg += f", est. prefill TFLOPS/s (per GPU): {tflops_per_s:.2f}" + + if ENABLE_METRICS_DEVICE_TIMER: + msg += f", fwd occupancy: {self.fwd_occupancy:.2f}%" + + if self.is_stats_logging_rank: + logger.info(msg) + if self.current_scheduler_metrics_enabled: + self.metrics_collector.increment_prefill_cuda_graph_pass( + value=can_run_cuda_graph + ) + self.metrics_collector.increment_realtime_tokens( + prefill_compute_tokens=prefill_stats.log_input_tokens, + prefill_cache_tokens=prefill_stats.log_hit_tokens, + dp_cooperation_info=dp_cooperation_info, + ) + if self.enable_mfu_metrics: + flops, read_bytes, write_bytes = self._estimate_prefill_perf( + prefill_stats.log_input_tokens + ) + self.metrics_collector.increment_estimated_perf( + num_flops_per_gpu=flops, + num_read_bytes_per_gpu=read_bytes, + num_write_bytes_per_gpu=write_bytes, + ) + + priority_enabled = self.scheduler.enable_priority_scheduling + total_tokens = prefill_stats.log_input_tokens + prefill_stats.log_hit_tokens + cache_hit_rate = ( + prefill_stats.log_hit_tokens / total_tokens if total_tokens > 0 else 0.0 + ) + + # Basics + self.stats.num_running_reqs = prefill_stats.num_running_reqs + self.stats.num_queue_reqs = QueueCount.from_reqs( + self.scheduler.waiting_queue, priority_enabled + ) + self.stats.num_grammar_queue_reqs = len(self.scheduler.grammar_manager) + self.stats.cache_hit_rate = cache_hit_rate + + # Memory pool usage ratios / Absolute token counts + pool_stats.update_scheduler_stats(self.stats) + + # Retract + self.stats.num_retracted_reqs = self.num_retracted_reqs + self.stats.num_paused_reqs = self.num_paused_reqs + self.num_retracted_reqs = self.num_paused_reqs = 0 + + # PD disaggregation + if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL: + self.stats.num_prefill_bootstrap_queue_reqs = QueueCount.from_reqs( + self.scheduler.disagg_prefill_bootstrap_queue.queue, + priority_enabled, + ) + self.stats.num_prefill_inflight_queue_reqs = QueueCount.from_reqs( + self.scheduler.disagg_prefill_inflight_queue, priority_enabled + ) + self.stats.kv_transfer_speed_gb_s = self.kv_transfer_speed_gb_s + self.stats.kv_transfer_latency_ms = self.kv_transfer_latency_ms + elif self.scheduler.disaggregation_mode == DisaggregationMode.DECODE: + self.stats.num_decode_prealloc_queue_reqs = QueueCount.from_reqs( + self.scheduler.disagg_decode_prealloc_queue.queue, priority_enabled + ) + self.stats.num_decode_transfer_queue_reqs = QueueCount.from_reqs( + self.scheduler.disagg_decode_transfer_queue.queue, priority_enabled + ) + + # Utilization / LoRA / HiCache + self._calculate_utilization() + self.stats.fwd_occupancy = self.fwd_occupancy + self._update_lora_metrics() + self._log_hicache_stats() + self.metrics_collector.log_stats(self.stats) + self.scheduler.kv_events_publisher.emit_kv_metrics() + self.scheduler.kv_events_publisher.publish_kv_events() + + def report_decode_stats( + self, + can_run_cuda_graph: bool, + running_batch: ScheduleBatch = None, + num_correct_drafts: int = 0, + ): + batch = running_batch or self.scheduler.running_batch + + # Every-iteration work: realtime token counting + status logger + if self.current_scheduler_metrics_enabled: + decode_tokens = batch.batch_size() + num_correct_drafts + self.metrics_collector.increment_realtime_tokens( + # TODO unify this w/ the bumping logic in `Scheduler.num_generated_tokens` accumulator + decode_tokens=decode_tokens, + dp_cooperation_info=batch.dp_cooperation_info, + ) + if self.enable_mfu_metrics: + flops, read_bytes, write_bytes = self._estimate_decode_perf( + batch, decode_tokens + ) + self.metrics_collector.increment_estimated_perf( + num_flops_per_gpu=flops, + num_read_bytes_per_gpu=read_bytes, + num_write_bytes_per_gpu=write_bytes, + ) + self._mfu_log_flops += flops + self._mfu_log_read_bytes += read_bytes + self._mfu_log_write_bytes += write_bytes + + if x := self.scheduler_status_logger: + x.maybe_dump(batch, self.scheduler.waiting_queue) + + # Periodic work: log + heavy metrics at decode_log_interval + if self.forward_ct_decode % self.scheduler.server_args.decode_log_interval != 0: + return + if ( + not self.is_stats_logging_rank + and not self.current_scheduler_metrics_enabled + ): + return + + gap_latency = time.perf_counter() - self.last_decode_stats_tic + self.last_decode_stats_tic = time.perf_counter() + self.last_gen_throughput = self.num_generated_tokens / gap_latency + + self.num_generated_tokens = 0 + num_running_reqs = len(batch.reqs) + + pool_stats = self.scheduler.pool_stats_observer.get_pool_stats() + token_usage_msg = ", ".join(pool_stats.get_decode_usage_msg_parts()) + ", " + + if RECORD_STEP_TIME: + self.step_time_dict[num_running_reqs].append( + gap_latency / self.scheduler.server_args.decode_log_interval + ) + + batch_iter = ( + batch.forward_iter + if batch is not None and batch.forward_iter is not None + else self.scheduler.forward_ct + ) + iter_msg = f" [{batch_iter}]" if LOG_FORWARD_ITERS else "" + msg = f"Decode batch{iter_msg}, #running-req: {num_running_reqs}, {token_usage_msg}" + + if self.scheduler.spec_algorithm.is_none(): + spec_accept_length = 0 + spec_accept_rate = 0 + else: + spec_accept_length = self.spec_num_accept_tokens / self.spec_num_forward_ct + num_correct_drafts = self.spec_num_accept_tokens - self.spec_num_forward_ct + if self.scheduler.server_args.speculative_num_draft_tokens: + draft_per_round = ( + self.scheduler.server_args.speculative_num_draft_tokens - 1 + ) + else: + draft_per_round = self.scheduler.server_args.speculative_num_steps or 0 + total_draft_tokens = self.spec_num_forward_ct * draft_per_round + spec_accept_rate = ( + num_correct_drafts / total_draft_tokens if total_draft_tokens > 0 else 0 + ) + self.spec_total_num_accept_tokens += self.spec_num_accept_tokens + self.spec_total_num_forward_ct += self.spec_num_forward_ct + self.spec_num_accept_tokens = self.spec_num_forward_ct = 0 + msg += f"accept len: {spec_accept_length:.2f}, accept rate: {spec_accept_rate:.2f}, " + cache_hit_rate = 0.0 + + if self.scheduler.disaggregation_mode == DisaggregationMode.DECODE: + msg += f"pre-allocated usage: {self.scheduler.disagg_decode_prealloc_queue.num_tokens_pre_allocated / self.scheduler.max_total_num_tokens:.2f}, " + msg += f"#prealloc-req: {len(self.scheduler.disagg_decode_prealloc_queue.queue)}, " + msg += f"#transfer-req: {len(self.scheduler.disagg_decode_transfer_queue.queue)}, " + msg += f"#retracted-req: {len(self.scheduler.disagg_decode_prealloc_queue.retracted_queue)}, " + + if ( + self.scheduler.server_args.language_only + and self.scheduler.server_args.encoder_transfer_backend + == "zmq_to_scheduler" + ): + msg += ( + f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, " + ) + + msg += ( + f"{self._graph_backend_label}: {can_run_cuda_graph}, " + f"gen throughput (token/s): {self.last_gen_throughput:.2f}, " + f"#queue-req: {len(self.scheduler.waiting_queue)}" + ) + + if self.enable_mfu_metrics and gap_latency > 0: + flops_per_s = self._mfu_log_flops / gap_latency + read_bytes_per_s = self._mfu_log_read_bytes / gap_latency + write_bytes_per_s = self._mfu_log_write_bytes / gap_latency + tflops_per_s = flops_per_s / 1e12 + read_gb_per_s = read_bytes_per_s / 1e9 + write_gb_per_s = write_bytes_per_s / 1e9 + msg += ( + f", est. decode TFLOPS/s (per GPU): {tflops_per_s:.2f}, " + f"est. read BW (GB/s per GPU): {read_gb_per_s:.2f}, " + f"est. write BW (GB/s per GPU): {write_gb_per_s:.2f}" + ) + self._mfu_log_flops = 0.0 + self._mfu_log_read_bytes = 0.0 + self._mfu_log_write_bytes = 0.0 + + if ENABLE_METRICS_DEVICE_TIMER: + msg += f", fwd occupancy: {self.fwd_occupancy:.2f}%" + + if self.is_stats_logging_rank: + logger.info(msg) + if self.current_scheduler_metrics_enabled: + priority_enabled = self.scheduler.enable_priority_scheduling + + # Basics + self.stats.num_running_reqs = QueueCount.from_reqs( + batch.reqs, priority_enabled + ) + self.stats.num_queue_reqs = QueueCount.from_reqs( + self.scheduler.waiting_queue, priority_enabled + ) + self.stats.num_grammar_queue_reqs = len(self.scheduler.grammar_manager) + self.stats.gen_throughput = self.last_gen_throughput + self.stats.cache_hit_rate = cache_hit_rate + self.stats.decode_sum_seq_lens = batch.seq_lens_cpu.sum().item() + + # Memory pool usage ratios / Absolute token counts + pool_stats.update_scheduler_stats(self.stats) + + # Speculative decoding + self.stats.spec_accept_length = spec_accept_length + self.stats.spec_accept_rate = spec_accept_rate + + # Retract + self.stats.num_retracted_reqs = self.num_retracted_reqs + self.stats.num_paused_reqs = self.num_paused_reqs + self.num_retracted_reqs = self.num_paused_reqs = 0 + + # PD disaggregation + if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL: + self.stats.num_prefill_bootstrap_queue_reqs = QueueCount.from_reqs( + self.scheduler.disagg_prefill_bootstrap_queue.queue, + priority_enabled, + ) + self.stats.num_prefill_inflight_queue_reqs = QueueCount.from_reqs( + self.scheduler.disagg_prefill_inflight_queue, priority_enabled + ) + elif self.scheduler.disaggregation_mode == DisaggregationMode.DECODE: + self.stats.num_decode_prealloc_queue_reqs = QueueCount.from_reqs( + self.scheduler.disagg_decode_prealloc_queue.queue, priority_enabled + ) + self.stats.num_decode_transfer_queue_reqs = QueueCount.from_reqs( + self.scheduler.disagg_decode_transfer_queue.queue, priority_enabled + ) + + # Streaming session metrics + self.stats.num_streaming_sessions = ( + self.scheduler.pool_stats_observer.streaming_session_count() + ) + self.stats.streaming_session_held_tokens = ( + self.scheduler.pool_stats_observer.session_held_tokens() + ) + + # Routing key metrics + # (to reduce the overhead, we only compute this when all requests have routing_key) + if all(r.routing_key is not None for r in batch.reqs): + running_routing_keys = [r.routing_key for r in batch.reqs] + waiting_routing_keys = [ + r.routing_key for r in self.scheduler.waiting_queue + ] + ( + self.stats.num_unique_running_routing_keys, + self.stats.routing_key_running_req_counts, + ) = compute_routing_key_stats(running_routing_keys) + _, self.stats.routing_key_all_req_counts = compute_routing_key_stats( + running_routing_keys + waiting_routing_keys + ) + + # Utilization / LoRA / HiCache + self._calculate_utilization() + self.stats.fwd_occupancy = self.fwd_occupancy + self._update_lora_metrics() + self._log_hicache_stats() + self.metrics_collector.log_stats(self.stats) + self.scheduler.kv_events_publisher.emit_kv_metrics() + self.scheduler.kv_events_publisher.publish_kv_events() + + def log_batch_result_stats( + self, + batch: ScheduleBatch, + result: Union[GenerationBatchResult, EmbeddingBatchResult], + ): + if not self.enable_metrics: + return + if not isinstance(result, GenerationBatchResult): + return + + if (m := result.expert_distribution_metrics) is not None: + self.metrics_collector.increment_eplb_balancedness( + forward_mode=batch.forward_mode.name.lower(), + balancedness=m.eplb_balancedness.item(), + ) + + def _emit_forward_pass_metrics( + self, + 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.scheduler.enable_fpm: + return + + from sglang.srt.observability.forward_pass_metrics import ( + ForwardPassMetrics, + ) + + if self.scheduler._fpm_uses_device_timer: + self.forward_pass_device_timer._report() + wall_time = self.scheduler._fpm_gpu_time_acc + self.scheduler._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.scheduler._fpm_worker_id, + dp_rank=self.scheduler._fpm_dp_rank, + wall_time=wall_time, + scheduled_requests=self._build_scheduled_request_metrics(batch), + queued_requests=self._build_queued_request_metrics(), + ) + self.scheduler._fpm_publisher.publish(fpm) + + def _shutdown_fpm(self): + """Shut down the FPM publisher thread.""" + if self.scheduler.enable_fpm: + self.scheduler._fpm_publisher.shutdown() + + def _log_hicache_stats(self): + """Populate HiCache host-tier stats on self.stats. + + These are pushed to Prometheus by SchedulerMetricsCollector.log_stats(). + """ + if not self.scheduler.enable_hierarchical_cache: + return + + host_pool = getattr( + self.scheduler.tree_cache, "token_to_kv_pool_host", None + ) or getattr(self.scheduler.tree_cache, "full_kv_pool_host", None) + assert host_pool is not None, "Host pool not found" + self.stats.hicache_host_used_tokens = ( + host_pool.size - host_pool.available_size() + ) + self.stats.hicache_host_total_tokens = host_pool.size + + def _update_lora_metrics(self): + """Update LoRA pool metrics for monitoring and autoscaling.""" + if not self.scheduler.enable_lora: + return + + try: + # Get LoRA memory pool stats + lora_manager = self.scheduler.tp_worker.model_runner.lora_manager + if lora_manager is None or lora_manager.memory_pool is None: + return + + mem_pool = lora_manager.memory_pool + slots_total = mem_pool.max_loras_per_batch + + # Calculate active adapters from running batch + # This gives a true measure of current load for autoscaling purposes + active_lora_ids = set() + + # For PP mode, check all running micro batches + if self.scheduler.server_args.pp_size > 1: + for batch in self.scheduler.running_mbs: + if batch and hasattr(batch, "reqs"): + for req in batch.reqs: + if hasattr(req, "lora_id") and req.lora_id is not None: + active_lora_ids.add(req.lora_id) + # For normal mode, check running_batch + elif self.scheduler.running_batch: + if hasattr(self.scheduler.running_batch, "reqs"): + for req in self.scheduler.running_batch.reqs: + if hasattr(req, "lora_id") and req.lora_id is not None: + active_lora_ids.add(req.lora_id) + + # Count active adapters (excluding None for base model) + slots_used = len(active_lora_ids) + utilization = slots_used / slots_total if slots_total > 0 else 0.0 + + # Update stats + self.stats.lora_pool_slots_used = slots_used + self.stats.lora_pool_slots_total = slots_total + self.stats.lora_pool_utilization = utilization + + except Exception as e: + logger.warning(f"Failed to update LoRA metrics: {e}") + + def _calculate_utilization(self): + if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL: + self.stats.utilization = -1 + else: + # TODO: max_running_requests_under_SLO has no setter — sglang:utilization stuck at 0 (regressed #22713). + max_under_slo = getattr( + self.scheduler, "max_running_requests_under_SLO", None + ) + if max_under_slo is not None and max_under_slo > 0: + self.stats.utilization = max( + self.stats.num_running_reqs.total / max_under_slo, + self.stats.token_usage / 0.9, + ) + + def update_device_timer(self): + if not ENABLE_METRICS_DEVICE_TIMER: + return + self.forward_pass_device_timer._report() + now = time.perf_counter() + if self._device_timer_window_batch_count == 0: + self._device_timer_window_start = now + self._device_timer_window_gpu_time = 0.0 + cpu_time = 0 + self.fwd_occupancy = float("nan") + else: + cpu_time = now - self._device_timer_window_start + self.fwd_occupancy = min( + self._device_timer_window_gpu_time / cpu_time * 100, 100 + ) + # ratio = self._device_timer_window_gpu_time / cpu_time if cpu_time > 0 else float("nan") + # print(f"{self._device_timer_window_batch_count=} {self.fwd_occupancy=}, {self._device_timer_window_gpu_time=}, {cpu_time=}, {ratio=}") + self._device_timer_window_batch_count += 1 + if ( + self._device_timer_window_batch_count + >= self.scheduler.server_args.decode_log_interval + ): + self._device_timer_window_batch_count = 0 + + def reset_device_timer_window(self): + if ENABLE_METRICS_DEVICE_TIMER: + self._device_timer_window_batch_count = 0 + self.fwd_occupancy = float("nan") diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py index b52f742ef..34551b1d2 100644 --- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py +++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py @@ -399,8 +399,7 @@ class SchedulerOutputProcessorMixin: self.stream_output(batch.reqs, batch.return_logprob, skip_stream_req) can_run_cuda_graph = getattr(result, "can_run_cuda_graph", False) - self.report_prefill_stats( - self.metrics_reporter, + self.metrics_reporter.report_prefill_stats( batch=batch, prefill_stats=batch.prefill_stats, can_run_cuda_graph=can_run_cuda_graph, @@ -514,8 +513,8 @@ class SchedulerOutputProcessorMixin: self.metrics_reporter.num_generated_tokens += len(batch.reqs) if not batch.spec_algorithm.is_none(): - self.update_spec_metrics( - self.metrics_reporter, batch.batch_size(), result.num_correct_drafts + self.metrics_reporter.update_spec_metrics( + batch.batch_size(), result.num_correct_drafts ) if self.metrics_reporter.enable_metrics: self.metrics_collector.increment_decode_cuda_graph_pass( @@ -630,8 +629,7 @@ class SchedulerOutputProcessorMixin: self.metrics_reporter.forward_ct_decode = ( self.metrics_reporter.forward_ct_decode + 1 ) % (1 << 30) - self.report_decode_stats( - self.metrics_reporter, + self.metrics_reporter.report_decode_stats( can_run_cuda_graph, running_batch=batch, num_correct_drafts=result.num_correct_drafts, diff --git a/python/sglang/srt/observability/scheduler_metrics_mixin.py b/python/sglang/srt/observability/scheduler_metrics_mixin.py deleted file mode 100644 index 81ab6514f..000000000 --- a/python/sglang/srt/observability/scheduler_metrics_mixin.py +++ /dev/null @@ -1,961 +0,0 @@ -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 - -from sglang.srt.disaggregation.utils import DisaggregationMode -from sglang.srt.environ import envs -from sglang.srt.managers.schedule_batch import ScheduleBatch -from sglang.srt.managers.utils import GenerationBatchResult -from sglang.srt.observability.metrics_collector import ( - DPCooperationInfo, - QueueCount, - SchedulerStats, - compute_routing_key_stats, -) -from sglang.srt.utils.device_timer import DeviceTimer -from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger - -if TYPE_CHECKING: - from sglang.srt.managers.schedule_batch import Req - from sglang.srt.managers.schedule_policy import PrefillAdder - from sglang.srt.managers.scheduler import EmbeddingBatchResult - -logger = logging.getLogger(__name__) - -RECORD_STEP_TIME = envs.SGLANG_RECORD_STEP_TIME.get() -LOG_FORWARD_ITERS = envs.SGLANG_LOG_FORWARD_ITERS.get() -ENABLE_METRICS_DEVICE_TIMER = envs.SGLANG_ENABLE_METRICS_DEVICE_TIMER.get() - - -@dataclasses.dataclass -class PrefillStats: - """Stats for logging prefill batch metrics.""" - - log_input_tokens: int - log_hit_tokens: int - new_token_ratio: float - num_running_reqs: QueueCount - num_new_seqs: int # len(can_run_list) - num_pending_tokens: int = 0 - - @classmethod - def from_adder( - cls, - adder: PrefillAdder, - running_reqs: List[Req], - enable_priority_scheduling: bool = False, - num_pending_tokens: int = 0, - ): - return cls( - log_input_tokens=adder.log_input_tokens, - log_hit_tokens=adder.log_hit_tokens, - new_token_ratio=adder.new_token_ratio, - num_running_reqs=QueueCount.from_reqs( - running_reqs, enable_priority_scheduling - ), - num_new_seqs=len(adder.can_run_list), - num_pending_tokens=num_pending_tokens, - ) - - -class SchedulerMetricsMixin: - enable_fpm: bool = False - - @staticmethod - def _init_metrics( - self: "SchedulerMetricsReporter", - tp_rank: int, - pp_rank: int, - dp_rank: Optional[int], - ): - # Basic stats - self.forward_ct_decode = 0 - self.num_generated_tokens = 0 - self.last_decode_stats_tic = time.perf_counter() - self.last_prefill_stats_tic = time.perf_counter() - self.last_gen_throughput: float = 0.0 - self.last_input_throughput: float = 0.0 - self.step_time_dict = defaultdict(list) # Dict[batch size -> step time] - self.stats = SchedulerStats() - self._graph_backend_label = { - "cpu": "cpu graph", - "npu": "npu graph", - "musa": "musa graph", - }.get(getattr(self.scheduler, "device", ""), "cuda graph") - - # Cumulative spec-decoding counters (reset every decode_log_interval). - # Each update adds (num_correct_drafts + bs, bs). - # `*_accept_tokens` = drafts + bonus; `*_correct_drafts` = drafts-only. - self.spec_num_accept_tokens = 0 # per-log-interval - self.spec_num_forward_ct = 0 - self.spec_total_num_accept_tokens = 0 # lifetime - self.spec_total_num_forward_ct = 0 - - # For PD disaggregation - self.kv_transfer_speed_gb_s: float = 0.0 - self.kv_transfer_latency_ms: float = 0.0 - - self.enable_mfu_metrics = False - - if self.enable_metrics: - self.enable_mfu_metrics = self.scheduler.server_args.enable_mfu_metrics - if self.enable_mfu_metrics: - SchedulerMetricsMixin._init_estimated_perf_constants(self) - self._mfu_log_flops = 0.0 - self._mfu_log_read_bytes = 0.0 - self._mfu_log_write_bytes = 0.0 - - self.fwd_occupancy = float("nan") - - self.forward_pass_device_timer: Optional[DeviceTimer] = None - - if ENABLE_METRICS_DEVICE_TIMER: - self._device_timer_window_batch_count = 0 - self._device_timer_window_gpu_time = 0.0 - self._device_timer_window_start = None - - def _wrap_execution_reporter(**kwargs): - self._device_timer_window_gpu_time += kwargs["t"] - if self.enable_metrics: - self.metrics_collector.increment_forward_execution_seconds(**kwargs) - - self.forward_pass_device_timer = DeviceTimer( - reporter=_wrap_execution_reporter, - ) - - SchedulerMetricsMixin._init_fpm(self) - - self.scheduler_status_logger = SchedulerStatusLogger.maybe_create( - enable_metrics=self.enable_metrics - ) - - @staticmethod - def _install_device_timer_on_runners(self: "SchedulerMetricsReporter"): - if self.forward_pass_device_timer is None: - return - timer = self.forward_pass_device_timer - self.scheduler.tp_worker.model_runner.device_timer = timer - if self.scheduler.draft_worker is not None: - dw = getattr(self.scheduler.draft_worker, "draft_worker", None) - if dw is not None: - if hasattr(dw, "draft_runner"): - dw.draft_runner.device_timer = timer - for r in getattr(dw, "draft_runner_list", []): - r.device_timer = timer - - @staticmethod - def _init_fpm(self: "SchedulerMetricsReporter"): - """Initialize Forward Pass Metrics (FPM) publisher if configured.""" - self.scheduler.enable_fpm = False - if ( - self.scheduler.server_args.enable_forward_pass_metrics - and self.scheduler.ps.attn_tp_rank == 0 - and self.scheduler.ps.pp_rank == self.scheduler.ps.pp_size - 1 - ): - from sglang.srt.observability.forward_pass_metrics import ( - _FpmPublisherThread, - ) - - self.scheduler._fpm_dp_rank = ( - self.scheduler.ps.dp_rank - if self.scheduler.ps.dp_rank is not None - else 0 - ) - self.scheduler._fpm_worker_id = ( - self.scheduler.server_args.forward_pass_metrics_worker_id - ) - base_endpoint = self.scheduler.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.scheduler.server_args.forward_pass_metrics_ipc_name = base_endpoint - endpoint = f"{base_endpoint}.{self.scheduler._fpm_dp_rank}" - self.scheduler._fpm_publisher = _FpmPublisherThread( - endpoint, - worker_id=self.scheduler._fpm_worker_id, - dp_rank=self.scheduler._fpm_dp_rank, - ) - self.scheduler._fpm_gpu_time_acc = 0.0 - - def _fpm_device_timer_reporter(t, **_kwargs): - self.scheduler._fpm_gpu_time_acc += t - - if self.forward_pass_device_timer is not None: - self.forward_pass_device_timer.add_reporter(_fpm_device_timer_reporter) - else: - self.forward_pass_device_timer = DeviceTimer( - reporter=_fpm_device_timer_reporter, - ) - self.scheduler._fpm_uses_device_timer = True - self.scheduler.enable_fpm = True - logger.info( - "FPM: ZMQ PUB bound on %s (dp_rank=%d, device_timer=%s)", - endpoint, - self.scheduler._fpm_dp_rank, - self.scheduler._fpm_uses_device_timer, - ) - - @staticmethod - def _build_scheduled_request_metrics( - self: "SchedulerMetricsReporter", 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(), - ) - - @staticmethod - def _build_queued_request_metrics(self: "SchedulerMetricsReporter"): - from sglang.srt.observability.forward_pass_metrics import ( - QueuedRequestMetrics, - WelfordAccumulator, - ) - - prefill_q = WelfordAccumulator() - decode_q = WelfordAccumulator() - if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL: - for req in self.scheduler.disagg_prefill_bootstrap_queue.queue: - prefill_q.add(len(req.origin_input_ids)) - elif self.scheduler.disaggregation_mode == DisaggregationMode.DECODE: - for req in self.scheduler.disagg_decode_prealloc_queue.queue: - decode_q.add(req.seqlen) - for req in self.scheduler.disagg_decode_transfer_queue.queue: - decode_q.add(req.seqlen) - else: - for req in self.scheduler.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(), - ) - - @staticmethod - def update_spec_metrics( - self: "SchedulerMetricsReporter", bs: int, num_correct_drafts: int - ): - self.spec_num_accept_tokens += num_correct_drafts + bs - self.spec_num_forward_ct += bs - - # Bonus tokens updated elsewhere - self.num_generated_tokens += num_correct_drafts - - @staticmethod - def _init_estimated_perf_constants(self: "SchedulerMetricsReporter") -> None: - model_config = self.scheduler.model_config - hf_text_config = model_config.hf_text_config - - hidden_size = float(model_config.hidden_size) - num_layers = float(getattr(model_config, "num_attention_layers", 0)) - head_dim = float(getattr(model_config, "head_dim", 0)) - num_attn_heads = float( - model_config.get_num_attention_heads(self.scheduler.ps.tp_size) - ) - num_kv_heads = float(model_config.get_num_kv_heads(self.scheduler.ps.tp_size)) - intermediate_size = getattr(hf_text_config, "intermediate_size", None) - if intermediate_size is None: - intermediate_size = getattr(hf_text_config, "ffn_hidden_size", 0) - intermediate_size = float(intermediate_size) - - dtype_num_bytes = getattr(model_config.dtype, "itemsize", None) - if dtype_num_bytes is None: - dtype_num_bytes = 2 - # Keep this estimator lightweight and consistent with current server dtype. - # KV cache quantization-aware bytes can be added in a follow-up. - act_bytes = float(dtype_num_bytes) - w_bytes = float(dtype_num_bytes) - cache_bytes = float(dtype_num_bytes) - - # Linear-layer FLOPs per token on one GPU. - attn_linear_flops = ( - 2.0 * hidden_size * head_dim * (num_attn_heads + 2.0 * num_kv_heads) - + 2.0 * hidden_size * head_dim * num_attn_heads - ) - mlp_flops = ( - 6.0 * hidden_size * intermediate_size if intermediate_size > 0 else 0.0 - ) - self._linear_flops_per_token = max( - 0.0, (attn_linear_flops + mlp_flops) * num_layers - ) - - # Attention dot-product FLOPs coefficient to multiply token-context product. - # attn_qk + attn_av = 4 * q * TC * d * L - self._attn_dot_flops_coeff = 4.0 * num_attn_heads * head_dim * num_layers - - # KV cache bytes (write one K and one V vector per generated token). - self._kv_cache_bytes_per_token = ( - 2.0 * num_layers * num_kv_heads * head_dim * cache_bytes - ) - - # Weight read bytes per token. - self._weight_read_bytes_per_token = ( - hidden_size - * head_dim - * (num_attn_heads + 2.0 * num_kv_heads) - * w_bytes - * num_layers - + hidden_size * head_dim * num_attn_heads * w_bytes * num_layers - + ( - 3.0 * hidden_size * intermediate_size * w_bytes * num_layers - if intermediate_size > 0 - else 0.0 - ) - ) - - # Activation movement bytes per token (coarse approximation). - self._qkv_act_bytes_per_token = ( - hidden_size * act_bytes * num_layers - + (num_attn_heads + 2.0 * num_kv_heads) * head_dim * act_bytes * num_layers - + head_dim * num_attn_heads * act_bytes * num_layers - + hidden_size * act_bytes * num_layers - ) - self._ffn_act_bytes_per_token = ( - 3.0 * intermediate_size * act_bytes * num_layers - if intermediate_size > 0 - else 0.0 - ) - - # Prefill reads Q/K/V activations from on-device memory. - self._prefill_attn_act_read_per_token = ( - (num_attn_heads + 2.0 * num_kv_heads) * head_dim * act_bytes * num_layers - ) - - # Decode reads Q from activation memory; K/V reads are from KV cache. - self._decode_q_read_bytes_per_token = ( - num_attn_heads * head_dim * act_bytes * num_layers - ) - - @staticmethod - def _estimate_prefill_perf( - self: "SchedulerMetricsReporter", num_tokens: int - ) -> Tuple[float, float, float]: - tokens = max(0, int(num_tokens)) - if tokens == 0: - return 0.0, 0.0, 0.0 - - # Causal prefill token-context product. - context_product = tokens * (tokens + 1) / 2.0 - flops = ( - tokens * self._linear_flops_per_token - + self._attn_dot_flops_coeff * context_product - ) - - read_bytes = ( - tokens * self._weight_read_bytes_per_token - + tokens * self._qkv_act_bytes_per_token - + tokens * self._prefill_attn_act_read_per_token - ) - write_bytes = ( - tokens * self._kv_cache_bytes_per_token - + tokens * self._qkv_act_bytes_per_token - + tokens * self._ffn_act_bytes_per_token - ) - return flops, read_bytes, write_bytes - - @staticmethod - def _estimate_decode_perf( - self: "SchedulerMetricsReporter", batch: ScheduleBatch, num_tokens: int - ) -> Tuple[float, float, float]: - tokens = max(0, int(num_tokens)) - if tokens == 0: - return 0.0, 0.0, 0.0 - - total_context = float(batch.seq_lens_cpu.sum().item()) - flops = ( - tokens * self._linear_flops_per_token - + self._attn_dot_flops_coeff * total_context - ) - read_bytes = ( - tokens * self._weight_read_bytes_per_token - + tokens * self._qkv_act_bytes_per_token - + tokens * self._decode_q_read_bytes_per_token - + total_context * self._kv_cache_bytes_per_token - ) - write_bytes = ( - tokens * self._kv_cache_bytes_per_token - + tokens * self._qkv_act_bytes_per_token - + tokens * self._ffn_act_bytes_per_token - ) - return flops, read_bytes, write_bytes - - @staticmethod - def reset_metrics(self: "SchedulerMetricsReporter"): - self.forward_ct_decode = 0 - self.num_generated_tokens = 0 - self.spec_num_accept_tokens = 0 - self.spec_num_forward_ct = 0 - self.spec_total_num_accept_tokens = 0 - self.spec_total_num_forward_ct = 0 - - @staticmethod - def report_prefill_stats( - self: "SchedulerMetricsReporter", - batch: Optional[ScheduleBatch], - prefill_stats: PrefillStats, - can_run_cuda_graph: bool, - dp_cooperation_info: Optional[DPCooperationInfo] = None, - ): - if ( - not self.is_stats_logging_rank - and not self.current_scheduler_metrics_enabled - ): - return - - now = time.perf_counter() - gap_latency = now - self.last_prefill_stats_tic - self.last_prefill_stats_tic = now - self.last_input_throughput = ( - prefill_stats.log_input_tokens / gap_latency if gap_latency > 0 else 0.0 - ) - - pool_stats = self.scheduler.pool_stats_observer.get_pool_stats() - token_usage_msg = ", ".join(pool_stats.get_prefill_usage_msg_parts()) + ", " - - self.stats.new_token_ratio = prefill_stats.new_token_ratio - batch_iter = ( - batch.forward_iter - if batch is not None and batch.forward_iter is not None - else self.scheduler.forward_ct - ) - iter_msg = f" [{batch_iter}]" if LOG_FORWARD_ITERS else "" - - msg = ( - f"Prefill batch{iter_msg}, " - f"#new-seq: {prefill_stats.num_new_seqs}, " - f"#new-token: {prefill_stats.log_input_tokens}, " - f"#cached-token: {prefill_stats.log_hit_tokens}, " - f"{token_usage_msg}" - f"#running-req: {prefill_stats.num_running_reqs.total}, " - f"#queue-req: {len(self.scheduler.waiting_queue)}, " - f"#pending-token: {prefill_stats.num_pending_tokens}, " - ) - - if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL: - msg += f"#bootstrap-req: {len(self.scheduler.disagg_prefill_bootstrap_queue.queue)}, " - msg += ( - f"#inflight-req: {len(self.scheduler.disagg_prefill_inflight_queue)}, " - ) - - if ( - self.scheduler.server_args.language_only - and self.scheduler.server_args.encoder_transfer_backend - == "zmq_to_scheduler" - ): - msg += ( - f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, " - ) - - msg += f"{self._graph_backend_label}: {can_run_cuda_graph}, " - msg += f"input throughput (token/s): {self.last_input_throughput:.2f}" - - if self.enable_mfu_metrics and gap_latency > 0: - flops, _, _ = SchedulerMetricsMixin._estimate_prefill_perf( - self, prefill_stats.log_input_tokens - ) - tflops_per_s = flops / gap_latency / 1e12 - msg += f", est. prefill TFLOPS/s (per GPU): {tflops_per_s:.2f}" - - if ENABLE_METRICS_DEVICE_TIMER: - msg += f", fwd occupancy: {self.fwd_occupancy:.2f}%" - - if self.is_stats_logging_rank: - logger.info(msg) - if self.current_scheduler_metrics_enabled: - self.metrics_collector.increment_prefill_cuda_graph_pass( - value=can_run_cuda_graph - ) - self.metrics_collector.increment_realtime_tokens( - prefill_compute_tokens=prefill_stats.log_input_tokens, - prefill_cache_tokens=prefill_stats.log_hit_tokens, - dp_cooperation_info=dp_cooperation_info, - ) - if self.enable_mfu_metrics: - flops, read_bytes, write_bytes = ( - SchedulerMetricsMixin._estimate_prefill_perf( - self, prefill_stats.log_input_tokens - ) - ) - self.metrics_collector.increment_estimated_perf( - num_flops_per_gpu=flops, - num_read_bytes_per_gpu=read_bytes, - num_write_bytes_per_gpu=write_bytes, - ) - - priority_enabled = self.scheduler.enable_priority_scheduling - total_tokens = prefill_stats.log_input_tokens + prefill_stats.log_hit_tokens - cache_hit_rate = ( - prefill_stats.log_hit_tokens / total_tokens if total_tokens > 0 else 0.0 - ) - - # Basics - self.stats.num_running_reqs = prefill_stats.num_running_reqs - self.stats.num_queue_reqs = QueueCount.from_reqs( - self.scheduler.waiting_queue, priority_enabled - ) - self.stats.num_grammar_queue_reqs = len(self.scheduler.grammar_manager) - self.stats.cache_hit_rate = cache_hit_rate - - # Memory pool usage ratios / Absolute token counts - pool_stats.update_scheduler_stats(self.stats) - - # Retract - self.stats.num_retracted_reqs = self.num_retracted_reqs - self.stats.num_paused_reqs = self.num_paused_reqs - self.num_retracted_reqs = self.num_paused_reqs = 0 - - # PD disaggregation - if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL: - self.stats.num_prefill_bootstrap_queue_reqs = QueueCount.from_reqs( - self.scheduler.disagg_prefill_bootstrap_queue.queue, - priority_enabled, - ) - self.stats.num_prefill_inflight_queue_reqs = QueueCount.from_reqs( - self.scheduler.disagg_prefill_inflight_queue, priority_enabled - ) - self.stats.kv_transfer_speed_gb_s = self.kv_transfer_speed_gb_s - self.stats.kv_transfer_latency_ms = self.kv_transfer_latency_ms - elif self.scheduler.disaggregation_mode == DisaggregationMode.DECODE: - self.stats.num_decode_prealloc_queue_reqs = QueueCount.from_reqs( - self.scheduler.disagg_decode_prealloc_queue.queue, priority_enabled - ) - self.stats.num_decode_transfer_queue_reqs = QueueCount.from_reqs( - self.scheduler.disagg_decode_transfer_queue.queue, priority_enabled - ) - - # Utilization / LoRA / HiCache - SchedulerMetricsMixin._calculate_utilization(self) - self.stats.fwd_occupancy = self.fwd_occupancy - SchedulerMetricsMixin._update_lora_metrics(self) - SchedulerMetricsMixin._log_hicache_stats(self) - self.metrics_collector.log_stats(self.stats) - self.scheduler.kv_events_publisher.emit_kv_metrics() - self.scheduler.kv_events_publisher.publish_kv_events() - - @staticmethod - def report_decode_stats( - self: "SchedulerMetricsReporter", - can_run_cuda_graph: bool, - running_batch: ScheduleBatch = None, - num_correct_drafts: int = 0, - ): - batch = running_batch or self.scheduler.running_batch - - # Every-iteration work: realtime token counting + status logger - if self.current_scheduler_metrics_enabled: - decode_tokens = batch.batch_size() + num_correct_drafts - self.metrics_collector.increment_realtime_tokens( - # TODO unify this w/ the bumping logic in `Scheduler.num_generated_tokens` accumulator - decode_tokens=decode_tokens, - dp_cooperation_info=batch.dp_cooperation_info, - ) - if self.enable_mfu_metrics: - flops, read_bytes, write_bytes = ( - SchedulerMetricsMixin._estimate_decode_perf( - self, batch, decode_tokens - ) - ) - self.metrics_collector.increment_estimated_perf( - num_flops_per_gpu=flops, - num_read_bytes_per_gpu=read_bytes, - num_write_bytes_per_gpu=write_bytes, - ) - self._mfu_log_flops += flops - self._mfu_log_read_bytes += read_bytes - self._mfu_log_write_bytes += write_bytes - - if x := self.scheduler_status_logger: - x.maybe_dump(batch, self.scheduler.waiting_queue) - - # Periodic work: log + heavy metrics at decode_log_interval - if self.forward_ct_decode % self.scheduler.server_args.decode_log_interval != 0: - return - if ( - not self.is_stats_logging_rank - and not self.current_scheduler_metrics_enabled - ): - return - - gap_latency = time.perf_counter() - self.last_decode_stats_tic - self.last_decode_stats_tic = time.perf_counter() - self.last_gen_throughput = self.num_generated_tokens / gap_latency - - self.num_generated_tokens = 0 - num_running_reqs = len(batch.reqs) - - pool_stats = self.scheduler.pool_stats_observer.get_pool_stats() - token_usage_msg = ", ".join(pool_stats.get_decode_usage_msg_parts()) + ", " - - if RECORD_STEP_TIME: - self.step_time_dict[num_running_reqs].append( - gap_latency / self.scheduler.server_args.decode_log_interval - ) - - batch_iter = ( - batch.forward_iter - if batch is not None and batch.forward_iter is not None - else self.scheduler.forward_ct - ) - iter_msg = f" [{batch_iter}]" if LOG_FORWARD_ITERS else "" - msg = f"Decode batch{iter_msg}, #running-req: {num_running_reqs}, {token_usage_msg}" - - if self.scheduler.spec_algorithm.is_none(): - spec_accept_length = 0 - spec_accept_rate = 0 - else: - spec_accept_length = self.spec_num_accept_tokens / self.spec_num_forward_ct - num_correct_drafts = self.spec_num_accept_tokens - self.spec_num_forward_ct - if self.scheduler.server_args.speculative_num_draft_tokens: - draft_per_round = ( - self.scheduler.server_args.speculative_num_draft_tokens - 1 - ) - else: - draft_per_round = self.scheduler.server_args.speculative_num_steps or 0 - total_draft_tokens = self.spec_num_forward_ct * draft_per_round - spec_accept_rate = ( - num_correct_drafts / total_draft_tokens if total_draft_tokens > 0 else 0 - ) - self.spec_total_num_accept_tokens += self.spec_num_accept_tokens - self.spec_total_num_forward_ct += self.spec_num_forward_ct - self.spec_num_accept_tokens = self.spec_num_forward_ct = 0 - msg += f"accept len: {spec_accept_length:.2f}, accept rate: {spec_accept_rate:.2f}, " - cache_hit_rate = 0.0 - - if self.scheduler.disaggregation_mode == DisaggregationMode.DECODE: - msg += f"pre-allocated usage: {self.scheduler.disagg_decode_prealloc_queue.num_tokens_pre_allocated / self.scheduler.max_total_num_tokens:.2f}, " - msg += f"#prealloc-req: {len(self.scheduler.disagg_decode_prealloc_queue.queue)}, " - msg += f"#transfer-req: {len(self.scheduler.disagg_decode_transfer_queue.queue)}, " - msg += f"#retracted-req: {len(self.scheduler.disagg_decode_prealloc_queue.retracted_queue)}, " - - if ( - self.scheduler.server_args.language_only - and self.scheduler.server_args.encoder_transfer_backend - == "zmq_to_scheduler" - ): - msg += ( - f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, " - ) - - msg += ( - f"{self._graph_backend_label}: {can_run_cuda_graph}, " - f"gen throughput (token/s): {self.last_gen_throughput:.2f}, " - f"#queue-req: {len(self.scheduler.waiting_queue)}" - ) - - if self.enable_mfu_metrics and gap_latency > 0: - flops_per_s = self._mfu_log_flops / gap_latency - read_bytes_per_s = self._mfu_log_read_bytes / gap_latency - write_bytes_per_s = self._mfu_log_write_bytes / gap_latency - tflops_per_s = flops_per_s / 1e12 - read_gb_per_s = read_bytes_per_s / 1e9 - write_gb_per_s = write_bytes_per_s / 1e9 - msg += ( - f", est. decode TFLOPS/s (per GPU): {tflops_per_s:.2f}, " - f"est. read BW (GB/s per GPU): {read_gb_per_s:.2f}, " - f"est. write BW (GB/s per GPU): {write_gb_per_s:.2f}" - ) - self._mfu_log_flops = 0.0 - self._mfu_log_read_bytes = 0.0 - self._mfu_log_write_bytes = 0.0 - - if ENABLE_METRICS_DEVICE_TIMER: - msg += f", fwd occupancy: {self.fwd_occupancy:.2f}%" - - if self.is_stats_logging_rank: - logger.info(msg) - if self.current_scheduler_metrics_enabled: - priority_enabled = self.scheduler.enable_priority_scheduling - - # Basics - self.stats.num_running_reqs = QueueCount.from_reqs( - batch.reqs, priority_enabled - ) - self.stats.num_queue_reqs = QueueCount.from_reqs( - self.scheduler.waiting_queue, priority_enabled - ) - self.stats.num_grammar_queue_reqs = len(self.scheduler.grammar_manager) - self.stats.gen_throughput = self.last_gen_throughput - self.stats.cache_hit_rate = cache_hit_rate - self.stats.decode_sum_seq_lens = batch.seq_lens_cpu.sum().item() - - # Memory pool usage ratios / Absolute token counts - pool_stats.update_scheduler_stats(self.stats) - - # Speculative decoding - self.stats.spec_accept_length = spec_accept_length - self.stats.spec_accept_rate = spec_accept_rate - - # Retract - self.stats.num_retracted_reqs = self.num_retracted_reqs - self.stats.num_paused_reqs = self.num_paused_reqs - self.num_retracted_reqs = self.num_paused_reqs = 0 - - # PD disaggregation - if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL: - self.stats.num_prefill_bootstrap_queue_reqs = QueueCount.from_reqs( - self.scheduler.disagg_prefill_bootstrap_queue.queue, - priority_enabled, - ) - self.stats.num_prefill_inflight_queue_reqs = QueueCount.from_reqs( - self.scheduler.disagg_prefill_inflight_queue, priority_enabled - ) - elif self.scheduler.disaggregation_mode == DisaggregationMode.DECODE: - self.stats.num_decode_prealloc_queue_reqs = QueueCount.from_reqs( - self.scheduler.disagg_decode_prealloc_queue.queue, priority_enabled - ) - self.stats.num_decode_transfer_queue_reqs = QueueCount.from_reqs( - self.scheduler.disagg_decode_transfer_queue.queue, priority_enabled - ) - - # Streaming session metrics - self.stats.num_streaming_sessions = ( - self.scheduler.pool_stats_observer.streaming_session_count() - ) - self.stats.streaming_session_held_tokens = ( - self.scheduler.pool_stats_observer.session_held_tokens() - ) - - # Routing key metrics - # (to reduce the overhead, we only compute this when all requests have routing_key) - if all(r.routing_key is not None for r in batch.reqs): - running_routing_keys = [r.routing_key for r in batch.reqs] - waiting_routing_keys = [ - r.routing_key for r in self.scheduler.waiting_queue - ] - ( - self.stats.num_unique_running_routing_keys, - self.stats.routing_key_running_req_counts, - ) = compute_routing_key_stats(running_routing_keys) - _, self.stats.routing_key_all_req_counts = compute_routing_key_stats( - running_routing_keys + waiting_routing_keys - ) - - # Utilization / LoRA / HiCache - SchedulerMetricsMixin._calculate_utilization(self) - self.stats.fwd_occupancy = self.fwd_occupancy - SchedulerMetricsMixin._update_lora_metrics(self) - SchedulerMetricsMixin._log_hicache_stats(self) - self.metrics_collector.log_stats(self.stats) - self.scheduler.kv_events_publisher.emit_kv_metrics() - self.scheduler.kv_events_publisher.publish_kv_events() - - @staticmethod - def log_batch_result_stats( - self: "SchedulerMetricsReporter", - batch: ScheduleBatch, - result: Union[GenerationBatchResult, EmbeddingBatchResult], - ): - if not self.enable_metrics: - return - if not isinstance(result, GenerationBatchResult): - return - - if (m := result.expert_distribution_metrics) is not None: - self.metrics_collector.increment_eplb_balancedness( - forward_mode=batch.forward_mode.name.lower(), - balancedness=m.eplb_balancedness.item(), - ) - - @staticmethod - def _emit_forward_pass_metrics( - self: "SchedulerMetricsReporter", - 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.scheduler.enable_fpm: - return - - from sglang.srt.observability.forward_pass_metrics import ( - ForwardPassMetrics, - ) - - if self.scheduler._fpm_uses_device_timer: - self.forward_pass_device_timer._report() - wall_time = self.scheduler._fpm_gpu_time_acc - self.scheduler._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.scheduler._fpm_worker_id, - dp_rank=self.scheduler._fpm_dp_rank, - wall_time=wall_time, - scheduled_requests=SchedulerMetricsMixin._build_scheduled_request_metrics( - self, batch - ), - queued_requests=SchedulerMetricsMixin._build_queued_request_metrics(self), - ) - self.scheduler._fpm_publisher.publish(fpm) - - @staticmethod - def _shutdown_fpm(self: "SchedulerMetricsReporter"): - """Shut down the FPM publisher thread.""" - if self.scheduler.enable_fpm: - self.scheduler._fpm_publisher.shutdown() - - @staticmethod - def _log_hicache_stats(self: "SchedulerMetricsReporter"): - """Populate HiCache host-tier stats on self.stats. - - These are pushed to Prometheus by SchedulerMetricsCollector.log_stats(). - """ - if not self.scheduler.enable_hierarchical_cache: - return - - host_pool = getattr( - self.scheduler.tree_cache, "token_to_kv_pool_host", None - ) or getattr(self.scheduler.tree_cache, "full_kv_pool_host", None) - assert host_pool is not None, "Host pool not found" - self.stats.hicache_host_used_tokens = ( - host_pool.size - host_pool.available_size() - ) - self.stats.hicache_host_total_tokens = host_pool.size - - @staticmethod - def _update_lora_metrics(self: "SchedulerMetricsReporter"): - """Update LoRA pool metrics for monitoring and autoscaling.""" - if not self.scheduler.enable_lora: - return - - try: - # Get LoRA memory pool stats - lora_manager = self.scheduler.tp_worker.model_runner.lora_manager - if lora_manager is None or lora_manager.memory_pool is None: - return - - mem_pool = lora_manager.memory_pool - slots_total = mem_pool.max_loras_per_batch - - # Calculate active adapters from running batch - # This gives a true measure of current load for autoscaling purposes - active_lora_ids = set() - - # For PP mode, check all running micro batches - if self.scheduler.server_args.pp_size > 1: - for batch in self.scheduler.running_mbs: - if batch and hasattr(batch, "reqs"): - for req in batch.reqs: - if hasattr(req, "lora_id") and req.lora_id is not None: - active_lora_ids.add(req.lora_id) - # For normal mode, check running_batch - elif self.scheduler.running_batch: - if hasattr(self.scheduler.running_batch, "reqs"): - for req in self.scheduler.running_batch.reqs: - if hasattr(req, "lora_id") and req.lora_id is not None: - active_lora_ids.add(req.lora_id) - - # Count active adapters (excluding None for base model) - slots_used = len(active_lora_ids) - utilization = slots_used / slots_total if slots_total > 0 else 0.0 - - # Update stats - self.stats.lora_pool_slots_used = slots_used - self.stats.lora_pool_slots_total = slots_total - self.stats.lora_pool_utilization = utilization - - except Exception as e: - logger.warning(f"Failed to update LoRA metrics: {e}") - - @staticmethod - def _calculate_utilization(self: "SchedulerMetricsReporter"): - if self.scheduler.disaggregation_mode == DisaggregationMode.PREFILL: - self.stats.utilization = -1 - else: - # TODO: max_running_requests_under_SLO has no setter — sglang:utilization stuck at 0 (regressed #22713). - max_under_slo = getattr( - self.scheduler, "max_running_requests_under_SLO", None - ) - if max_under_slo is not None and max_under_slo > 0: - self.stats.utilization = max( - self.stats.num_running_reqs.total / max_under_slo, - self.stats.token_usage / 0.9, - ) - - @staticmethod - def update_device_timer(self: "SchedulerMetricsReporter"): - if not ENABLE_METRICS_DEVICE_TIMER: - return - self.forward_pass_device_timer._report() - now = time.perf_counter() - if self._device_timer_window_batch_count == 0: - self._device_timer_window_start = now - self._device_timer_window_gpu_time = 0.0 - cpu_time = 0 - self.fwd_occupancy = float("nan") - else: - cpu_time = now - self._device_timer_window_start - self.fwd_occupancy = min( - self._device_timer_window_gpu_time / cpu_time * 100, 100 - ) - # ratio = self._device_timer_window_gpu_time / cpu_time if cpu_time > 0 else float("nan") - # print(f"{self._device_timer_window_batch_count=} {self.fwd_occupancy=}, {self._device_timer_window_gpu_time=}, {cpu_time=}, {ratio=}") - self._device_timer_window_batch_count += 1 - if ( - self._device_timer_window_batch_count - >= self.scheduler.server_args.decode_log_interval - ): - self._device_timer_window_batch_count = 0 - - @staticmethod - def reset_device_timer_window(self: "SchedulerMetricsReporter"): - if ENABLE_METRICS_DEVICE_TIMER: - self._device_timer_window_batch_count = 0 - self.fwd_occupancy = float("nan") diff --git a/test/registered/unit/observability/test_forward_pass_metrics.py b/test/registered/unit/observability/test_forward_pass_metrics.py index a34470951..6ed3f992a 100644 --- a/test/registered/unit/observability/test_forward_pass_metrics.py +++ b/test/registered/unit/observability/test_forward_pass_metrics.py @@ -8,9 +8,9 @@ from unittest.mock import patch from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.distributed.parallel_state_wrapper import ParallelState -from sglang.srt.observability.scheduler_metrics_mixin import ( +from sglang.srt.managers.scheduler_components.metrics_reporter import ( PrefillStats, - SchedulerMetricsMixin, + SchedulerMetricsReporter, ) @@ -85,14 +85,49 @@ class _DummyPublisherThread: pass -class _DummyScheduler(SchedulerMetricsMixin): - pass +def _make_reporter(scheduler) -> SchedulerMetricsReporter: + if not hasattr(scheduler, "server_args"): + scheduler.server_args = types.SimpleNamespace( + enable_metrics=False, + enable_metrics_for_all_schedulers=False, + kv_events_config=None, + enable_mfu_metrics=False, + enable_forward_pass_metrics=False, + ) + if not hasattr(scheduler, "ps"): + scheduler.ps = types.SimpleNamespace(attn_tp_rank=0, attn_cp_rank=0) + if not hasattr(scheduler, "kv_events_publisher"): + scheduler.kv_events_publisher = types.SimpleNamespace( + init_kv_events=lambda *a, **kw: None, + ) + if not hasattr(scheduler, "tp_workers"): + scheduler.tp_workers = [] + if not hasattr(scheduler, "tp_worker"): + scheduler.tp_worker = types.SimpleNamespace( + model_runner=types.SimpleNamespace(), + ) + if not hasattr(scheduler, "draft_worker"): + scheduler.draft_worker = None + context = types.SimpleNamespace( + enable_metrics=False, + is_stats_logging_rank=True, + current_scheduler_metrics_enabled=False, + enable_kv_cache_events=False, + collector=None, + ) + return SchedulerMetricsReporter( + scheduler=scheduler, + tp_rank=0, + pp_rank=0, + dp_rank=0, + metrics_collector_context=context, + metrics_collector=None, + ) class TestForwardPassMetrics(unittest.TestCase): def setUp(self): - self.scheduler = _DummyScheduler() - self.scheduler.enable_fpm = True + self.scheduler = types.SimpleNamespace() self.scheduler._fpm_worker_id = "worker-7" self.scheduler._fpm_dp_rank = 0 self.scheduler._fpm_publisher = _CollectingPublisher() @@ -100,6 +135,8 @@ class TestForwardPassMetrics(unittest.TestCase): self.scheduler._fpm_gpu_time_acc = 0.0 self.scheduler.waiting_queue = [] self.scheduler.disaggregation_mode = DisaggregationMode.NULL + self.reporter = _make_reporter(self.scheduler) + self.scheduler.enable_fpm = True def _make_batch(self, **overrides): defaults = dict( @@ -135,10 +172,10 @@ class TestForwardPassMetrics(unittest.TestCase): ) with patch( - "sglang.srt.observability.scheduler_metrics_mixin.time.monotonic", + "sglang.srt.managers.scheduler_components.metrics_reporter.time.monotonic", return_value=104.5, ): - self.scheduler._emit_forward_pass_metrics(batch) + self.reporter._emit_forward_pass_metrics(batch) self.assertEqual(len(self.scheduler._fpm_publisher.metrics), 1) metrics = self.scheduler._fpm_publisher.metrics[0] @@ -158,12 +195,12 @@ class TestForwardPassMetrics(unittest.TestCase): 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( + self.reporter.forward_pass_device_timer = types.SimpleNamespace( _report=lambda: None, ) batch = self._make_batch() - self.scheduler._emit_forward_pass_metrics(batch) + self.reporter._emit_forward_pass_metrics(batch) self.assertEqual(len(self.scheduler._fpm_publisher.metrics), 1) self.assertAlmostEqual( @@ -174,12 +211,12 @@ class TestForwardPassMetrics(unittest.TestCase): 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( + self.reporter.forward_pass_device_timer = types.SimpleNamespace( _report=lambda: None, ) batch = self._make_batch() - self.scheduler._emit_forward_pass_metrics(batch) + self.reporter._emit_forward_pass_metrics(batch) self.assertEqual(len(self.scheduler._fpm_publisher.metrics), 0) @@ -187,10 +224,10 @@ class TestForwardPassMetrics(unittest.TestCase): batch = self._make_batch() with patch( - "sglang.srt.observability.scheduler_metrics_mixin.time.monotonic", + "sglang.srt.managers.scheduler_components.metrics_reporter.time.monotonic", return_value=100.035, ): - self.scheduler._emit_forward_pass_metrics(batch, result=None) + self.reporter._emit_forward_pass_metrics(batch, result=None) self.assertEqual(len(self.scheduler._fpm_publisher.metrics), 1) self.assertAlmostEqual( @@ -205,10 +242,10 @@ class TestForwardPassMetrics(unittest.TestCase): batch = self._make_batch() with patch( - "sglang.srt.observability.scheduler_metrics_mixin.time.monotonic", + "sglang.srt.managers.scheduler_components.metrics_reporter.time.monotonic", return_value=101.0, ): - self.scheduler._emit_forward_pass_metrics(batch) + self.reporter._emit_forward_pass_metrics(batch) metrics = self.scheduler._fpm_publisher.metrics[0] self.assertEqual(metrics.queued_requests.num_prefill_requests, 3) @@ -226,10 +263,10 @@ class TestForwardPassMetrics(unittest.TestCase): batch = self._make_batch() with patch( - "sglang.srt.observability.scheduler_metrics_mixin.time.monotonic", + "sglang.srt.managers.scheduler_components.metrics_reporter.time.monotonic", return_value=101.0, ): - self.scheduler._emit_forward_pass_metrics(batch) + self.reporter._emit_forward_pass_metrics(batch) metrics = self.scheduler._fpm_publisher.metrics[0] self.assertEqual(metrics.queued_requests.num_prefill_requests, 0) @@ -237,7 +274,7 @@ class TestForwardPassMetrics(unittest.TestCase): self.assertEqual(metrics.queued_requests.sum_decode_kv_tokens, 15 + 30 + 45) def test_init_metrics_uses_server_worker_id(self): - scheduler = _DummyScheduler() + scheduler = types.SimpleNamespace() scheduler.server_args = types.SimpleNamespace( enable_metrics=False, enable_metrics_for_all_schedulers=False, @@ -254,7 +291,7 @@ class TestForwardPassMetrics(unittest.TestCase): "sglang.srt.observability.forward_pass_metrics._FpmPublisherThread", _DummyPublisherThread, ): - scheduler.init_metrics(tp_rank=0, pp_rank=0, dp_rank=2) + reporter = _make_reporter(scheduler) self.assertTrue(scheduler.enable_fpm) self.assertEqual(scheduler._fpm_worker_id, "endpoint-42") @@ -265,7 +302,7 @@ class TestForwardPassMetrics(unittest.TestCase): self.assertIsNotNone(scheduler.server_args.forward_pass_metrics_ipc_name) def test_init_fpm_disabled_on_non_last_pp_rank(self): - scheduler = _DummyScheduler() + scheduler = types.SimpleNamespace() scheduler.server_args = types.SimpleNamespace( enable_metrics=False, enable_metrics_for_all_schedulers=False, @@ -282,7 +319,7 @@ class TestForwardPassMetrics(unittest.TestCase): "sglang.srt.observability.forward_pass_metrics._FpmPublisherThread", _DummyPublisherThread, ): - scheduler.init_metrics(tp_rank=0, pp_rank=0, dp_rank=0) + reporter = _make_reporter(scheduler) self.assertFalse(scheduler.enable_fpm)