From 4dc9cda5f977e7cf86e980f9cb2f8f456b584fbb Mon Sep 17 00:00:00 2001 From: Yonghao Zhuang Date: Thu, 3 Sep 2026 15:42:21 -0700 Subject: [PATCH] [PD] Gate deferred decode KV release on backend capability (#37454) Co-authored-by: yhzhuang --- python/sglang/srt/disaggregation/base/conn.py | 2 ++ python/sglang/srt/disaggregation/decode.py | 1 + python/sglang/srt/disaggregation/fake/conn.py | 2 ++ .../test_decode_queue_cleanup.py | 29 ++++++++++++++++++- 4 files changed, 33 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py index 1dedb16eb..4996cc6eb 100644 --- a/python/sglang/srt/disaggregation/base/conn.py +++ b/python/sglang/srt/disaggregation/base/conn.py @@ -101,6 +101,8 @@ class KVPoll: class BaseKVManager(ABC): """Base class for managing transfer states""" + enable_deferred_decode_kv_release: bool = False + @abstractmethod def __init__( self, diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index c03483ec6..9e0ea879c 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -2326,6 +2326,7 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin): self.scheduler.hisparse_coordinator.request_finished(decode_req.req) if ( self.enable_deferred_kv_release + and decode_req.kv_receiver.kv_mgr.enable_deferred_decode_kv_release and decode_req.kv_receiver.abort_notified ): # Decode-initiated abort: a prefill write may still target diff --git a/python/sglang/srt/disaggregation/fake/conn.py b/python/sglang/srt/disaggregation/fake/conn.py index 15a76c039..9fc18af80 100644 --- a/python/sglang/srt/disaggregation/fake/conn.py +++ b/python/sglang/srt/disaggregation/fake/conn.py @@ -99,6 +99,8 @@ class FakeKVReceiver(BaseKVReceiver): bootstrap_addr: str, bootstrap_room: Optional[int] = None, ): + self.kv_mgr = mgr + self.abort_notified: bool = False self.bootstrap_done = False self.has_sent_metadata = False self.require_staging: bool = False diff --git a/test/registered/unit/disaggregation/test_decode_queue_cleanup.py b/test/registered/unit/disaggregation/test_decode_queue_cleanup.py index 0aabdbeef..98115b6b5 100644 --- a/test/registered/unit/disaggregation/test_decode_queue_cleanup.py +++ b/test/registered/unit/disaggregation/test_decode_queue_cleanup.py @@ -9,6 +9,7 @@ from sglang.srt.disaggregation.decode import ( DecodeTransferQueue, HiCacheRestoreResult, ) +from sglang.srt.disaggregation.fake.conn import FakeKVManager, FakeKVReceiver from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.managers.schedule_batch import FINISH_ABORT @@ -318,7 +319,7 @@ class TestDecodeQueueCleanup(CustomTestCase): @patch("sglang.srt.disaggregation.decode.release_kv_cache") @patch("sglang.srt.disaggregation.decode.prepare_abort") @patch("sglang.srt.disaggregation.decode.poll_and_all_reduce") - def test_transfer_failure_clears_receiver_before_removing_request( + def test_transfer_failure_cleanup_respects_deferred_release_gates( self, mock_poll, mock_prepare_abort, mock_release_kv_cache ): receiver = FakeReceiver() @@ -371,6 +372,32 @@ class TestDecodeQueueCleanup(CustomTestCase): req, queue.tree_cache, is_insert=False ) + receiver = FakeReceiver() + receiver.kv_mgr = FakeKVManager.__new__(FakeKVManager) + decode_req.kv_receiver = receiver + queue.queue = [decode_req] + queue.enable_deferred_kv_release = True + queue.req_to_metadata_buffer_idx_allocator.reset_mock() + mock_release_kv_cache.reset_mock() + + transferred = queue.pop_transferred() + + self.assertEqual(transferred, []) + self.assertEqual(queue.queue, []) + self.assertTrue(receiver.clear_called) + self.assertIsNone(decode_req.kv_receiver) + queue.req_to_metadata_buffer_idx_allocator.free.assert_called_once_with(3) + mock_release_kv_cache.assert_called_once_with( + req, queue.tree_cache, is_insert=False + ) + + def test_fake_receiver_initializes_deferred_release_state(self): + manager = MagicMock() + receiver = FakeKVReceiver(manager, "") + + self.assertIs(receiver.kv_mgr, manager) + self.assertFalse(receiver.abort_notified) + def test_retracted_decode_requests_keep_scheduler_non_idle(self): scheduler = Scheduler.__new__(Scheduler) scheduler.running_batch = MagicMock()