Stand up SchedulerKvEventsPublisher; migrate KV-event state to it (#25625)

This commit is contained in:
fzyzcjy
2026-05-18 18:39:50 +08:00
committed by GitHub
parent f3dce08283
commit 0f888442c2
3 changed files with 100 additions and 28 deletions
+24 -2
View File
@@ -170,6 +170,9 @@ from sglang.srt.managers.scheduler_components.dp_attn import (
from sglang.srt.managers.scheduler_components.invariant_checker import (
SchedulerInvariantChecker,
)
from sglang.srt.managers.scheduler_components.kv_events_publisher import (
SchedulerKvEventsPublisher,
)
from sglang.srt.managers.scheduler_components.pool_stats_observer import (
SchedulerPoolStatsObserver,
)
@@ -491,7 +494,11 @@ class Scheduler(
tp_cpu_group=self.tp_cpu_group,
attn_cp_cpu_group=self.attn_cp_cpu_group,
enable_metrics=self.enable_metrics,
enable_kv_cache_events=self.enable_kv_cache_events,
enable_kv_cache_events=bool(
self.server_args.kv_events_config
and self.ps.attn_tp_rank == 0
and self.ps.attn_cp_rank == 0
),
ps=self.ps,
tp_group=self.tp_group,
enable_hierarchical_cache=self.enable_hierarchical_cache,
@@ -674,6 +681,20 @@ class Scheduler(
get_running_batch=lambda: self.running_batch,
)
self.kv_events_publisher = SchedulerKvEventsPublisher(
kv_events_config=self.server_args.kv_events_config,
ps=self.ps,
attn_tp_rank=self.ps.attn_tp_rank,
attn_cp_rank=self.ps.attn_cp_rank,
attn_dp_rank=self.ps.attn_dp_rank,
dp_rank=self.ps.dp_rank,
tree_cache=self.tree_cache,
send_metrics_from_scheduler=self.send_metrics_from_scheduler,
max_running_requests=self.max_running_requests,
max_total_num_tokens=self.max_total_num_tokens,
get_stats=lambda: self.stats,
)
self.is_initializing = False
def init_zbal_on_npu(self):
@@ -710,6 +731,7 @@ class Scheduler(
def init_ipc_channels(self, port_args: PortArgs):
context = zmq.Context(2)
self.idle_sleeper = None
self.send_metrics_from_scheduler = None
if (
self.ps.pp_rank == 0
@@ -3118,7 +3140,7 @@ class Scheduler(
self._maybe_log_idle_metrics()
# kv event publishing
self._publish_kv_events()
self.publish_kv_events(self.kv_events_publisher)
# reset token ratio
self.new_token_ratio = self.init_new_token_ratio
@@ -0,0 +1,55 @@
from __future__ import annotations
import dataclasses
from dataclasses import dataclass
from typing import (
TYPE_CHECKING,
Any,
Callable,
Optional,
)
import zmq
if TYPE_CHECKING:
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
class SchedulerStats: ... # type: ignore[no-redef]
@dataclasses.dataclass
class KvMetrics:
request_active_slots: int = 0
request_total_slots: int = 0
kv_active_blocks: int = 0
kv_total_blocks: int = 0
num_requests_waiting: int = 0
gpu_cache_usage_perc: float = 0.0
gpu_prefix_cache_hit_rate: float = 0.0
data_parallel_rank: int = 0
@dataclass(kw_only=True, slots=True)
class SchedulerKvEventsPublisher:
kv_events_config: Optional[str]
ps: "ParallelState"
attn_tp_rank: int
attn_cp_rank: int
attn_dp_rank: int
dp_rank: Optional[int]
tree_cache: "BasePrefixCache"
send_metrics_from_scheduler: Optional["zmq.Socket"]
max_running_requests: int
max_total_num_tokens: int
get_stats: Callable
enable_kv_cache_events: bool = False
kv_event_publisher: Any = None
def __post_init__(self) -> None:
from sglang.srt.observability.scheduler_metrics_mixin import (
SchedulerMetricsMixin,
)
SchedulerMetricsMixin.init_kv_events(self, self.kv_events_config)
@@ -20,6 +20,7 @@ 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,
@@ -35,6 +36,9 @@ 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__)
@@ -74,18 +78,6 @@ class PrefillStats:
)
@dataclasses.dataclass
class KvMetrics:
request_active_slots: int = 0
request_total_slots: int = 0
kv_active_blocks: int = 0
kv_total_blocks: int = 0
num_requests_waiting: int = 0
gpu_cache_usage_perc: float = 0.0
gpu_prefix_cache_hit_rate: float = 0.0
data_parallel_rank: int = 0
class SchedulerMetricsMixin:
enable_fpm: bool = False
@@ -178,8 +170,6 @@ class SchedulerMetricsMixin:
reporter=_wrap_execution_reporter,
)
self.init_kv_events(self.server_args.kv_events_config)
self._init_fpm()
self.scheduler_status_logger = SchedulerStatusLogger.maybe_create(
@@ -199,7 +189,10 @@ class SchedulerMetricsMixin:
for r in getattr(dw, "draft_runner_list", []):
r.device_timer = timer
def init_kv_events(self: Scheduler, kv_events_config: Optional[str]):
@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
)
@@ -609,8 +602,8 @@ class SchedulerMetricsMixin:
self.update_lora_metrics()
self._log_hicache_stats()
self.metrics_collector.log_stats(self.stats)
self._emit_kv_metrics()
self._publish_kv_events()
self.emit_kv_metrics(self.kv_events_publisher)
self.publish_kv_events(self.kv_events_publisher)
def report_decode_stats(
self: Scheduler,
@@ -805,8 +798,8 @@ class SchedulerMetricsMixin:
self.update_lora_metrics()
self._log_hicache_stats()
self.metrics_collector.log_stats(self.stats)
self._emit_kv_metrics()
self._publish_kv_events()
self.emit_kv_metrics(self.kv_events_publisher)
self.publish_kv_events(self.kv_events_publisher)
def log_batch_result_stats(
self: Scheduler,
@@ -824,20 +817,21 @@ class SchedulerMetricsMixin:
balancedness=m.eplb_balancedness.item(),
)
def _emit_kv_metrics(self: Scheduler):
@staticmethod
def emit_kv_metrics(self: "SchedulerKvEventsPublisher"):
if not self.enable_kv_cache_events:
return
kv_metrics = KvMetrics()
kv_metrics.request_active_slots = self.stats.num_running_reqs.total
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.stats.token_usage * self.max_total_num_tokens
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.stats.num_queue_reqs.total
kv_metrics.gpu_cache_usage_perc = self.stats.token_usage
kv_metrics.gpu_prefix_cache_hit_rate = self.stats.cache_hit_rate
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
)
@@ -845,7 +839,8 @@ class SchedulerMetricsMixin:
if not self.send_metrics_from_scheduler.closed:
self.send_metrics_from_scheduler.send_pyobj(kv_metrics)
def _publish_kv_events(self: Scheduler):
@staticmethod
def publish_kv_events(self: "SchedulerKvEventsPublisher"):
if not self.enable_kv_cache_events:
return