Move KV-cache event emission to SchedulerKvEventsPublisher (#25626)

This commit is contained in:
fzyzcjy
2026-05-18 18:40:20 +08:00
committed by GitHub
parent 0f888442c2
commit 1213277879
3 changed files with 50 additions and 58 deletions
+1 -1
View File
@@ -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
@@ -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)
@@ -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,