[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:
co-authored by
Lianmin Zheng
Shangming Cai
parent
11b0e510aa
commit
163bf1ba71
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user