[PD] Gate deferred decode KV release on backend capability (#37454)
Co-authored-by: yhzhuang <yhzhuang@fb.com>
This commit is contained in:
co-authored by
yhzhuang
parent
fd70325c10
commit
4dc9cda5f9
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user