[PD] Handle abort requests in PP mode (#32797)
This commit is contained in:
@@ -898,7 +898,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
self._resolve_pending_reqs()
|
self._resolve_pending_reqs()
|
||||||
self._update_handshake_waiters(rids_to_check, pp_good_rids, pp_bad_rids)
|
self._update_handshake_waiters(rids_to_check, pp_good_rids, pp_bad_rids)
|
||||||
if is_pp_mode:
|
if is_pp_mode:
|
||||||
rids_to_check = pp_good_rids + pp_bad_rids
|
rids_to_check = set(pp_good_rids) | set(pp_bad_rids)
|
||||||
|
|
||||||
failed_reqs = []
|
failed_reqs = []
|
||||||
preallocated_reqs = []
|
preallocated_reqs = []
|
||||||
|
|||||||
@@ -4224,6 +4224,8 @@ class Scheduler(
|
|||||||
|
|
||||||
if hasattr(req.disagg_kv_sender, "abort"):
|
if hasattr(req.disagg_kv_sender, "abort"):
|
||||||
req.disagg_kv_sender.abort()
|
req.disagg_kv_sender.abort()
|
||||||
|
if self.ps.pp_size > 1:
|
||||||
|
prepare_abort(req, "Aborted by AbortReq.")
|
||||||
|
|
||||||
# Abort in-flight requests
|
# Abort in-flight requests
|
||||||
for req in self.disagg_prefill_inflight_queue:
|
for req in self.disagg_prefill_inflight_queue:
|
||||||
@@ -4238,6 +4240,8 @@ 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 self.ps.pp_size > 1:
|
||||||
|
prepare_abort(decode_req.req, "Aborted by AbortReq.")
|
||||||
|
|
||||||
# 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:
|
||||||
|
|||||||
@@ -24,7 +24,7 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
set_is_extend_in_batch,
|
set_is_extend_in_batch,
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.overlap_utils import RelayPayload
|
from sglang.srt.managers.overlap_utils import RelayPayload
|
||||||
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
from sglang.srt.managers.schedule_batch import FINISH_ABORT, Req, ScheduleBatch
|
||||||
from sglang.srt.managers.utils import (
|
from sglang.srt.managers.utils import (
|
||||||
GenerationBatchResult,
|
GenerationBatchResult,
|
||||||
get_logprob_dict_from_result,
|
get_logprob_dict_from_result,
|
||||||
@@ -847,6 +847,18 @@ class SchedulerPPMixin:
|
|||||||
bad_bootstrapped_rids = list(
|
bad_bootstrapped_rids = list(
|
||||||
set(prev_bad_bootstrapped_rids) | set(curr_bad_bootstrapped_rids)
|
set(prev_bad_bootstrapped_rids) | set(curr_bad_bootstrapped_rids)
|
||||||
)
|
)
|
||||||
|
# Route locally-aborted reqs through the bad-union consensus so every PP
|
||||||
|
# rank flushes them in the same consensus round, regardless of when the
|
||||||
|
# AbortReq reaches each rank and regardless of whether
|
||||||
|
# disagg_kv_sender.abort() drives the poll to Failed (it is optional).
|
||||||
|
aborted_rids = {
|
||||||
|
req.rid
|
||||||
|
for req in self.disagg_prefill_bootstrap_queue.queue
|
||||||
|
if isinstance(req.finished_reason, FINISH_ABORT)
|
||||||
|
}
|
||||||
|
good_bootstrapped_rids, bad_bootstrapped_rids = self._route_aborts_to_bad(
|
||||||
|
good_bootstrapped_rids, bad_bootstrapped_rids, aborted_rids
|
||||||
|
)
|
||||||
return [good_bootstrapped_rids, bad_bootstrapped_rids]
|
return [good_bootstrapped_rids, bad_bootstrapped_rids]
|
||||||
|
|
||||||
def _pp_pd_get_prefill_transferred_ids(self: Scheduler):
|
def _pp_pd_get_prefill_transferred_ids(self: Scheduler):
|
||||||
@@ -1369,8 +1381,31 @@ class SchedulerPPMixin:
|
|||||||
bad_prealloc_rids = list(
|
bad_prealloc_rids = list(
|
||||||
set(prev_bad_prealloc_rids) | set(curr_bad_prealloc_rids)
|
set(prev_bad_prealloc_rids) | set(curr_bad_prealloc_rids)
|
||||||
)
|
)
|
||||||
|
# Same abort routing as the prefill bootstrap consensus above.
|
||||||
|
aborted_rids = {
|
||||||
|
decode_req.req.rid
|
||||||
|
for decode_req in self.disagg_decode_prealloc_queue.queue
|
||||||
|
if isinstance(decode_req.req.finished_reason, FINISH_ABORT)
|
||||||
|
}
|
||||||
|
good_prealloc_rids, bad_prealloc_rids = self._route_aborts_to_bad(
|
||||||
|
good_prealloc_rids, bad_prealloc_rids, aborted_rids
|
||||||
|
)
|
||||||
return [good_prealloc_rids, bad_prealloc_rids]
|
return [good_prealloc_rids, bad_prealloc_rids]
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _route_aborts_to_bad(good_rids, bad_rids, aborted_rids):
|
||||||
|
"""Move aborted rids out of the good (intersection) set and into the
|
||||||
|
bad (union) set, so PP consensus fails them uniformly on every rank.
|
||||||
|
|
||||||
|
This also flushes aborted reqs that never reached good/bad consensus
|
||||||
|
(e.g. stuck in Bootstrapping with a sender that has no working abort()).
|
||||||
|
"""
|
||||||
|
if not aborted_rids:
|
||||||
|
return good_rids, bad_rids
|
||||||
|
good_rids = [rid for rid in good_rids if rid not in aborted_rids]
|
||||||
|
bad_rids = list(set(bad_rids) | set(aborted_rids))
|
||||||
|
return good_rids, bad_rids
|
||||||
|
|
||||||
def _pp_pd_get_decode_transferred_ids(self: Scheduler):
|
def _pp_pd_get_decode_transferred_ids(self: Scheduler):
|
||||||
# get the current stage transfer success
|
# get the current stage transfer success
|
||||||
if self.pp_group.is_first_rank:
|
if self.pp_group.is_first_rank:
|
||||||
|
|||||||
Reference in New Issue
Block a user