diff --git a/python/sglang/srt/managers/scheduler_metrics_mixin.py b/python/sglang/srt/managers/scheduler_metrics_mixin.py index 4735e8f5d..87a5f92e8 100644 --- a/python/sglang/srt/managers/scheduler_metrics_mixin.py +++ b/python/sglang/srt/managers/scheduler_metrics_mixin.py @@ -4,7 +4,7 @@ import logging import time from collections import defaultdict from contextlib import contextmanager -from typing import TYPE_CHECKING, List, Optional, Union +from typing import TYPE_CHECKING, Dict, List, Optional, Union from sglang.srt.disaggregation.kv_events import EventPublisherFactory, KVEventBatch from sglang.srt.disaggregation.utils import DisaggregationMode @@ -49,7 +49,6 @@ class SchedulerMetricsMixin: self.last_decode_stats_tic = time.perf_counter() self.last_prefill_stats_tic = time.perf_counter() self.last_prefill_tokens = 0 - self.last_prefill_cache_tokens = 0 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] @@ -68,6 +67,9 @@ class SchedulerMetricsMixin: self.kv_transfer_alloc_ms: float = 0.0 self.kv_transfer_total_mb: float = 0.0 + # Only for `log_prefill_stats` to pass information to `log_prefill_stats_late` + self.temp_prefill_info: Optional[Dict] = None + self.stats = SchedulerStats() # Metrics @@ -128,7 +130,12 @@ class SchedulerMetricsMixin: self.last_prefill_stats_tic = time.perf_counter() self.last_input_throughput = self.last_prefill_tokens / gap_latency self.last_prefill_tokens = adder.log_input_tokens - self.last_prefill_cache_tokens = adder.log_hit_tokens + + assert self.temp_prefill_info is None + self.temp_prefill_info = dict( + adder_log_input_tokens=adder.log_input_tokens, + adder_log_hit_tokens=adder.log_hit_tokens, + ) # TODO: generalize this for various memory pools if self.is_hybrid_swa: @@ -245,10 +252,14 @@ class SchedulerMetricsMixin: def log_prefill_stats_late(self: Scheduler, batch: Optional[ScheduleBatch]): """This should be called after `batch` has gathered enough metadata.""" - if self.enable_metrics and batch is not None: + + info = self.temp_prefill_info + self.temp_prefill_info = None + + if self.enable_metrics and batch is not None and info is not None: self.metrics_collector.increment_realtime_tokens( - prefill_compute_tokens=self.last_prefill_tokens, - prefill_cache_tokens=self.last_prefill_cache_tokens, + prefill_compute_tokens=info["adder_log_input_tokens"], + prefill_cache_tokens=info["adder_log_hit_tokens"], dp_cooperation_info=batch.dp_cooperation_info, )