Move KV-cache event emission to SchedulerKvEventsPublisher (#25626)
This commit is contained in:
@@ -3140,7 +3140,7 @@ class Scheduler(
|
|||||||
self._maybe_log_idle_metrics()
|
self._maybe_log_idle_metrics()
|
||||||
|
|
||||||
# kv event publishing
|
# kv event publishing
|
||||||
self.publish_kv_events(self.kv_events_publisher)
|
self.kv_events_publisher.publish_kv_events()
|
||||||
|
|
||||||
# reset token ratio
|
# reset token ratio
|
||||||
self.new_token_ratio = self.init_new_token_ratio
|
self.new_token_ratio = self.init_new_token_ratio
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import dataclasses
|
import dataclasses
|
||||||
|
import time
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import (
|
from typing import (
|
||||||
TYPE_CHECKING,
|
TYPE_CHECKING,
|
||||||
@@ -11,6 +12,11 @@ from typing import (
|
|||||||
|
|
||||||
import zmq
|
import zmq
|
||||||
|
|
||||||
|
from sglang.srt.disaggregation.kv_events import (
|
||||||
|
EventPublisherFactory,
|
||||||
|
KVEventBatch,
|
||||||
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
||||||
@@ -48,8 +54,44 @@ class SchedulerKvEventsPublisher:
|
|||||||
kv_event_publisher: Any = None
|
kv_event_publisher: Any = None
|
||||||
|
|
||||||
def __post_init__(self) -> None:
|
def __post_init__(self) -> None:
|
||||||
from sglang.srt.observability.scheduler_metrics_mixin import (
|
self.init_kv_events(self.kv_events_config)
|
||||||
SchedulerMetricsMixin,
|
|
||||||
|
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 collections import defaultdict
|
||||||
from typing import TYPE_CHECKING, List, Optional, Tuple, Union
|
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.disaggregation.utils import DisaggregationMode
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.managers.io_struct import (
|
from sglang.srt.managers.io_struct import (
|
||||||
@@ -20,7 +19,6 @@ from sglang.srt.managers.io_struct import (
|
|||||||
SpeculativeMetrics,
|
SpeculativeMetrics,
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
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.managers.utils import GenerationBatchResult
|
||||||
from sglang.srt.observability.metrics_collector import (
|
from sglang.srt.observability.metrics_collector import (
|
||||||
DPCooperationInfo,
|
DPCooperationInfo,
|
||||||
@@ -36,9 +34,6 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.managers.schedule_batch import Req
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
from sglang.srt.managers.schedule_policy import PrefillAdder
|
from sglang.srt.managers.schedule_policy import PrefillAdder
|
||||||
from sglang.srt.managers.scheduler import EmbeddingBatchResult, Scheduler
|
from sglang.srt.managers.scheduler import EmbeddingBatchResult, Scheduler
|
||||||
from sglang.srt.managers.scheduler_components.kv_events_publisher import (
|
|
||||||
SchedulerKvEventsPublisher,
|
|
||||||
)
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -189,19 +184,6 @@ class SchedulerMetricsMixin:
|
|||||||
for r in getattr(dw, "draft_runner_list", []):
|
for r in getattr(dw, "draft_runner_list", []):
|
||||||
r.device_timer = timer
|
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):
|
def _init_fpm(self: Scheduler):
|
||||||
"""Initialize Forward Pass Metrics (FPM) publisher if configured."""
|
"""Initialize Forward Pass Metrics (FPM) publisher if configured."""
|
||||||
self.enable_fpm = False
|
self.enable_fpm = False
|
||||||
@@ -602,8 +584,8 @@ class SchedulerMetricsMixin:
|
|||||||
self.update_lora_metrics()
|
self.update_lora_metrics()
|
||||||
self._log_hicache_stats()
|
self._log_hicache_stats()
|
||||||
self.metrics_collector.log_stats(self.stats)
|
self.metrics_collector.log_stats(self.stats)
|
||||||
self.emit_kv_metrics(self.kv_events_publisher)
|
self.kv_events_publisher.emit_kv_metrics()
|
||||||
self.publish_kv_events(self.kv_events_publisher)
|
self.kv_events_publisher.publish_kv_events()
|
||||||
|
|
||||||
def report_decode_stats(
|
def report_decode_stats(
|
||||||
self: Scheduler,
|
self: Scheduler,
|
||||||
@@ -798,8 +780,8 @@ class SchedulerMetricsMixin:
|
|||||||
self.update_lora_metrics()
|
self.update_lora_metrics()
|
||||||
self._log_hicache_stats()
|
self._log_hicache_stats()
|
||||||
self.metrics_collector.log_stats(self.stats)
|
self.metrics_collector.log_stats(self.stats)
|
||||||
self.emit_kv_metrics(self.kv_events_publisher)
|
self.kv_events_publisher.emit_kv_metrics()
|
||||||
self.publish_kv_events(self.kv_events_publisher)
|
self.kv_events_publisher.publish_kv_events()
|
||||||
|
|
||||||
def log_batch_result_stats(
|
def log_batch_result_stats(
|
||||||
self: Scheduler,
|
self: Scheduler,
|
||||||
@@ -817,38 +799,6 @@ class SchedulerMetricsMixin:
|
|||||||
balancedness=m.eplb_balancedness.item(),
|
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(
|
def _emit_forward_pass_metrics(
|
||||||
self: Scheduler,
|
self: Scheduler,
|
||||||
batch: ScheduleBatch,
|
batch: ScheduleBatch,
|
||||||
|
|||||||
Reference in New Issue
Block a user