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>
123 lines
4.7 KiB
Python
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()
|