[PD] Gate deferred decode KV release on backend capability (#37454)

Co-authored-by: yhzhuang <yhzhuang@fb.com>
This commit is contained in:
Yonghao Zhuang
2026-09-03 15:42:21 -07:00
committed by GitHub
co-authored by yhzhuang
parent fd70325c10
commit 4dc9cda5f9
4 changed files with 33 additions and 1 deletions
@@ -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,
@@ -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
@@ -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
@@ -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()