Move KV-cache event emission to SchedulerKvEventsPublisher (#25626)
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user