Files
sglang/test/registered/unit/disaggregation/test_kv_transfer_metrics.py
leavelet c6c88c2617 fix(disagg): account KV transfer metrics by actually-sent bytes
Port upstream sgl-project/sglang #24416 (staging-free). Our transfer metrics
were computed from the full prompt length, so total_mb / speed_gb_s were
systematically over-reported on every prefix-cache hit and every CP shared-KV
per-rank page filter. Now each sender accumulates the actually-sent KV/state
indices and reports bytes = sent_pages * per-page item bytes via a new
KVTransferMetric returned by get_transfer_metric(); prefill consumes that
instead of estimating from len(origin_input_ids), and skips fake-bootstrap and
dummy-CP-rank senders (which transfer nothing).

Also route convert_to_duration() phase deltas through duration_between(), which
returns 0 when a phase timestamp is uninitialized, fixing nonsensical negative
durations in the time-stats log.

Unit-tested (test_kv_transfer_metrics.py, 10 cases) in the CUDA-13 container.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-06-07 09:51:04 +00:00

123 lines
4.7 KiB
Python

"""Unit tests for the KV-transfer metric accounting (port of upstream #24416).
These exercise the pure metric logic without standing up a real PD transfer:
- bytes are counted from the *actually-sent* indices (not prompt length),
- duration_between guards uninitialized (<=0) timestamps,
- compute_and_observe_kv_transfer_metrics consumes a KVTransferMetric.
"""
import unittest
import numpy as np
from sglang.srt.disaggregation.base.conn import KVTransferMetric
from sglang.srt.disaggregation.common.conn import CommonKVSender
from sglang.srt.observability.req_time_stats import SchedulerReqTimeStats
class _FakeMgr:
def __init__(self, kv_sum, state_sum):
self.kv_item_lens_sum = kv_sum
self.state_item_lens_sum = state_sum
def _make_sender(kv_sum, state_sum):
s = CommonKVSender.__new__(CommonKVSender) # bypass heavy __init__
s.kv_mgr = _FakeMgr(kv_sum, state_sum)
s._transfer_metric = KVTransferMetric()
s._transfer_num_kv_indices = 0
s._transfer_num_state_indices = 0
return s
def _make_stats():
st = SchedulerReqTimeStats.__new__(SchedulerReqTimeStats)
st.enable_metrics = False
st.transfer_total_mb = 0.0
st.transfer_speed_gb_s = 0.0
st.prefill_transfer_queue_entry_time = 0.0
st.completion_time = 0.0
st.prefill_bootstrap_queue_entry_time = 0.0
st.bootstrap_done_time = 0.0
st.wait_queue_entry_time = 0.0
return st
class TestKVTransferMetricBytes(unittest.TestCase):
def test_records_actually_sent_indices(self):
s = _make_sender(kv_sum=100, state_sum=10)
s._record_transfer_indices(np.arange(5, dtype=np.int32), [1, 2, 3])
s._record_transfer_indices(np.arange(2, dtype=np.int32), None)
m = s.get_transfer_metric()
# 7 kv pages * 100 + 3 state pages * 10
self.assertEqual(m.transfer_total_bytes, 7 * 100 + 3 * 10)
def test_no_state_pool(self):
s = _make_sender(kv_sum=64, state_sum=0)
s._record_transfer_indices(np.arange(4, dtype=np.int32), None)
self.assertEqual(s.get_transfer_metric().transfer_total_bytes, 4 * 64)
def test_bytes_track_sent_pages_not_prompt_length(self):
# The whole point of #24416: under cache hits / CP shared-KV filtering only
# a subset of pages is sent, and the metric must reflect that subset.
s = _make_sender(kv_sum=100, state_sum=0)
s._record_transfer_indices(np.arange(3, dtype=np.int32), None)
self.assertEqual(s.get_transfer_metric().transfer_total_bytes, 3 * 100)
class TestDurationBetween(unittest.TestCase):
def test_guards_nonpositive(self):
st = _make_stats()
self.assertEqual(st.duration_between(0, 5), 0.0)
self.assertEqual(st.duration_between(5, 0), 0.0)
self.assertEqual(st.duration_between(-1.0, 5.0), 0.0)
def test_positive(self):
st = _make_stats()
self.assertAlmostEqual(st.duration_between(2.0, 5.0), 3.0)
class TestComputeAndObserve(unittest.TestCase):
def test_none_bytes_returns_none(self):
st = _make_stats()
m = KVTransferMetric(transfer_latency_s=1.0, transfer_total_bytes=None)
self.assertIsNone(st.compute_and_observe_kv_transfer_metrics(m))
def test_explicit_latency_and_bytes(self):
st = _make_stats()
mb = 1024 * 1024
m = KVTransferMetric(transfer_latency_s=2.0, transfer_total_bytes=20 * mb)
out = st.compute_and_observe_kv_transfer_metrics(m)
self.assertAlmostEqual(out["total_mb"], 20.0)
self.assertAlmostEqual(out["latency_ms"], 2000.0)
self.assertAlmostEqual(out["speed_gb_s"], (20.0 / 1024) / 2.0)
def test_latency_derived_from_timestamps(self):
st = _make_stats()
st.prefill_transfer_queue_entry_time = 1.0
st.completion_time = 1.5
m = KVTransferMetric(transfer_latency_s=None, transfer_total_bytes=1024 * 1024)
out = st.compute_and_observe_kv_transfer_metrics(m)
self.assertAlmostEqual(out["latency_ms"], 500.0)
def test_missing_timestamps_returns_none(self):
st = _make_stats() # both timestamps are 0
m = KVTransferMetric(transfer_latency_s=None, transfer_total_bytes=1024)
self.assertIsNone(st.compute_and_observe_kv_transfer_metrics(m))
class TestSendersImplementMetric(unittest.TestCase):
def test_senders_have_get_transfer_metric(self):
from sglang.srt.disaggregation.fake.conn import FakeKVSender
from sglang.srt.disaggregation.mooncake.conn import MooncakeKVSender
for cls in (CommonKVSender, FakeKVSender, MooncakeKVSender):
self.assertTrue(
callable(getattr(cls, "get_transfer_metric", None)),
f"{cls.__name__} missing get_transfer_metric",
)
if __name__ == "__main__":
unittest.main()