[PD] Fix KV transfer metrics (#24416)

Co-authored-by: Lianmin Zheng <lianminzheng@gmail.com>
Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
cctry
2026-05-06 03:44:48 -07:00
committed by GitHub
co-authored by Lianmin Zheng Shangming Cai
parent 11b0e510aa
commit 163bf1ba71
8 changed files with 122 additions and 46 deletions
@@ -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:
"""
@@ -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],
@@ -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],
@@ -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:
@@ -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
@@ -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
+6 -9
View File
@@ -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
@@ -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