fix(disagg): unstuck decode aborts under prealloc pressure (#25561)
This commit is contained in:
@@ -609,7 +609,11 @@ class DecodePreallocQueue:
|
|||||||
if not self.queue:
|
if not self.queue:
|
||||||
return
|
return
|
||||||
|
|
||||||
if all(decode_req.waiting_for_input for decode_req in self.queue):
|
# Still poll if any receiver was aborted, otherwise it stays stuck.
|
||||||
|
if all(decode_req.waiting_for_input for decode_req in self.queue) and not any(
|
||||||
|
getattr(decode_req.kv_receiver, "conclude_state", None) == KVPoll.Failed
|
||||||
|
for decode_req in self.queue
|
||||||
|
):
|
||||||
return
|
return
|
||||||
|
|
||||||
polls = poll_and_all_reduce(
|
polls = poll_and_all_reduce(
|
||||||
|
|||||||
@@ -3499,12 +3499,16 @@ class Scheduler(
|
|||||||
if recv_req.abort_all or decode_req.req.rid.startswith(recv_req.rid):
|
if recv_req.abort_all or decode_req.req.rid.startswith(recv_req.rid):
|
||||||
logger.debug(f"Abort prealloc queue request. {decode_req.req.rid=}")
|
logger.debug(f"Abort prealloc queue request. {decode_req.req.rid=}")
|
||||||
decode_req.kv_receiver.abort()
|
decode_req.kv_receiver.abort()
|
||||||
|
if not isinstance(decode_req.req.finished_reason, FINISH_ABORT):
|
||||||
|
decode_req.req.finished_reason = FINISH_ABORT()
|
||||||
|
|
||||||
# Abort requests waiting for kvcache to release tree cache
|
# Abort requests waiting for kvcache to release tree cache
|
||||||
for decode_req in self.disagg_decode_transfer_queue.queue:
|
for decode_req in self.disagg_decode_transfer_queue.queue:
|
||||||
if recv_req.abort_all or decode_req.req.rid.startswith(recv_req.rid):
|
if recv_req.abort_all or decode_req.req.rid.startswith(recv_req.rid):
|
||||||
logger.debug(f"Abort transfer queue request. {decode_req.req.rid=}")
|
logger.debug(f"Abort transfer queue request. {decode_req.req.rid=}")
|
||||||
decode_req.kv_receiver.abort()
|
decode_req.kv_receiver.abort()
|
||||||
|
if not isinstance(decode_req.req.finished_reason, FINISH_ABORT):
|
||||||
|
decode_req.req.finished_reason = FINISH_ABORT()
|
||||||
|
|
||||||
# Abort requests already retracted to CPU cache
|
# Abort requests already retracted to CPU cache
|
||||||
if self.disagg_decode_prealloc_queue.retracted_queue:
|
if self.disagg_decode_prealloc_queue.retracted_queue:
|
||||||
|
|||||||
@@ -1491,7 +1491,15 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
def abort_request(self, rid: str = "", abort_all: bool = False):
|
def abort_request(self, rid: str = "", abort_all: bool = False):
|
||||||
if not abort_all and rid not in self.rid_to_state:
|
# Empty rid would startswith-match every request on the scheduler.
|
||||||
|
if not abort_all and not rid:
|
||||||
|
logger.warning("Ignore abort_request with empty rid and abort_all=False")
|
||||||
|
return
|
||||||
|
if (
|
||||||
|
not abort_all
|
||||||
|
and self.server_args.tokenizer_worker_num == 1
|
||||||
|
and rid not in self.rid_to_state
|
||||||
|
):
|
||||||
return
|
return
|
||||||
req = AbortReq(rid=rid, abort_all=abort_all)
|
req = AbortReq(rid=rid, abort_all=abort_all)
|
||||||
self.send_to_scheduler.send_pyobj(req)
|
self.send_to_scheduler.send_pyobj(req)
|
||||||
|
|||||||
Reference in New Issue
Block a user