From 2e6670739973533ac4bc2810613d4d4d3d564b1c Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Fri, 10 Jul 2026 08:53:18 +0800 Subject: [PATCH] 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> --- python/sglang/srt/managers/scheduler.py | 9 +- .../scheduler/test_scripted_pp_abort.py | 94 +++++++++++++++++++ 2 files changed, 99 insertions(+), 4 deletions(-) create mode 100644 test/manual/scheduler/test_scripted_pp_abort.py diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index a872dbbcb..c43c8c5eb 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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) ): diff --git a/test/manual/scheduler/test_scripted_pp_abort.py b/test/manual/scheduler/test_scripted_pp_abort.py new file mode 100644 index 000000000..06ab4505a --- /dev/null +++ b/test/manual/scheduler/test_scripted_pp_abort.py @@ -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()