diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 8062123e3..261f818b0 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -127,6 +127,11 @@ class CommonKVManager(BaseKVManager): self.kv_item_lens_sum = sum(args.kv_item_lens) self.state_item_lens_sum = sum(x for comp in args.state_item_lens for x in comp) self.is_mla_backend = is_mla_backend + # Per-sender fan-out of a KV copy onto N decode destinations + # (MLA under Prefill-CP + Decode-TP, or decode_tp > prefill_tp). + # MLA is resolved lazily at bootstrap (see resolve_kv_replica_factor); + # MHA never replicates, so it stays pinned at 1. + self._kv_replica_factor: Optional[int] = None if is_mla_backend else 1 self.is_hybrid_mla_backend = getattr(args, "is_hybrid_mla_backend", False) self.disaggregation_mode = disaggregation_mode self.server_args = server_args @@ -263,6 +268,30 @@ class CommonKVManager(BaseKVManager): with self.failure_lock: self.failure_records[bootstrap_room] = failure_reason + def get_kv_replica_factor(self) -> int: + if self._kv_replica_factor is None: + logger.warning_once( + "get_kv_replica_factor called before resolve_kv_replica_factor; " + "assuming 1, but the metrics may be inaccurate." + ) + return 1 + return self._kv_replica_factor + + def resolve_kv_replica_factor(self, transfer_infos: Dict) -> None: + if not self.is_mla_backend: + # Only MLA replicates its KV across decode ranks; non-MLA head slices are + # disjoint and stay pinned at the factor of 1 set in __init__. + return + + info = next(iter(transfer_infos.values()), None) + if info is None or info.required_dst_info_num is None: + logger.warning_once( + "resolve_kv_replica_factor: no decode destinations available; " + "KV transfer metrics may be inaccurate." + ) + return + self._kv_replica_factor = info.required_dst_info_num + def _ensure_prefill_recompute_executor( self, ) -> concurrent.futures.ThreadPoolExecutor: @@ -1033,6 +1062,8 @@ class CommonKVSender(BaseKVSender): total_bytes += ( self._transfer_num_state_indices * self.kv_mgr.state_item_lens_sum ) + # Pinned to 1 for MHA (disjoint slices); only MLA replication makes it > 1. + total_bytes *= self.kv_mgr.get_kv_replica_factor() self._transfer_metric.transfer_total_bytes = total_bytes return self._transfer_metric diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 200646d09..f14d98aba 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -1560,6 +1560,7 @@ class MooncakeKVManager(CommonKVManager): ) # NOTE: after bootstrapping we can mark the req as waiting for input if len(self.transfer_infos[room]) == required_dst_info_num: + self.resolve_kv_replica_factor(self.transfer_infos[room]) self.req_to_decode_prefix_len[room] = next( ( info.decode_prefix_len diff --git a/python/sglang/srt/disaggregation/mori/conn.py b/python/sglang/srt/disaggregation/mori/conn.py index 61835a9f2..c8f90f987 100644 --- a/python/sglang/srt/disaggregation/mori/conn.py +++ b/python/sglang/srt/disaggregation/mori/conn.py @@ -506,6 +506,7 @@ class MoriKVManager(CommonKVManager): infos[transfer_info.engine_key] = transfer_info if len(infos) >= transfer_info.required_dst_info_num: + self.resolve_kv_replica_factor(infos) # All decode peers reported their dst metadata; pick a # non-None decode_prefix_len if any peer set it (they # should all agree, but be defensive). 0 means "no diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index f4a7529f9..1f2597178 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -2412,6 +2412,7 @@ class NixlKVManager(CommonKVManager): ].required_dst_info_num logger.debug(f"got info {room=} {agent_name=} {required_dst_info_num=}") if len(self.transfer_infos[room]) == required_dst_info_num: + self.resolve_kv_replica_factor(self.transfer_infos[room]) self.req_to_decode_prefix_len[room] = next( ( info.decode_prefix_len diff --git a/test/registered/unit/disaggregation/test_kv_transfer_replica_metric.py b/test/registered/unit/disaggregation/test_kv_transfer_replica_metric.py new file mode 100644 index 000000000..ed17cc085 --- /dev/null +++ b/test/registered/unit/disaggregation/test_kv_transfer_replica_metric.py @@ -0,0 +1,104 @@ +"""Unit tests for KV-transfer replication accounting in PD disaggregation. + +In Prefill-CP + Decode-TP (e.g. prefill CP8, decode TP8), by default only prefill +CP rank 0 transfers, replicating its KV to every decode TP rank. get_transfer_metric() +must report bytes put on the wire (logical KV size x fan-out). The fan-out equals +required_dst_info_num -- a topological invariant -- so it is resolved once and cached +on the shared CommonKVManager. +""" + +import unittest +from types import SimpleNamespace + +import numpy as np + +from sglang.srt.disaggregation.base.conn import KVTransferMetric +from sglang.srt.disaggregation.common.conn import CommonKVManager, CommonKVSender +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + +KV_ITEM_LENS_SUM = 100 +STATE_ITEM_LENS_SUM = 7 + + +def _room(fan_out): + """A room of `fan_out` destinations, each reporting required_dst_info_num. + + All infos in a room report the same required_dst_info_num (the registration + barrier waits until that many arrive), so the resolver reads it off any one. + """ + return { + f"sess{i}": SimpleNamespace(required_dst_info_num=fan_out) + for i in range(fan_out) + } + + +def _make_kv_mgr(is_mla_backend): + """CommonKVManager bypassing __init__, wiring only the fields the path reads.""" + mgr = CommonKVManager.__new__(CommonKVManager) + mgr.is_mla_backend = is_mla_backend + mgr.kv_item_lens_sum = KV_ITEM_LENS_SUM + mgr.state_item_lens_sum = STATE_ITEM_LENS_SUM + mgr._kv_replica_factor = None if is_mla_backend else 1 + return mgr + + +def _make_sender(kv_mgr): + """CommonKVSender bypassing __init__, wiring only the fields the path reads.""" + sender = CommonKVSender.__new__(CommonKVSender) + sender._transfer_metric = KVTransferMetric() + sender._transfer_num_kv_indices = 0 + sender._transfer_num_state_indices = 0 + sender.kv_mgr = kv_mgr + return sender + + +class TestKVTransferReplicaMetric(CustomTestCase): + def test_mla_scales_kv_and_state_bytes_by_fan_out(self): + # CP rank 0 replicates to 4 decode TP ranks; both kv and state scale. + mgr = _make_kv_mgr(is_mla_backend=True) + sender = _make_sender(mgr) + + mgr.resolve_kv_replica_factor(_room(4)) + self.assertEqual(mgr._kv_replica_factor, 4) + + sender._record_transfer_indices( + np.arange(8, dtype=np.int32), [np.arange(5, dtype=np.int32)] + ) + expected = (8 * KV_ITEM_LENS_SUM + 5 * STATE_ITEM_LENS_SUM) * 4 + self.assertEqual(sender.get_transfer_metric().transfer_total_bytes, expected) + + def test_non_mla_factor_is_one_regardless_of_destinations(self): + # Non-MLA head slices sum to one logical copy: factor stays 1. + mgr = _make_kv_mgr(is_mla_backend=False) + sender = _make_sender(mgr) + + mgr.resolve_kv_replica_factor(_room(8)) + self.assertEqual(mgr._kv_replica_factor, 1) + + sender._record_transfer_indices(np.arange(6, dtype=np.int32), None) + self.assertEqual( + sender.get_transfer_metric().transfer_total_bytes, 6 * KV_ITEM_LENS_SUM + ) + + def test_unresolved_factor_does_not_crash_metric(self): + # An empty room or a missing required_dst_info_num leaves the factor + # unresolved. get_transfer_metric() must still compute -- it must not + # multiply bytes by None. + for room in ({}, {"sess0": SimpleNamespace(required_dst_info_num=None)}): + mgr = _make_kv_mgr(is_mla_backend=True) + sender = _make_sender(mgr) + + mgr.resolve_kv_replica_factor(room) + self.assertIsNone(mgr._kv_replica_factor) + + sender._record_transfer_indices(np.arange(4, dtype=np.int32), None) + self.assertIsInstance( + sender.get_transfer_metric().transfer_total_bytes, int + ) + + +if __name__ == "__main__": + unittest.main()