[Bug fix] Account for KV replication fan-out in transfer-byte metrics (#30351)
Co-authored-by: Hun-ho Kim <hunho.kim@samsung.com>
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user