From 243f1b964c623a1982ceb258489453354826c928 Mon Sep 17 00:00:00 2001 From: leavelet Date: Sun, 7 Jun 2026 08:59:02 +0000 Subject: [PATCH] fix(cp_per_layer_transfer): gate per-layer path on a single non-dummy decode info The transfer worker iterates every non-dummy decode info for a room and calls per_layer_mgr.finish() once per info, but register_per_layer_transfer registers exactly one context per room/chunk (built for one info's dst_kv_indices). This is only sound when there is exactly one non-dummy info (required_dst_info_num == 1). With decode attn_tp < prefill attn_tp a single prefill rank holds >1 non-dummy infos; finishing once-per-info would over-pop chunk contexts and under-deliver KV to the other infos. Make the assumption explicit: register only when there is one non-dummy info, otherwise fall back to the monolithic post-forward transfer (which fans out to all infos correctly). Found by an independent first-principles audit. Adds TestRegisterGuardSingleInfo (2-info fallback, 1-info register, all-dummy fallback) exercising the real MooncakeKVManager.register_per_layer_transfer. Co-Authored-By: Claude Opus 4.8 (1M context) --- .../srt/disaggregation/mooncake/conn.py | 50 ++++++++------ .../test_cp_per_layer_transfer.py | 66 +++++++++++++++++++ 2 files changed, 96 insertions(+), 20 deletions(-) diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index e93a3b2e7..bc379c6ea 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -408,6 +408,20 @@ class MooncakeKVManager(CommonKVManager): infos = self.transfer_infos.get(room) if not infos: return False + # The per-layer path registers exactly ONE context per room/chunk, but the + # transfer worker iterates every non-dummy decode info for the room and calls + # finish() per info (conn.py reqs_to_be_processed loop). That is only sound + # when there is exactly one non-dummy info (required_dst_info_num == 1). For + # decode attn_tp < prefill attn_tp a single prefill rank holds >1 non-dummy + # infos; finishing once-per-info would over-pop chunk contexts and the single + # ctx only carries one info's dst_kv_indices. Fall back to the monolithic + # post-forward transfer (which fans out to all infos) in that case. + non_dummy = [ + info for info in infos.values() if not getattr(info, "is_dummy", False) + ] + if len(non_dummy) != 1: + return False + info = non_dummy[0] from sglang.srt.disaggregation.utils import filter_kv_pages_for_cp_shared_kv from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout @@ -425,27 +439,23 @@ class MooncakeKVManager(CommonKVManager): # page_positions]). Must offset by the chunk's absolute page start, else chunk # N>0 writes its KV onto chunk 0's decode pages. page_size divides chunk_key. chunk_page_start = int(chunk_key) // self.kv_args.page_size - for info in infos.values(): - if getattr(info, "is_dummy", False): - continue - owned_pages, positions = filter_kv_pages_for_cp_shared_kv( - layout=layout, logical_pages=pages, chunk_page_start=chunk_page_start - ) - dst_indices = np.asarray(info.dst_kv_indices, dtype=np.int32)[positions] - ctx = self.build_per_layer_context( - info.mooncake_session_id, owned_pages, dst_indices - ) - if ctx is not None: - mgr.register(room, ctx, chunk_key=chunk_key) - logger.info( - "[CP_PER_LAYER_TRANSFER] registered room=%s chunk=%s owned_pages=%d", - room, - chunk_key, - len(owned_pages), - ) - return True + owned_pages, positions = filter_kv_pages_for_cp_shared_kv( + layout=layout, logical_pages=pages, chunk_page_start=chunk_page_start + ) + dst_indices = np.asarray(info.dst_kv_indices, dtype=np.int32)[positions] + ctx = self.build_per_layer_context( + info.mooncake_session_id, owned_pages, dst_indices + ) + if ctx is None: return False - return False + mgr.register(room, ctx, chunk_key=chunk_key) + logger.info( + "[CP_PER_LAYER_TRANSFER] registered room=%s chunk=%s owned_pages=%d", + room, + chunk_key, + len(owned_pages), + ) + return True def _send_kvcache_generic( self, diff --git a/test/registered/unit/disaggregation/test_cp_per_layer_transfer.py b/test/registered/unit/disaggregation/test_cp_per_layer_transfer.py index 020edf465..dd87e7879 100644 --- a/test/registered/unit/disaggregation/test_cp_per_layer_transfer.py +++ b/test/registered/unit/disaggregation/test_cp_per_layer_transfer.py @@ -395,5 +395,71 @@ class TestChunkedDstMapping(unittest.TestCase): self.fail("no rank owned a chunk-1 page; adjust fixture") +class TestRegisterGuardSingleInfo(unittest.TestCase): + """register_per_layer_transfer must register a context only when there is + exactly one non-dummy decode info. The transfer worker calls finish() once + per non-dummy info for the room while the per-layer path registers a single + context, so >1 non-dummy info (decode attn_tp < prefill attn_tp) must fall + back to the monolithic transfer.""" + + def _fake_self(self, infos, recorder): + import types + + mgr = types.SimpleNamespace( + register=lambda room, ctx, chunk_key=0: recorder.append( + (room, chunk_key) + ), + ) + return types.SimpleNamespace( + per_layer_transfer_manager=mgr, + server_args=types.SimpleNamespace(enable_nsa_prefill_cp_shared_kv=True), + transfer_infos={7: infos}, + kv_args=types.SimpleNamespace(page_size=64), + attn_cp_size=8, + attn_cp_rank=0, + build_per_layer_context=lambda sid, owned, dst: ("ctx", sid), + ) + + def _info(self, is_dummy): + import types + + import numpy as np + + return types.SimpleNamespace( + is_dummy=is_dummy, + dst_kv_indices=np.arange(64, dtype=np.int32), + mooncake_session_id="s", + ) + + def _call(self, infos, recorder): + try: + from sglang.srt.disaggregation.mooncake.conn import MooncakeKVManager + except Exception as e: # pragma: no cover - env without mooncake + self.skipTest(f"mooncake conn not importable: {e}") + s = self._fake_self(infos, recorder) + # logical pages incl. page 1 which cp_rank 0 owns ((1-1)%8==0) + return MooncakeKVManager.register_per_layer_transfer( + s, 7, [0, 1, 2, 3], chunk_key=0 + ) + + def test_two_non_dummy_infos_falls_back(self): + rec = [] + infos = {"a": self._info(False), "b": self._info(False)} + self.assertFalse(self._call(infos, rec)) + self.assertEqual(rec, []) # nothing registered + + def test_single_non_dummy_info_registers(self): + rec = [] + infos = {"a": self._info(True), "b": self._info(False)} # 1 dummy + 1 real + self.assertTrue(self._call(infos, rec)) + self.assertEqual(len(rec), 1) + + def test_all_dummy_falls_back(self): + rec = [] + infos = {"a": self._info(True), "b": self._info(True)} + self.assertFalse(self._call(infos, rec)) + self.assertEqual(rec, []) + + if __name__ == "__main__": unittest.main()