[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):
|
class BaseKVManager(ABC):
|
||||||
"""Base class for managing transfer states"""
|
"""Base class for managing transfer states"""
|
||||||
|
|
||||||
|
enable_deferred_decode_kv_release: bool = False
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -2326,6 +2326,7 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
|
|||||||
self.scheduler.hisparse_coordinator.request_finished(decode_req.req)
|
self.scheduler.hisparse_coordinator.request_finished(decode_req.req)
|
||||||
if (
|
if (
|
||||||
self.enable_deferred_kv_release
|
self.enable_deferred_kv_release
|
||||||
|
and decode_req.kv_receiver.kv_mgr.enable_deferred_decode_kv_release
|
||||||
and decode_req.kv_receiver.abort_notified
|
and decode_req.kv_receiver.abort_notified
|
||||||
):
|
):
|
||||||
# Decode-initiated abort: a prefill write may still target
|
# Decode-initiated abort: a prefill write may still target
|
||||||
|
|||||||
@@ -99,6 +99,8 @@ class FakeKVReceiver(BaseKVReceiver):
|
|||||||
bootstrap_addr: str,
|
bootstrap_addr: str,
|
||||||
bootstrap_room: Optional[int] = None,
|
bootstrap_room: Optional[int] = None,
|
||||||
):
|
):
|
||||||
|
self.kv_mgr = mgr
|
||||||
|
self.abort_notified: bool = False
|
||||||
self.bootstrap_done = False
|
self.bootstrap_done = False
|
||||||
self.has_sent_metadata = False
|
self.has_sent_metadata = False
|
||||||
self.require_staging: bool = False
|
self.require_staging: bool = False
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ from sglang.srt.disaggregation.decode import (
|
|||||||
DecodeTransferQueue,
|
DecodeTransferQueue,
|
||||||
HiCacheRestoreResult,
|
HiCacheRestoreResult,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.disaggregation.fake.conn import FakeKVManager, FakeKVReceiver
|
||||||
from sglang.srt.disaggregation.utils import DisaggregationMode
|
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||||
from sglang.srt.managers.schedule_batch import FINISH_ABORT
|
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.release_kv_cache")
|
||||||
@patch("sglang.srt.disaggregation.decode.prepare_abort")
|
@patch("sglang.srt.disaggregation.decode.prepare_abort")
|
||||||
@patch("sglang.srt.disaggregation.decode.poll_and_all_reduce")
|
@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
|
self, mock_poll, mock_prepare_abort, mock_release_kv_cache
|
||||||
):
|
):
|
||||||
receiver = FakeReceiver()
|
receiver = FakeReceiver()
|
||||||
@@ -371,6 +372,32 @@ class TestDecodeQueueCleanup(CustomTestCase):
|
|||||||
req, queue.tree_cache, is_insert=False
|
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):
|
def test_retracted_decode_requests_keep_scheduler_non_idle(self):
|
||||||
scheduler = Scheduler.__new__(Scheduler)
|
scheduler = Scheduler.__new__(Scheduler)
|
||||||
scheduler.running_batch = MagicMock()
|
scheduler.running_batch = MagicMock()
|
||||||
|
|||||||
Reference in New Issue
Block a user