From 163bf1ba7143fe612124598f485b670fdd5bf1c1 Mon Sep 17 00:00:00 2001 From: cctry Date: Wed, 6 May 2026 03:44:48 -0700 Subject: [PATCH] [PD] Fix KV transfer metrics (#24416) Co-authored-by: Lianmin Zheng Co-authored-by: Shangming Cai --- python/sglang/srt/disaggregation/base/conn.py | 13 +++ .../sglang/srt/disaggregation/common/conn.py | 24 ++++++ python/sglang/srt/disaggregation/fake/conn.py | 4 + .../srt/disaggregation/mooncake/conn.py | 15 ++-- python/sglang/srt/disaggregation/mori/conn.py | 1 + python/sglang/srt/disaggregation/nixl/conn.py | 14 ++++ python/sglang/srt/disaggregation/prefill.py | 15 ++-- .../srt/observability/req_time_stats.py | 82 ++++++++++++------- 8 files changed, 122 insertions(+), 46 deletions(-) diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py index 2b9ddde75..87244e044 100644 --- a/python/sglang/srt/disaggregation/base/conn.py +++ b/python/sglang/srt/disaggregation/base/conn.py @@ -1,5 +1,6 @@ from __future__ import annotations +import dataclasses from abc import ABC, abstractmethod from typing import TYPE_CHECKING, List, Optional @@ -12,6 +13,13 @@ if TYPE_CHECKING: from sglang.srt.disaggregation.utils import DisaggregationMode +@dataclasses.dataclass +class KVTransferMetric: + # Backends that cannot isolate transfer latency can leave this as None. + transfer_latency_s: Optional[float] = None + transfer_total_bytes: Optional[int] = None + + class KVArgs: engine_rank: int kv_data_ptrs: List[int] @@ -101,6 +109,11 @@ class BaseKVSender(ABC): def should_send_kv_chunk(self, num_pages: int, last_chunk: bool) -> bool: return num_pages > 0 + @abstractmethod + def get_transfer_metric(self) -> KVTransferMetric: + """Return backend-specific transfer metrics for this sender.""" + ... + @abstractmethod def poll(self) -> KVPoll: """ diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 3c916229a..79784ffce 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -22,6 +22,7 @@ from sglang.srt.disaggregation.base.conn import ( BaseKVSender, KVArgs, KVPoll, + KVTransferMetric, ) from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.distributed import get_pp_group @@ -94,6 +95,8 @@ class CommonKVManager(BaseKVManager): is_mla_backend: Optional[bool] = False, ): self.kv_args = args + self.kv_item_lens_sum = sum(args.kv_item_lens) + self.state_item_lens_sum = sum(args.state_item_lens) self.is_mla_backend = is_mla_backend self.disaggregation_mode = disaggregation_mode self.server_args = server_args @@ -442,6 +445,10 @@ class CommonKVSender(BaseKVSender): self.bootstrap_room = bootstrap_room self.aux_index = None self.bootstrap_server_url = bootstrap_addr + self.conclude_state: Optional[KVPoll] = None + self._transfer_metric = KVTransferMetric() + self._transfer_num_kv_indices = 0 + self._transfer_num_state_indices = 0 # inner state self.curr_idx = 0 if self.kv_mgr.is_dummy_cp_rank: @@ -502,6 +509,23 @@ class CommonKVSender(BaseKVSender): def should_send_kv_chunk(self, num_pages: int, last_chunk: bool) -> bool: return num_pages > 0 + def get_transfer_metric(self) -> KVTransferMetric: + total_bytes = self._transfer_num_kv_indices * self.kv_mgr.kv_item_lens_sum + total_bytes += ( + self._transfer_num_state_indices * self.kv_mgr.state_item_lens_sum + ) + self._transfer_metric.transfer_total_bytes = total_bytes + return self._transfer_metric + + def _record_transfer_indices( + self, + kv_indices: npt.NDArray[np.int32], + state_indices: Optional[List[int]], + ): + self._transfer_num_kv_indices += len(kv_indices) + if state_indices is not None: + self._transfer_num_state_indices += len(state_indices) + def send( self, kv_indices: npt.NDArray[np.int32], diff --git a/python/sglang/srt/disaggregation/fake/conn.py b/python/sglang/srt/disaggregation/fake/conn.py index 073faecfb..638834207 100644 --- a/python/sglang/srt/disaggregation/fake/conn.py +++ b/python/sglang/srt/disaggregation/fake/conn.py @@ -10,6 +10,7 @@ from sglang.srt.disaggregation.base.conn import ( BaseKVSender, KVArgs, KVPoll, + KVTransferMetric, ) from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.server_args import ServerArgs @@ -54,6 +55,9 @@ class FakeKVSender(BaseKVSender): logger.debug("FakeKVSender poll success") return KVPoll.Success + def get_transfer_metric(self) -> KVTransferMetric: + return KVTransferMetric() + def init( self, kv_indices: list[int], diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index fe0d65103..2bd1fe461 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -21,6 +21,12 @@ from sglang.srt.disaggregation.common.conn import ( CommonKVReceiver, CommonKVSender, ) +from sglang.srt.disaggregation.common.staging_handler import ( + DecodeStagingContext, + PrefillStagingContext, + StagingRegisterInfo, + StagingTransferInfo, +) from sglang.srt.disaggregation.common.utils import ( FastQueue, group_concurrent_contiguous, @@ -61,14 +67,6 @@ class TransferKVChunk: state_indices: Optional[List[int]] -from sglang.srt.disaggregation.common.staging_handler import ( - DecodeStagingContext, - PrefillStagingContext, - StagingRegisterInfo, - StagingTransferInfo, -) - - # decode @dataclasses.dataclass class TransferInfo: @@ -1710,6 +1708,7 @@ class MooncakeKVSender(CommonKVSender): aux_index=self.aux_index, state_indices=state_indices, ) + self._record_transfer_indices(kv_indices, state_indices) def poll(self) -> KVPoll: if self.conclude_state is None: diff --git a/python/sglang/srt/disaggregation/mori/conn.py b/python/sglang/srt/disaggregation/mori/conn.py index 6226c19df..d657c4c68 100644 --- a/python/sglang/srt/disaggregation/mori/conn.py +++ b/python/sglang/srt/disaggregation/mori/conn.py @@ -882,6 +882,7 @@ class MoriKVSender(CommonKVSender): aux_index=self.aux_index if is_last else None, ) self.transfer_statuses.extend(statuses) + self._record_transfer_indices(kv_indices, None) if infos is not None: self.pending_infos = infos self.sent_last_chunk = True diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index ff706f336..ceca38782 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -1123,6 +1123,7 @@ class NixlKVSender(CommonKVSender): self.chunk_id = 0 self._send_failed = False self._send_error: Optional[Exception] = None + self._transfer_start_time: Optional[float] = None def pop_decode_prefix_len(self) -> int: return self.kv_mgr.req_to_decode_prefix_len.pop(self.bootstrap_room, 0) @@ -1156,6 +1157,11 @@ class NixlKVSender(CommonKVSender): self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Success) return + if self._transfer_start_time is None and ( + len(kv_indices) > 0 or state_indices is not None + ): + self._transfer_start_time = time.perf_counter() + try: new_xfer_handles = self.kv_mgr.add_transfer_request( self.bootstrap_room, @@ -1174,6 +1180,7 @@ class NixlKVSender(CommonKVSender): self._send_error = e return + self._record_transfer_indices(kv_indices, state_indices) self.xfer_handles.extend(new_xfer_handles) self.chunk_id += 1 if is_last: @@ -1195,6 +1202,13 @@ class NixlKVSender(CommonKVSender): self._send_error = e return KVPoll.Failed # type: ignore if all(x == "DONE" for x in states): + if ( + self._transfer_start_time is not None + and self._transfer_metric.transfer_latency_s is None + ): + self._transfer_metric.transfer_latency_s = ( + time.perf_counter() - self._transfer_start_time + ) return KVPoll.Success # type: ignore if any(x == "ERR" for x in states): self._send_failed = True diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 1619e3e28..1a089e8ff 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -654,19 +654,16 @@ class SchedulerDisaggregationPrefillMixin: for req in done_reqs: req.time_stats.set_completion_time() - page_size = self.token_to_kv_pool_allocator.page_size - kv_item_lens = ( - self.disagg_prefill_bootstrap_queue.kv_manager.kv_args.kv_item_lens - ) - bytes_per_page_all_layers = sum(kv_item_lens) - for req in done_reqs: if isinstance(req.finished_reason, FINISH_ABORT): continue + if req.bootstrap_host == FAKE_BOOTSTRAP_HOST: + continue + kv_mgr = getattr(req.disagg_kv_sender, "kv_mgr", None) + if kv_mgr and getattr(kv_mgr, "is_dummy_cp_rank", False): + continue metrics = req.time_stats.compute_and_observe_kv_transfer_metrics( - num_tokens=len(req.origin_input_ids), - page_size=page_size, - bytes_per_page_all_layers=bytes_per_page_all_layers, + req.disagg_kv_sender.get_transfer_metric() ) if metrics: # Update last-value for REST API diff --git a/python/sglang/srt/observability/req_time_stats.py b/python/sglang/srt/observability/req_time_stats.py index 0710fce85..0ad222392 100644 --- a/python/sglang/srt/observability/req_time_stats.py +++ b/python/sglang/srt/observability/req_time_stats.py @@ -37,6 +37,7 @@ from sglang.srt.observability.trace import ( from sglang.srt.utils import get_bool_env_var if TYPE_CHECKING: + from sglang.srt.disaggregation.base.conn import KVTransferMetric from sglang.srt.managers.schedule_batch import ScheduleBatch SGLANG_TEST_REQUEST_TIME_STATS = get_bool_env_var("SGLANG_TEST_REQUEST_TIME_STATS") @@ -842,27 +843,30 @@ class SchedulerReqTimeStats(ReqTimeStatsBase): def compute_and_observe_kv_transfer_metrics( self, - num_tokens: int, - page_size: int, - bytes_per_page_all_layers: int, + transfer_metric: KVTransferMetric, ) -> Optional[dict]: """Compute KV transfer metrics and observe them via the metrics collector. Returns a dict with latency_ms, total_mb, speed_gb_s if computable, else None. """ - from sglang.srt.mem_cache.common import kv_to_page_num - result = {} + if transfer_metric.transfer_total_bytes is None: + return result if result else None # Transfer latency, size, and speed - if self.prefill_transfer_queue_entry_time > 0 and self.completion_time > 0: + if transfer_metric.transfer_latency_s is not None: + transfer_latency_s = transfer_metric.transfer_latency_s + else: + if self.prefill_transfer_queue_entry_time <= 0 or self.completion_time <= 0: + return result if result else None transfer_latency_s = ( self.completion_time - self.prefill_transfer_queue_entry_time ) + + if transfer_latency_s > 0: latency_ms = transfer_latency_s * 1000 - num_pages = kv_to_page_num(num_tokens, page_size) - total_bytes = bytes_per_page_all_layers * num_pages + total_bytes = transfer_metric.transfer_total_bytes total_mb = total_bytes / (1024 * 1024) self.transfer_total_mb = total_mb @@ -980,8 +984,12 @@ class SchedulerReqTimeStats(ReqTimeStatsBase): def convert_to_duration(self) -> str: if self.disagg_mode == DisaggregationMode.NULL: - queue_duration = self.forward_entry_time - self.wait_queue_entry_time - forward_duration = self.completion_time - self.forward_entry_time + queue_duration = self.duration_between( + self.wait_queue_entry_time, self.forward_entry_time + ) + forward_duration = self.duration_between( + self.forward_entry_time, self.completion_time + ) if SGLANG_TEST_REQUEST_TIME_STATS: assert ( @@ -990,11 +998,15 @@ class SchedulerReqTimeStats(ReqTimeStatsBase): return f"queue_duration={self.format_duration(queue_duration)}, forward_duration={self.format_duration(forward_duration)}, start_time={self.wait_queue_entry_time:.3f}" elif self.disagg_mode == DisaggregationMode.PREFILL: - bootstrap_queue_duration = ( - self.wait_queue_entry_time - self.prefill_bootstrap_queue_entry_time + bootstrap_queue_duration = self.duration_between( + self.prefill_bootstrap_queue_entry_time, self.wait_queue_entry_time + ) + queue_duration = self.duration_between( + self.wait_queue_entry_time, self.forward_entry_time + ) + forward_duration = self.duration_between( + self.forward_entry_time, self.completion_time ) - queue_duration = self.forward_entry_time - self.wait_queue_entry_time - forward_duration = self.completion_time - self.forward_entry_time if SGLANG_TEST_REQUEST_TIME_STATS: if self.wait_queue_entry_time > 0: @@ -1006,11 +1018,11 @@ class SchedulerReqTimeStats(ReqTimeStatsBase): # Break down bootstrap_queue_duration into sub-phases if self.bootstrap_done_time > 0: - bootstrap_duration = ( - self.bootstrap_done_time - self.prefill_bootstrap_queue_entry_time + bootstrap_duration = self.duration_between( + self.prefill_bootstrap_queue_entry_time, self.bootstrap_done_time ) - alloc_wait_duration = ( - self.wait_queue_entry_time - self.bootstrap_done_time + alloc_wait_duration = self.duration_between( + self.bootstrap_done_time, self.wait_queue_entry_time ) if SGLANG_TEST_REQUEST_TIME_STATS: assert ( @@ -1034,15 +1046,22 @@ class SchedulerReqTimeStats(ReqTimeStatsBase): f"#retries={self.prefill_retry_count}" ) elif self.disagg_mode == DisaggregationMode.DECODE: - prealloc_duration = ( - self.decode_transfer_queue_entry_time - - self.decode_prealloc_queue_entry_time + prealloc_duration = self.duration_between( + self.decode_prealloc_queue_entry_time, + self.decode_transfer_queue_entry_time, ) - transfer_duration = ( - self.wait_queue_entry_time - self.decode_transfer_queue_entry_time + transfer_duration = self.duration_between( + self.decode_transfer_queue_entry_time, + self.wait_queue_entry_time, + ) + queue_duration = self.duration_between( + self.wait_queue_entry_time, + self.forward_entry_time, + ) + forward_duration = self.duration_between( + self.forward_entry_time, + self.completion_time, ) - queue_duration = self.forward_entry_time - self.wait_queue_entry_time - forward_duration = self.completion_time - self.forward_entry_time if SGLANG_TEST_REQUEST_TIME_STATS: if self.wait_queue_entry_time > 0: @@ -1055,11 +1074,11 @@ class SchedulerReqTimeStats(ReqTimeStatsBase): # Break down prealloc_duration into sub-phases if self.bootstrap_done_time > 0: - bootstrap_duration = ( - self.bootstrap_done_time - self.decode_prealloc_queue_entry_time + bootstrap_duration = self.duration_between( + self.decode_prealloc_queue_entry_time, self.bootstrap_done_time ) - alloc_wait_duration = ( - self.decode_transfer_queue_entry_time - self.bootstrap_done_time + alloc_wait_duration = self.duration_between( + self.bootstrap_done_time, self.decode_transfer_queue_entry_time ) if SGLANG_TEST_REQUEST_TIME_STATS: assert ( @@ -1105,6 +1124,11 @@ class SchedulerReqTimeStats(ReqTimeStatsBase): def format_duration(self, duration: float) -> str: return f"{duration * 1e3:.2f}ms" + def duration_between(self, start: float, end: float) -> float: + if start <= 0 or end <= 0: + return 0.0 + return end - start + def set_schedule_time_batch(batch: ScheduleBatch): # only for tracing