fix(disagg): poll receivers during decode preallocation (#37483)
This commit is contained in:
@@ -869,17 +869,8 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
if not self.queue:
|
if not self.queue:
|
||||||
return
|
return
|
||||||
|
|
||||||
# Still poll if any receiver was aborted, otherwise it stays stuck.
|
# Receiver polling observes asynchronous failures while KV allocation
|
||||||
if (
|
# is blocked.
|
||||||
self.pp_size <= 1
|
|
||||||
and all(decode_req.waiting_for_input for decode_req in self.queue)
|
|
||||||
and not any(
|
|
||||||
decode_req.kv_receiver.conclude_state == KVPoll.Failed
|
|
||||||
for decode_req in self.queue
|
|
||||||
)
|
|
||||||
):
|
|
||||||
return
|
|
||||||
|
|
||||||
if self.pp_size > 1:
|
if self.pp_size > 1:
|
||||||
polls = poll_and_all_reduce_pp(
|
polls = poll_and_all_reduce_pp(
|
||||||
(decode_req.req.rid for decode_req in self.queue),
|
(decode_req.req.rid for decode_req in self.queue),
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
|||||||
class FakeReceiver:
|
class FakeReceiver:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.clear_called = False
|
self.clear_called = False
|
||||||
|
self.conclude_state = None
|
||||||
|
|
||||||
def clear(self):
|
def clear(self):
|
||||||
self.clear_called = True
|
self.clear_called = True
|
||||||
@@ -94,18 +95,22 @@ class TestDecodeQueueCleanup(CustomTestCase):
|
|||||||
receiver = FakeReceiver()
|
receiver = FakeReceiver()
|
||||||
req = SimpleNamespace(
|
req = SimpleNamespace(
|
||||||
rid="abort-prealloc",
|
rid="abort-prealloc",
|
||||||
finished_reason=FINISH_ABORT("aborted"),
|
bootstrap_room=42,
|
||||||
|
finished_reason=None,
|
||||||
return_logprob=False,
|
return_logprob=False,
|
||||||
)
|
)
|
||||||
decode_req = SimpleNamespace(req=req, kv_receiver=receiver)
|
decode_req = SimpleNamespace(
|
||||||
|
req=req, kv_receiver=receiver, waiting_for_input=True
|
||||||
|
)
|
||||||
|
|
||||||
queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
|
queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
|
||||||
queue.pp_size = 1
|
queue.pp_size = 1
|
||||||
|
queue.tp_rank = 0
|
||||||
|
queue.gloo_group = object()
|
||||||
queue.queue = [decode_req]
|
queue.queue = [decode_req]
|
||||||
queue.pending_reqs = []
|
queue.pending_reqs = []
|
||||||
queue.retracted_queue = []
|
queue.retracted_queue = []
|
||||||
queue._resolve_pending_reqs = MagicMock()
|
queue._resolve_pending_reqs = MagicMock()
|
||||||
queue._update_handshake_waiters = MagicMock()
|
|
||||||
queue._uses_swa_tail_prealloc = MagicMock(return_value=False)
|
queue._uses_swa_tail_prealloc = MagicMock(return_value=False)
|
||||||
queue._allocatable_token_budgets = MagicMock(return_value=0)
|
queue._allocatable_token_budgets = MagicMock(return_value=0)
|
||||||
queue._hicache_pending_restore_tokens = MagicMock(return_value=0)
|
queue._hicache_pending_restore_tokens = MagicMock(return_value=0)
|
||||||
@@ -114,16 +119,23 @@ class TestDecodeQueueCleanup(CustomTestCase):
|
|||||||
scheduler.running_batch.reqs = []
|
scheduler.running_batch.reqs = []
|
||||||
scheduler.enable_priority_scheduling = False
|
scheduler.enable_priority_scheduling = False
|
||||||
scheduler.enable_hisparse = False
|
scheduler.enable_hisparse = False
|
||||||
|
scheduler.metrics_reporter.enable_metrics = False
|
||||||
scheduler.output_streamer = MagicMock()
|
scheduler.output_streamer = MagicMock()
|
||||||
queue.scheduler = scheduler
|
queue.scheduler = scheduler
|
||||||
|
|
||||||
preallocated, failed = queue.pop_preallocated()
|
with patch(
|
||||||
|
"sglang.srt.disaggregation.decode.poll_and_all_reduce",
|
||||||
|
return_value=[KVPoll.Failed],
|
||||||
|
) as poll:
|
||||||
|
preallocated, failed = queue.pop_preallocated()
|
||||||
|
|
||||||
|
poll.assert_called_once_with([receiver], queue.gloo_group)
|
||||||
self.assertEqual(preallocated, [])
|
self.assertEqual(preallocated, [])
|
||||||
self.assertEqual(failed, [decode_req])
|
self.assertEqual(failed, [decode_req])
|
||||||
self.assertEqual(queue.queue, [])
|
self.assertEqual(queue.queue, [])
|
||||||
self.assertTrue(receiver.clear_called)
|
self.assertTrue(receiver.clear_called)
|
||||||
self.assertIsNone(decode_req.kv_receiver)
|
self.assertIsNone(decode_req.kv_receiver)
|
||||||
|
self.assertIsInstance(req.finished_reason, FINISH_ABORT)
|
||||||
scheduler.output_streamer.stream_output.assert_called_once_with(
|
scheduler.output_streamer.stream_output.assert_called_once_with(
|
||||||
[req], req.return_logprob
|
[req], req.return_logprob
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user