diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index c11394497..44f205c8f 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -3140,7 +3140,7 @@ class Scheduler( self._maybe_log_idle_metrics() # kv event publishing - self.publish_kv_events(self.kv_events_publisher) + self.kv_events_publisher.publish_kv_events() # reset token ratio self.new_token_ratio = self.init_new_token_ratio diff --git a/python/sglang/srt/managers/scheduler_components/kv_events_publisher.py b/python/sglang/srt/managers/scheduler_components/kv_events_publisher.py index 1f8dda3f9..85531e49e 100644 --- a/python/sglang/srt/managers/scheduler_components/kv_events_publisher.py +++ b/python/sglang/srt/managers/scheduler_components/kv_events_publisher.py @@ -1,6 +1,7 @@ from __future__ import annotations import dataclasses +import time from dataclasses import dataclass from typing import ( TYPE_CHECKING, @@ -11,6 +12,11 @@ from typing import ( import zmq +from sglang.srt.disaggregation.kv_events import ( + EventPublisherFactory, + KVEventBatch, +) + if TYPE_CHECKING: from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache @@ -48,8 +54,44 @@ class SchedulerKvEventsPublisher: kv_event_publisher: Any = None def __post_init__(self) -> None: - from sglang.srt.observability.scheduler_metrics_mixin import ( - SchedulerMetricsMixin, + self.init_kv_events(self.kv_events_config) + + def init_kv_events(self, kv_events_config: Optional[str]): + self.enable_kv_cache_events = bool( + kv_events_config and self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0 ) - SchedulerMetricsMixin.init_kv_events(self, self.kv_events_config) + if self.enable_kv_cache_events: + self.kv_event_publisher = EventPublisherFactory.create( + kv_events_config, self.ps.attn_dp_rank + ) + + def emit_kv_metrics(self): + if not self.enable_kv_cache_events: + return + + kv_metrics = KvMetrics() + kv_metrics.request_active_slots = self.get_stats().num_running_reqs.total + kv_metrics.request_total_slots = self.max_running_requests + kv_metrics.kv_active_blocks = int( + self.get_stats().token_usage * self.max_total_num_tokens + ) + kv_metrics.kv_total_blocks = self.max_total_num_tokens + kv_metrics.num_requests_waiting = self.get_stats().num_queue_reqs.total + kv_metrics.gpu_cache_usage_perc = self.get_stats().token_usage + kv_metrics.gpu_prefix_cache_hit_rate = self.get_stats().cache_hit_rate + kv_metrics.data_parallel_rank = ( + self.ps.dp_rank if self.ps.dp_rank is not None else 0 + ) + + if not self.send_metrics_from_scheduler.closed: + self.send_metrics_from_scheduler.send_pyobj(kv_metrics) + + def publish_kv_events(self): + if not self.enable_kv_cache_events: + return + + events = self.tree_cache.take_events() + if events: + batch = KVEventBatch(ts=time.time(), events=events) + self.kv_event_publisher.publish(batch) diff --git a/python/sglang/srt/observability/scheduler_metrics_mixin.py b/python/sglang/srt/observability/scheduler_metrics_mixin.py index 44aefddb0..263556d5c 100644 --- a/python/sglang/srt/observability/scheduler_metrics_mixin.py +++ b/python/sglang/srt/observability/scheduler_metrics_mixin.py @@ -7,7 +7,6 @@ import time from collections import defaultdict from typing import TYPE_CHECKING, List, Optional, Tuple, Union -from sglang.srt.disaggregation.kv_events import EventPublisherFactory, KVEventBatch from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.environ import envs from sglang.srt.managers.io_struct import ( @@ -20,7 +19,6 @@ from sglang.srt.managers.io_struct import ( SpeculativeMetrics, ) from sglang.srt.managers.schedule_batch import ScheduleBatch -from sglang.srt.managers.scheduler_components.kv_events_publisher import KvMetrics from sglang.srt.managers.utils import GenerationBatchResult from sglang.srt.observability.metrics_collector import ( DPCooperationInfo, @@ -36,9 +34,6 @@ 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, Scheduler - from sglang.srt.managers.scheduler_components.kv_events_publisher import ( - SchedulerKvEventsPublisher, - ) logger = logging.getLogger(__name__) @@ -189,19 +184,6 @@ class SchedulerMetricsMixin: for r in getattr(dw, "draft_runner_list", []): r.device_timer = timer - @staticmethod - def init_kv_events( - self: "SchedulerKvEventsPublisher", kv_events_config: Optional[str] - ): - self.enable_kv_cache_events = bool( - kv_events_config and self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0 - ) - - if self.enable_kv_cache_events: - self.kv_event_publisher = EventPublisherFactory.create( - kv_events_config, self.ps.attn_dp_rank - ) - def _init_fpm(self: Scheduler): """Initialize Forward Pass Metrics (FPM) publisher if configured.""" self.enable_fpm = False @@ -602,8 +584,8 @@ class SchedulerMetricsMixin: self.update_lora_metrics() self._log_hicache_stats() self.metrics_collector.log_stats(self.stats) - self.emit_kv_metrics(self.kv_events_publisher) - self.publish_kv_events(self.kv_events_publisher) + self.kv_events_publisher.emit_kv_metrics() + self.kv_events_publisher.publish_kv_events() def report_decode_stats( self: Scheduler, @@ -798,8 +780,8 @@ class SchedulerMetricsMixin: self.update_lora_metrics() self._log_hicache_stats() self.metrics_collector.log_stats(self.stats) - self.emit_kv_metrics(self.kv_events_publisher) - self.publish_kv_events(self.kv_events_publisher) + self.kv_events_publisher.emit_kv_metrics() + self.kv_events_publisher.publish_kv_events() def log_batch_result_stats( self: Scheduler, @@ -817,38 +799,6 @@ class SchedulerMetricsMixin: balancedness=m.eplb_balancedness.item(), ) - @staticmethod - def emit_kv_metrics(self: "SchedulerKvEventsPublisher"): - if not self.enable_kv_cache_events: - return - - kv_metrics = KvMetrics() - kv_metrics.request_active_slots = self.get_stats().num_running_reqs.total - kv_metrics.request_total_slots = self.max_running_requests - kv_metrics.kv_active_blocks = int( - self.get_stats().token_usage * self.max_total_num_tokens - ) - kv_metrics.kv_total_blocks = self.max_total_num_tokens - kv_metrics.num_requests_waiting = self.get_stats().num_queue_reqs.total - kv_metrics.gpu_cache_usage_perc = self.get_stats().token_usage - kv_metrics.gpu_prefix_cache_hit_rate = self.get_stats().cache_hit_rate - kv_metrics.data_parallel_rank = ( - self.ps.dp_rank if self.ps.dp_rank is not None else 0 - ) - - if not self.send_metrics_from_scheduler.closed: - self.send_metrics_from_scheduler.send_pyobj(kv_metrics) - - @staticmethod - def publish_kv_events(self: "SchedulerKvEventsPublisher"): - if not self.enable_kv_cache_events: - return - - events = self.tree_cache.take_events() - if events: - batch = KVEventBatch(ts=time.time(), events=events) - self.kv_event_publisher.publish(batch) - def _emit_forward_pass_metrics( self: Scheduler, batch: ScheduleBatch,