[P/D disagg] Decode-side radix cache for SWA hybrid models (unified radix tree) (#27770)
Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
co-authored by
Shangming Cai
parent
aa3f766799
commit
978244d671
@@ -72,6 +72,9 @@ class TestDecodeQueueCleanup(CustomTestCase):
|
||||
queue._swa_aware_allocatable_token_budgets = MagicMock(
|
||||
return_value=(physical_available, physical_available)
|
||||
)
|
||||
queue._swa_tail_allocatable_token_budget = MagicMock(
|
||||
side_effect=lambda **_: physical_available
|
||||
)
|
||||
|
||||
def pre_alloc(_req):
|
||||
nonlocal physical_available
|
||||
@@ -170,6 +173,72 @@ class TestDecodeQueueCleanup(CustomTestCase):
|
||||
self.assertTrue(all(r is not decode_req for r in queue.pending_reqs))
|
||||
self.assertIsNone(decode_req.kv_receiver)
|
||||
|
||||
def test_swa_reclaim_failure_rejects_only_request(self):
|
||||
receiver = FakeReceiver()
|
||||
req = SimpleNamespace(
|
||||
rid="swa-reclaim-failed",
|
||||
origin_input_ids=[1, 2, 3],
|
||||
output_ids=[],
|
||||
finished_reason=None,
|
||||
return_logprob=False,
|
||||
sampling_params=SimpleNamespace(max_new_tokens=1),
|
||||
)
|
||||
decode_req = SimpleNamespace(
|
||||
req=req,
|
||||
kv_receiver=receiver,
|
||||
waiting_for_input=True,
|
||||
is_rebootstrap=False,
|
||||
)
|
||||
|
||||
queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
|
||||
queue.pp_size = 1
|
||||
queue.queue = [decode_req]
|
||||
queue.pending_reqs = [decode_req]
|
||||
queue.retracted_queue = []
|
||||
queue.num_reserved_decode_tokens = 0
|
||||
queue._resolve_pending_reqs = MagicMock()
|
||||
queue._update_handshake_waiters = MagicMock()
|
||||
queue._uses_swa_tail_prealloc = MagicMock(return_value=True)
|
||||
queue._swa_aware_allocatable_token_budgets = MagicMock(
|
||||
return_value=(1024, 1024)
|
||||
)
|
||||
queue._prealloc_required_tokens = MagicMock(return_value=(3, 3))
|
||||
queue._prealloc_kv_lens = MagicMock(return_value=(3, 3))
|
||||
queue._reclaim_swa_tail_capacity = MagicMock(
|
||||
return_value=(
|
||||
"SWA eviction insufficient: needed=64, available=0, "
|
||||
"req=swa-reclaim-failed"
|
||||
)
|
||||
)
|
||||
queue._hicache_pending_restore_tokens = MagicMock(return_value=0)
|
||||
queue._pre_alloc = MagicMock()
|
||||
queue.req_to_token_pool = MagicMock()
|
||||
queue.req_to_token_pool.available_size.return_value = 1
|
||||
queue.req_to_metadata_buffer_idx_allocator = MagicMock()
|
||||
queue.req_to_metadata_buffer_idx_allocator.available_size.return_value = 1
|
||||
|
||||
scheduler = MagicMock()
|
||||
scheduler.running_batch.reqs = []
|
||||
scheduler.enable_priority_scheduling = False
|
||||
scheduler.enable_hisparse = False
|
||||
scheduler.server_args.disaggregation_decode_enable_radix_cache = False
|
||||
scheduler.output_streamer = MagicMock()
|
||||
queue.scheduler = scheduler
|
||||
|
||||
preallocated, failed = queue.pop_preallocated()
|
||||
|
||||
self.assertEqual(preallocated, [])
|
||||
self.assertEqual(failed, [decode_req])
|
||||
self.assertEqual(queue.queue, [])
|
||||
self.assertEqual(queue.pending_reqs, [])
|
||||
self.assertTrue(receiver.clear_called)
|
||||
self.assertIsNone(decode_req.kv_receiver)
|
||||
self.assertIsInstance(req.finished_reason, FINISH_ABORT)
|
||||
queue._pre_alloc.assert_not_called()
|
||||
scheduler.output_streamer.stream_output.assert_called_once_with(
|
||||
[req], req.return_logprob
|
||||
)
|
||||
|
||||
def test_ensure_prefill_info_tolerates_cleared_receiver(self):
|
||||
# A req whose kv_receiver was already cleared must not crash on .abort().
|
||||
queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
|
||||
|
||||
Reference in New Issue
Block a user