diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index c339b88f0..20f370fe3 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -4083,46 +4083,36 @@ class Scheduler( tmp_batch, tmp_result = self.result_queue.popleft() self.process_batch_result(tmp_batch, tmp_result) - if self.last_batch and self.last_batch.forward_mode.is_extend(): - chunked_req_to_exclude = set() - self.last_batch.filter_batch( - chunked_req_to_exclude=list(chunked_req_to_exclude) - ) + retract_reqs = [r for r in self.running_batch.reqs if not r.finished()] + if ( + self.last_batch is not None + and self.last_batch.forward_mode.is_extend() # Skip merge for disagg prefill: completed prefill requests are # already in disagg_prefill_inflight_queue. Merging them into # running_batch leaks them, since the prefill event loop never # calls update_running_batch to clean them up. - if ( - not self.last_batch.is_empty() - and self.disaggregation_mode != DisaggregationMode.PREFILL - ): - if self.running_batch.is_empty(): - self.running_batch = self.last_batch - else: - self.running_batch.merge_batch(self.last_batch) + and self.disaggregation_mode != DisaggregationMode.PREFILL + ): + retract_reqs += [r for r in self.last_batch.reqs if not r.finished()] + + if ( + self.chunked_req is not None + and not self.chunked_req.finished() + and self.chunked_req not in retract_reqs + and self.disaggregation_mode != DisaggregationMode.PREFILL + ): + retract_reqs.append(self.chunked_req) self.last_batch = None self.cur_batch_for_debug = None - if not self.running_batch.is_empty(): - self.running_batch.filter_batch() - - retracted_reqs = list(self.running_batch.reqs) - if ( - self.chunked_req is not None - and not self.chunked_req.finished() - and self.chunked_req not in retracted_reqs - and self.disaggregation_mode != DisaggregationMode.PREFILL - ): - retracted_reqs.append(self.chunked_req) - - if retracted_reqs: + if retract_reqs: # Decode-side retract always rebootstraps (recomputes the KV from # the prefill), so skip the device->host KV offload that release_req # would otherwise do; the offloaded copy would be immediately # discarded. Non-decode modes ignore offload_kv (they never offload). retract_all( - reqs=retracted_reqs, + reqs=retract_reqs, server_args=self.server_args, req_to_token_pool=self.req_to_token_pool, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, @@ -4130,17 +4120,16 @@ class Scheduler( hisparse_coordinator=self.hisparse_coordinator, offload_kv=False, ) - self.running_batch.reqs = [] - for req in retracted_reqs: - if self.disaggregation_mode == DisaggregationMode.DECODE: - if req.output_ids: - req.pd_rebootstrap_forced_output_id = req.output_ids.pop() - req.pd_rebootstrap_in_progress = True - req.time_stats.set_retract_time() - self.disagg_decode_prealloc_queue.hold_rebootstrap(req) - else: - self._add_request_to_queue(req) - + self.running_batch.reqs = [] + for req in retract_reqs: + if self.disaggregation_mode == DisaggregationMode.DECODE: + if req.output_ids: + req.pd_rebootstrap_forced_output_id = req.output_ids.pop() + req.pd_rebootstrap_in_progress = True + req.time_stats.set_retract_time() + self.disagg_decode_prealloc_queue.hold_rebootstrap(req) + else: + self._add_request_to_queue(req) self.running_batch.batch_is_full = False # In disagg-PREFILL, keep a live mid-chunk chunked_req rather than retract it: # freeing its KV under a live disagg KV-sender crashes pop_bootstrapped or diff --git a/test/registered/unit/managers/test_scheduler_pause_generation.py b/test/registered/unit/managers/test_scheduler_pause_generation.py index 63883e07b..e9a24cef5 100644 --- a/test/registered/unit/managers/test_scheduler_pause_generation.py +++ b/test/registered/unit/managers/test_scheduler_pause_generation.py @@ -1,8 +1,11 @@ import unittest from collections import deque from types import SimpleNamespace +from typing import List, Optional from unittest.mock import MagicMock, patch +import torch + from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import maybe_stub_sgl_kernel @@ -13,8 +16,11 @@ from sglang.srt.managers.io_struct import ( ContinueGenerationReqInput, PauseGenerationReqInput, ) +from sglang.srt.managers.schedule_batch import Req, ScheduleBatch from sglang.srt.managers.scheduler import Scheduler from sglang.srt.managers.scheduler_components.pool_stats_observer import PoolStats +from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.srt.sampling.sampling_params import SamplingParams register_cpu_ci(est_time=15, suite="base-a-test-cpu") register_cpu_ci(est_time=9, suite="base-c-test-cpu") @@ -50,12 +56,64 @@ class TestSchedulerPauseGeneration(unittest.TestCase): full_evictable_size=0, ) ) + scheduler.disaggregation_mode = DisaggregationMode.NULL + scheduler.hisparse_coordinator = None + scheduler.server_args = MagicMock() + scheduler.waiting_queue = [] # pause_generation zeros gen_throughput and flushes KV events. scheduler.metrics_reporter = MagicMock() scheduler.metrics_reporter.current_scheduler_metrics_enabled = False scheduler.kv_events_publisher = MagicMock() return scheduler + def _make_req(self, rid: str, finished: bool = False) -> Req: + req = Req( + rid=rid, + origin_input_text="", + origin_input_ids=[1, 2, 3], + sampling_params=SamplingParams(), + ) + if finished: + req.finished_reason = MagicMock() + return req + + def _make_batch( + self, + scheduler: Scheduler, + reqs: List[Req], + forward_mode: Optional[ForwardMode] = None, + with_tensors: bool = False, + ) -> ScheduleBatch: + batch = ScheduleBatch(reqs=reqs) + batch.device = "cpu" + batch.forward_mode = forward_mode + batch.req_to_token_pool = scheduler.req_to_token_pool + batch.token_to_kv_pool_allocator = scheduler.token_to_kv_pool_allocator + batch.tree_cache = scheduler.tree_cache + batch.hisparse_coordinator = None + batch.model_config = MagicMock(is_encoder_decoder=False) + batch.sampling_info = MagicMock() + batch.spec_info = None + batch.multimodal_inputs = None + if with_tensors: + batch_size = len(reqs) + batch.req_pool_indices = torch.arange(batch_size, dtype=torch.int64) + batch.req_pool_indices_cpu = torch.arange(batch_size, dtype=torch.int64) + batch.seq_lens = torch.full((batch_size,), 4, dtype=torch.int64) + batch.orig_seq_lens = torch.full((batch_size,), 4, dtype=torch.int32) + batch.seq_lens_cpu = torch.full((batch_size,), 4, dtype=torch.int64) + batch.input_ids = None + return batch + + def _spy_requeue(self, scheduler: Scheduler) -> List[dict]: + requeue_log: List[dict] = [] + + def record(req): + requeue_log.append({"req": req, "is_retracted": req.is_retracted}) + + scheduler._add_request_to_queue = MagicMock(side_effect=record) + return requeue_log + def test_inplace_only_sets_flag(self): """in_place pause should only set _engine_paused and return.""" scheduler = self._new_scheduler() @@ -126,82 +184,222 @@ class TestSchedulerPauseGeneration(unittest.TestCase): self.assertIsNone(scheduler.last_batch) self.assertIsNone(scheduler.cur_batch_for_debug) - def test_retract_clears_running_batch(self): - """retract mode should retract all requests from running_batch.""" + def test_retract_requeues_running_then_last_fold_in(self): + """retract requeues running reqs first, then last extend reqs, all released.""" scheduler = self._new_scheduler() - scheduler.last_batch = None - scheduler.running_batch.reqs = [MagicMock(), MagicMock()] - scheduler.running_batch.__len__ = lambda self: len(self.reqs) - scheduler.running_batch.is_empty.return_value = False - scheduler.waiting_queue = [] - scheduler._add_request_to_queue = MagicMock() - - scheduler.running_batch.filter_batch = MagicMock() - scheduler.server_args = MagicMock() - reqs_before = scheduler.running_batch.reqs - - with patch("sglang.srt.managers.scheduler.retract_all") as mock_retract_all: - scheduler.pause_generation(PauseGenerationReqInput(mode="retract")) - - self.assertTrue(scheduler._engine_paused) - mock_retract_all.assert_called_once() - self.assertIs(mock_retract_all.call_args.kwargs["reqs"], reqs_before) - self.assertEqual(scheduler.running_batch.reqs, []) - self.assertEqual(scheduler._add_request_to_queue.call_count, 2) - self.assertEqual( - [call.args[0] for call in scheduler._add_request_to_queue.call_args_list], - reqs_before, + run_req_a = self._make_req("run-a") + run_req_b = self._make_req("run-b") + last_req = self._make_req("last") + scheduler.running_batch = self._make_batch( + scheduler, reqs=[run_req_a, run_req_b], with_tensors=True ) + scheduler.running_batch.batch_is_full = True + scheduler.last_batch = self._make_batch( + scheduler, + reqs=[last_req], + forward_mode=ForwardMode.EXTEND, + with_tensors=True, + ) + scheduler.chunked_req = MagicMock() + requeue_log = self._spy_requeue(scheduler) + + scheduler.pause_generation(PauseGenerationReqInput(mode="retract")) + + self.assertEqual( + [entry["req"] for entry in requeue_log], [run_req_a, run_req_b, last_req] + ) + self.assertTrue(all(entry["is_retracted"] for entry in requeue_log)) + self.assertEqual( + [req.retraction_count for req in (run_req_a, run_req_b, last_req)], + [1, 1, 1], + ) + self.assertEqual(scheduler.running_batch.reqs, []) + self.assertFalse(scheduler.running_batch.batch_is_full) + self.assertIsNone(scheduler.chunked_req) + self.assertIsNone(scheduler.last_batch) + + def test_retract_with_empty_running_uses_last_batch_reqs(self): + """retract with empty running batch releases and requeues the last extend reqs.""" + scheduler = self._new_scheduler() + last_req = self._make_req("last") + scheduler.running_batch = ScheduleBatch(reqs=[], batch_is_full=True) + scheduler.last_batch = self._make_batch( + scheduler, + reqs=[last_req], + forward_mode=ForwardMode.EXTEND, + with_tensors=True, + ) + scheduler.chunked_req = MagicMock() + requeue_log = self._spy_requeue(scheduler) + + scheduler.pause_generation(PauseGenerationReqInput(mode="retract")) + + self.assertEqual([entry["req"] for entry in requeue_log], [last_req]) + self.assertTrue(requeue_log[0]["is_retracted"]) + self.assertEqual(last_req.retraction_count, 1) + self.assertEqual(scheduler.running_batch.reqs, []) + self.assertFalse(scheduler.running_batch.batch_is_full) self.assertIsNone(scheduler.chunked_req) def test_retract_fold_in_releases_via_scheduler_hisparse_coordinator(self): """retract of a folded-in last extend batch must release through the scheduler-owned hisparse coordinator.""" scheduler = self._new_scheduler() - scheduler.disaggregation_mode = DisaggregationMode.NULL - scheduler.waiting_queue = [] - scheduler._add_request_to_queue = MagicMock() - scheduler.server_args = MagicMock() - - req = MagicMock() - req.finished.return_value = False - req.req_pool_idx = None - last_batch = MagicMock() - last_batch.forward_mode.is_extend.return_value = True - last_batch.is_empty.return_value = False - last_batch.reqs = [req] - scheduler.last_batch = last_batch + scheduler.hisparse_coordinator = MagicMock() + last_req = self._make_req("last") + scheduler.running_batch = ScheduleBatch(reqs=[], batch_is_full=True) + scheduler.last_batch = self._make_batch( + scheduler, + reqs=[last_req], + forward_mode=ForwardMode.EXTEND, + with_tensors=True, + ) + requeue_log = self._spy_requeue(scheduler) scheduler.pause_generation(PauseGenerationReqInput(mode="retract")) - scheduler.hisparse_coordinator.retract_req.assert_called_once_with(req) - self.assertEqual( - [call.args[0] for call in scheduler._add_request_to_queue.call_args_list], - [req], + scheduler.hisparse_coordinator.retract_req.assert_called_once_with(last_req) + self.assertEqual([entry["req"] for entry in requeue_log], [last_req]) + + def test_retract_disagg_prefill_excludes_last_batch(self): + """retract under disagg prefill must not release or requeue last extend reqs.""" + scheduler = self._new_scheduler() + scheduler.disaggregation_mode = DisaggregationMode.PREFILL + run_req = self._make_req("run") + last_req = self._make_req("last") + scheduler.running_batch = self._make_batch( + scheduler, reqs=[run_req], with_tensors=True ) + scheduler.last_batch = self._make_batch( + scheduler, + reqs=[last_req], + forward_mode=ForwardMode.EXTEND, + with_tensors=True, + ) + requeue_log = self._spy_requeue(scheduler) + + scheduler.pause_generation(PauseGenerationReqInput(mode="retract")) + + self.assertEqual([entry["req"] for entry in requeue_log], [run_req]) + self.assertEqual(run_req.retraction_count, 1) + self.assertEqual(last_req.retraction_count, 0) + self.assertFalse(last_req.is_retracted) + + def test_retract_decode_last_batch_only_retracts_running(self): + """retract with a decode last batch only releases and requeues running reqs.""" + scheduler = self._new_scheduler() + run_req = self._make_req("run") + running = self._make_batch( + scheduler, + reqs=[run_req], + forward_mode=ForwardMode.DECODE, + with_tensors=True, + ) + scheduler.running_batch = running + scheduler.last_batch = running + requeue_log = self._spy_requeue(scheduler) + + scheduler.pause_generation(PauseGenerationReqInput(mode="retract")) + + self.assertEqual([entry["req"] for entry in requeue_log], [run_req]) + self.assertEqual(run_req.retraction_count, 1) + self.assertEqual(scheduler.running_batch.reqs, []) + + def test_retract_partial_finished_running_batch(self): + """retract with mixed finished/unfinished reqs only releases the unfinished ones.""" + scheduler = self._new_scheduler() + req_unfinished_a = self._make_req("unfinished-a") + req_finished = self._make_req("finished", finished=True) + req_unfinished_b = self._make_req("unfinished-b") + scheduler.running_batch = self._make_batch( + scheduler, + reqs=[req_unfinished_a, req_finished, req_unfinished_b], + with_tensors=True, + ) + requeue_log = self._spy_requeue(scheduler) + + scheduler.pause_generation(PauseGenerationReqInput(mode="retract")) + + self.assertEqual( + [entry["req"] for entry in requeue_log], + [req_unfinished_a, req_unfinished_b], + ) + self.assertEqual(req_unfinished_a.retraction_count, 1) + self.assertEqual(req_unfinished_b.retraction_count, 1) + self.assertEqual(req_finished.retraction_count, 0) + self.assertFalse(req_finished.is_retracted) + self.assertEqual(scheduler.running_batch.reqs, []) + + def test_retract_empty_post_fold_clears_chunked_req_and_batch_is_full(self): + """retract with nothing to retract still clears chunked_req and batch_is_full.""" + scheduler = self._new_scheduler() + scheduler.running_batch = ScheduleBatch(reqs=[], batch_is_full=True) + scheduler.chunked_req = MagicMock() + requeue_log = self._spy_requeue(scheduler) + + scheduler.pause_generation(PauseGenerationReqInput(mode="retract")) + + self.assertEqual(requeue_log, []) + self.assertIsNone(scheduler.chunked_req) + self.assertFalse(scheduler.running_batch.batch_is_full) + + def test_retract_all_finished_clears_fields_without_requeue(self): + """retract with only finished reqs clears fields but releases nothing.""" + scheduler = self._new_scheduler() + req_finished_a = self._make_req("finished-a", finished=True) + req_finished_b = self._make_req("finished-b", finished=True) + scheduler.running_batch = self._make_batch( + scheduler, reqs=[req_finished_a, req_finished_b] + ) + scheduler.running_batch.batch_is_full = True + scheduler.chunked_req = MagicMock() + requeue_log = self._spy_requeue(scheduler) + + scheduler.pause_generation(PauseGenerationReqInput(mode="retract")) + + self.assertEqual(requeue_log, []) + self.assertEqual(req_finished_a.retraction_count, 0) + self.assertEqual(req_finished_b.retraction_count, 0) + self.assertEqual(scheduler.running_batch.reqs, []) + self.assertFalse(scheduler.running_batch.batch_is_full) + self.assertIsNone(scheduler.chunked_req) + + def test_retract_drain_happens_once_before_release(self): + """retract with overlap drains the result_queue once before releasing reqs.""" + scheduler = self._new_scheduler() + scheduler.enable_overlap = True + last_req = self._make_req("last") + scheduler.running_batch = ScheduleBatch(reqs=[]) + scheduler.last_batch = self._make_batch( + scheduler, + reqs=[last_req], + forward_mode=ForwardMode.EXTEND, + with_tensors=True, + ) + scheduler.result_queue = deque([(MagicMock(), MagicMock())]) + event_log: List[str] = [] + scheduler.process_batch_result = MagicMock( + side_effect=lambda *args, **kwargs: event_log.append("drain") + ) + scheduler._add_request_to_queue = MagicMock( + side_effect=lambda req: event_log.append( + "requeue-released" if req.is_retracted else "requeue-unreleased" + ) + ) + + scheduler.pause_generation(PauseGenerationReqInput(mode="retract")) + + self.assertEqual(event_log, ["drain", "requeue-released"]) + self.assertEqual(len(scheduler.result_queue), 0) def test_retract_empty_running_batch_requeues_nothing(self): """retract with empty running_batch must not release or requeue any request.""" scheduler = self._new_scheduler() - scheduler.waiting_queue = [] - original_reqs = scheduler.running_batch.reqs scheduler.pause_generation(PauseGenerationReqInput(mode="retract")) self.assertTrue(scheduler._engine_paused) self.assertEqual(len(scheduler.waiting_queue), 0) - self.assertIs(scheduler.running_batch.reqs, original_reqs) - - def test_retract_empty_clears_chunked_req_and_batch_is_full(self): - """retract with everything empty must still clear chunked_req and batch_is_full.""" - scheduler = self._new_scheduler() - scheduler.waiting_queue = [] - scheduler.chunked_req = MagicMock() - scheduler.running_batch.batch_is_full = True - - scheduler.pause_generation(PauseGenerationReqInput(mode="retract")) - - self.assertIsNone(scheduler.chunked_req) - self.assertFalse(scheduler.running_batch.batch_is_full) + self.assertEqual(scheduler.running_batch.reqs, []) def test_retract_disagg_prefill_keeps_live_chunked_req(self): """disagg-PREFILL retract must leave a live mid-chunk chunked_req untouched.""" @@ -241,17 +439,15 @@ class TestSchedulerPauseGeneration(unittest.TestCase): scheduler = self._new_scheduler() scheduler.disaggregation_mode = DisaggregationMode.DECODE scheduler.last_batch = None - scheduler.running_batch.is_empty.return_value = False scheduler._add_request_to_queue = MagicMock() scheduler.disagg_decode_prealloc_queue = MagicMock() req = SimpleNamespace( + finished=lambda: False, output_ids=[10, 11, 12], time_stats=MagicMock(), ) scheduler.running_batch.reqs = [req] - scheduler.running_batch.filter_batch = MagicMock() - scheduler.server_args = MagicMock() with patch("sglang.srt.managers.scheduler.retract_all") as mock_retract_all: scheduler.pause_generation(PauseGenerationReqInput(mode="retract"))