Fix scheduler crash on prefill-unreachable decode abort (#29834)
This commit is contained in:
@@ -725,6 +725,8 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
error_msg = f"Could not fetch prefill parallel info from {bootstrap_addr} after {count} attempts"
|
error_msg = f"Could not fetch prefill parallel info from {bootstrap_addr} after {count} attempts"
|
||||||
logger.error(error_msg)
|
logger.error(error_msg)
|
||||||
for decode_req in reqs:
|
for decode_req in reqs:
|
||||||
|
# kv_receiver may be None from a prior self.queue cleanup
|
||||||
|
if decode_req.kv_receiver is not None:
|
||||||
decode_req.kv_receiver.abort()
|
decode_req.kv_receiver.abort()
|
||||||
del self._ensure_retry_count[bootstrap_addr]
|
del self._ensure_retry_count[bootstrap_addr]
|
||||||
del self._ensure_last_attempt_time[bootstrap_addr]
|
del self._ensure_last_attempt_time[bootstrap_addr]
|
||||||
@@ -837,6 +839,14 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
failed_reqs.append(decode_req)
|
failed_reqs.append(decode_req)
|
||||||
indices_to_remove.add(i)
|
indices_to_remove.add(i)
|
||||||
|
|
||||||
|
# DecodeRequest is shared between self.queue and self.pending_reqs;
|
||||||
|
# drop failed reqs from both
|
||||||
|
if failed_reqs:
|
||||||
|
failed_ids = {id(r) for r in failed_reqs}
|
||||||
|
self.pending_reqs = [
|
||||||
|
r for r in self.pending_reqs if id(r) not in failed_ids
|
||||||
|
]
|
||||||
|
|
||||||
# HiSparse physical constraint: max requests by device buffer capacity.
|
# HiSparse physical constraint: max requests by device buffer capacity.
|
||||||
# Each admitted req needs padded_buffer_size from hisparse device pool.
|
# Each admitted req needs padded_buffer_size from hisparse device pool.
|
||||||
# waiting_queue reqs already have device buffers (allocated in admit_request_direct),
|
# waiting_queue reqs already have device buffers (allocated in admit_request_direct),
|
||||||
|
|||||||
@@ -66,6 +66,70 @@ class TestDecodeQueueCleanup(CustomTestCase):
|
|||||||
[req], req.return_logprob
|
[req], req.return_logprob
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_prealloc_abort_also_drops_from_pending_reqs(self):
|
||||||
|
# Same DecodeRequest lives in both queue and pending_reqs (add() slow
|
||||||
|
# path). Aborting must drop it from both, and compare by identity since
|
||||||
|
# DecodeRequest's dataclass __eq__ would compare the tensor receiver.
|
||||||
|
class BadEqReceiver(FakeReceiver):
|
||||||
|
def __eq__(self, other):
|
||||||
|
raise TypeError("use identity comparison, not value equality")
|
||||||
|
|
||||||
|
__hash__ = object.__hash__
|
||||||
|
|
||||||
|
receiver = BadEqReceiver()
|
||||||
|
req = SimpleNamespace(
|
||||||
|
rid="abort-shared",
|
||||||
|
finished_reason=FINISH_ABORT("aborted"),
|
||||||
|
return_logprob=False,
|
||||||
|
)
|
||||||
|
decode_req = SimpleNamespace(req=req, kv_receiver=receiver)
|
||||||
|
|
||||||
|
queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
|
||||||
|
queue.queue = [decode_req]
|
||||||
|
queue.pending_reqs = [decode_req] # same object, dual ownership
|
||||||
|
queue.retracted_queue = []
|
||||||
|
queue._resolve_pending_reqs = MagicMock()
|
||||||
|
queue._update_handshake_waiters = MagicMock()
|
||||||
|
queue._uses_swa_tail_prealloc = MagicMock(return_value=False)
|
||||||
|
queue._allocatable_token_budgets = MagicMock(return_value=0)
|
||||||
|
queue._hicache_pending_restore_tokens = MagicMock(return_value=0)
|
||||||
|
|
||||||
|
scheduler = MagicMock()
|
||||||
|
scheduler.running_batch.reqs = []
|
||||||
|
scheduler.enable_priority_scheduling = False
|
||||||
|
scheduler.enable_hisparse = False
|
||||||
|
scheduler.output_streamer = MagicMock()
|
||||||
|
queue.scheduler = scheduler
|
||||||
|
|
||||||
|
# Must not raise on the receiver __eq__ above.
|
||||||
|
preallocated, failed = queue.pop_preallocated()
|
||||||
|
|
||||||
|
self.assertEqual(preallocated, [])
|
||||||
|
self.assertEqual(failed, [decode_req])
|
||||||
|
self.assertEqual(queue.queue, [])
|
||||||
|
self.assertTrue(all(r is not decode_req for r in queue.pending_reqs))
|
||||||
|
self.assertIsNone(decode_req.kv_receiver)
|
||||||
|
|
||||||
|
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)
|
||||||
|
queue._max_ensure_retries = 1
|
||||||
|
queue._ensure_retry_interval = 0
|
||||||
|
queue._ensure_retry_count = {"127.0.0.1:11500": 0}
|
||||||
|
queue._ensure_last_attempt_time = {}
|
||||||
|
queue.kv_manager = MagicMock()
|
||||||
|
queue.kv_manager.try_ensure_parallel_info.return_value = False
|
||||||
|
|
||||||
|
cleared_req = SimpleNamespace(
|
||||||
|
req=SimpleNamespace(rid="cleared"), kv_receiver=None
|
||||||
|
)
|
||||||
|
addr_to_reqs = {"127.0.0.1:11500": [cleared_req]}
|
||||||
|
|
||||||
|
ready, remaining = queue._ensure_prefill_info(addr_to_reqs)
|
||||||
|
|
||||||
|
self.assertEqual(ready, {})
|
||||||
|
self.assertEqual(remaining, [])
|
||||||
|
|
||||||
@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")
|
||||||
|
|||||||
Reference in New Issue
Block a user