fix(disagg): poll receivers during decode preallocation (#37483)

This commit is contained in:
cctry
2026-09-02 14:26:40 -07:00
committed by GitHub
parent ad6e830858
commit 3a855b050a
2 changed files with 18 additions and 15 deletions
+2 -11
View File
@@ -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
) )