Fix pipeline-parallel abort missing in-flight requests in non-current microbatch slots (#29405)

Co-authored-by: burling <3637497+burling@users.noreply.github.com>
Co-authored-by: zhaotyer <89376832+zhaotyer@users.noreply.github.com>
This commit is contained in:
fzyzcjy
2026-07-10 08:53:18 +08:00
committed by GitHub
co-authored by burling zhaotyer
parent 2c6cd1ef41
commit 2e66707399
2 changed files with 99 additions and 4 deletions
+5 -4
View File
@@ -3966,12 +3966,13 @@ class Scheduler(
self.disagg_decode_prealloc_queue.retracted_queue = remaining_retracted
# Delete requests in the running batch
if self.cur_batch is self.running_batch or self.cur_batch is None:
reqs = self.running_batch.reqs
if self.ps.pp_size == 1:
inflight_batches = [self.running_batch, self.cur_batch]
else:
reqs = self.running_batch.reqs + self.cur_batch.reqs
inflight_batches = [*self.running_mbs, *self.mbs]
for req in reqs:
inflight_reqs = {r for b in inflight_batches if b is not None for r in b.reqs}
for req in inflight_reqs:
if not req.finished() and (
recv_req.abort_all or req.rid.startswith(recv_req.rid)
):
@@ -0,0 +1,94 @@
import unittest
from sglang.test.scripted_runtime.context import ScriptedContext
from sglang.test.scripted_runtime.req_handle import ScriptedReqHandle
from sglang.test.scripted_runtime.test_case import ScriptedTestCase
from sglang.test.scripted_runtime_chunked_helpers import (
DEFAULT_CHUNK_SIZE,
base_engine_kwargs,
)
def _drain_until_released(t: ScriptedContext, *handles: ScriptedReqHandle):
for _ in range(12):
if all(
h.kv_pages == 0
and h.lock_refs == 0
and (h.req is None or h.req.req_pool_idx is None)
for h in handles
):
return
yield
def _advance_until_in_running_mbs(
t: ScriptedContext, *handles: ScriptedReqHandle, max_steps: int = 800
):
rids = {h.rid for h in handles}
present: set[str] = set()
for _ in range(max_steps):
present = {
req.rid
for mb in t.scheduler.running_mbs
for req in mb.reqs
if req.rid in rids
}
if present == rids:
return
yield
raise AssertionError(
f"reqs did not all reach a running_mbs decode slot within {max_steps} "
f"steps; present={present!r} wanted={rids!r}"
)
class TestAbortPPCrossSlot(ScriptedTestCase):
ENGINE_KWARGS = base_engine_kwargs(
chunked_prefill_size=DEFAULT_CHUNK_SIZE,
pp_size=4,
pp_max_micro_batch_size=1,
)
def test_abort_all_reaches_running_reqs_in_all_microbatch_slots(self):
"""abort_all must abort running reqs in every PP microbatch slot, not just the current one (needs >2 slots)."""
self.server.execute_script(
self._script_abort_all_reaches_running_reqs_in_all_microbatch_slots
)
@staticmethod
def _script_abort_all_reaches_running_reqs_in_all_microbatch_slots(
t: ScriptedContext,
):
reqs = [
t.start_req(
prompt_len=16,
max_new_tokens=512,
ignore_eos=True,
prompt_token=310 + i,
)
for i in range(4)
]
yield from _advance_until_in_running_mbs(t, *reqs)
slot_of = {}
for slot_id, mb in enumerate(t.scheduler.running_mbs):
for req in mb.reqs:
slot_of[req.rid] = slot_id
slots = {slot_of[r.rid] for r in reqs}
assert len(slots) == len(reqs), (
f"setup invalid: reqs must each occupy a distinct mb slot to exercise "
f"the cross-slot abort scan; slot_of={slot_of!r}"
)
t.abort_all()
yield from _drain_until_released(t, *reqs)
alive = {r.rid: r.kv_pages for r in reqs if r.kv_pages != 0}
assert not alive, (
f"abort_all left running reqs alive in non-current mb slots (only the "
f"current slot + stale cur_batch were scanned): still_holding_kv={alive!r}"
)
if __name__ == "__main__":
unittest.main()