Fix non-existent abort mode in Scheduler.pause_generation and inline retract_all (#30673)

This commit is contained in:
fzyzcjy
2026-07-15 14:27:48 +08:00
committed by GitHub
parent 52a88fb212
commit b6cc897fea
3 changed files with 75 additions and 49 deletions
+1 -16
View File
@@ -1724,8 +1724,7 @@ def retract_all(
tree_cache: BasePrefixCache, tree_cache: BasePrefixCache,
hisparse_coordinator: Optional[HiSparseCoordinator], hisparse_coordinator: Optional[HiSparseCoordinator],
offload_kv: bool = True, offload_kv: bool = True,
) -> List[Req]: ) -> None:
retracted_reqs = reqs
for idx in range(len(reqs)): for idx in range(len(reqs)):
release_req( release_req(
req=reqs[idx], req=reqs[idx],
@@ -1737,7 +1736,6 @@ def retract_all(
hisparse_coordinator=hisparse_coordinator, hisparse_coordinator=hisparse_coordinator,
offload_kv=offload_kv, offload_kv=offload_kv,
) )
return retracted_reqs
def compute_extend_logprob_start_len( def compute_extend_logprob_start_len(
@@ -2571,19 +2569,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
evict_from_tree_cache(self.tree_cache, num_tokens) evict_from_tree_cache(self.tree_cache, num_tokens)
return self.token_to_kv_pool_allocator.available_size() >= num_tokens return self.token_to_kv_pool_allocator.available_size() >= num_tokens
def retract_all(self, server_args: ServerArgs, offload_kv: bool = True):
retracted_reqs = retract_all(
reqs=self.reqs,
server_args=server_args,
req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
tree_cache=self.tree_cache,
hisparse_coordinator=self.hisparse_coordinator,
offload_kv=offload_kv,
)
self.reqs = []
return retracted_reqs
def retract_decode( def retract_decode(
self, server_args: ServerArgs self, server_args: ServerArgs
) -> Tuple[List[Req], float, List[Req]]: ) -> Tuple[List[Req], float, List[Req]]:
+13 -3
View File
@@ -166,6 +166,7 @@ from sglang.srt.managers.schedule_batch import (
NextBatchPlan, NextBatchPlan,
Req, Req,
ScheduleBatch, ScheduleBatch,
retract_all,
) )
from sglang.srt.managers.schedule_policy import ( from sglang.srt.managers.schedule_policy import (
AddReqResult, AddReqResult,
@@ -4064,6 +4065,7 @@ class Scheduler(
raise NotImplementedError() raise NotImplementedError()
def pause_generation(self, recv_req: PauseGenerationReqInput): def pause_generation(self, recv_req: PauseGenerationReqInput):
assert recv_req.mode in ("in_place", "retract")
self._engine_paused = True self._engine_paused = True
if recv_req.mode == "in_place": if recv_req.mode == "in_place":
@@ -4102,16 +4104,24 @@ class Scheduler(
self.last_batch = None self.last_batch = None
self.cur_batch_for_debug = None self.cur_batch_for_debug = None
if recv_req.mode == "retract" and not self.running_batch.is_empty(): if not self.running_batch.is_empty():
self.running_batch.filter_batch() self.running_batch.filter_batch()
if len(self.running_batch.reqs) != 0: if len(self.running_batch.reqs) != 0:
# Decode-side retract always rebootstraps (recomputes the KV from # Decode-side retract always rebootstraps (recomputes the KV from
# the prefill), so skip the device->host KV offload that release_req # the prefill), so skip the device->host KV offload that release_req
# would otherwise do; the offloaded copy would be immediately # would otherwise do; the offloaded copy would be immediately
# discarded. Non-decode modes ignore offload_kv (they never offload). # discarded. Non-decode modes ignore offload_kv (they never offload).
retracted_reqs = self.running_batch.retract_all( retracted_reqs = self.running_batch.reqs
self.server_args, offload_kv=False retract_all(
reqs=retracted_reqs,
server_args=self.server_args,
req_to_token_pool=self.running_batch.req_to_token_pool,
token_to_kv_pool_allocator=self.running_batch.token_to_kv_pool_allocator,
tree_cache=self.running_batch.tree_cache,
hisparse_coordinator=self.running_batch.hisparse_coordinator,
offload_kv=False,
) )
self.running_batch.reqs = []
for req in retracted_reqs: for req in retracted_reqs:
if self.disaggregation_mode == DisaggregationMode.DECODE: if self.disaggregation_mode == DisaggregationMode.DECODE:
if req.output_ids: if req.output_ids:
@@ -1,7 +1,7 @@
import unittest import unittest
from collections import deque from collections import deque
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import MagicMock from unittest.mock import MagicMock, patch
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import maybe_stub_sgl_kernel from sglang.test.test_utils import maybe_stub_sgl_kernel
@@ -98,14 +98,28 @@ class TestSchedulerPauseGeneration(unittest.TestCase):
last_batch.filter_batch.assert_not_called() last_batch.filter_batch.assert_not_called()
scheduler.running_batch.merge_batch.assert_not_called() scheduler.running_batch.merge_batch.assert_not_called()
def test_abort_clears_state(self): def test_abort_mode_rejected_at_scheduler(self):
"""abort mode should clear last_batch and cur_batch_for_debug.""" """abort mode must be rejected by the scheduler-side assert."""
scheduler = self._new_scheduler()
with self.assertRaises(AssertionError):
scheduler.pause_generation(PauseGenerationReqInput(mode="abort"))
def test_default_mode_rejected_at_scheduler(self):
"""bare PauseGenerationReqInput defaults to abort and must be rejected."""
scheduler = self._new_scheduler()
with self.assertRaises(AssertionError):
scheduler.pause_generation(PauseGenerationReqInput())
def test_retract_clears_last_batch_state(self):
"""retract mode should clear last_batch and cur_batch_for_debug."""
scheduler = self._new_scheduler() scheduler = self._new_scheduler()
scheduler.last_batch = MagicMock() scheduler.last_batch = MagicMock()
scheduler.last_batch.forward_mode.is_extend.return_value = False scheduler.last_batch.forward_mode.is_extend.return_value = False
scheduler.cur_batch_for_debug = MagicMock() scheduler.cur_batch_for_debug = MagicMock()
scheduler.pause_generation(PauseGenerationReqInput(mode="abort")) scheduler.pause_generation(PauseGenerationReqInput(mode="retract"))
self.assertTrue(scheduler._engine_paused) self.assertTrue(scheduler._engine_paused)
self.assertIsNone(scheduler.last_batch) self.assertIsNone(scheduler.last_batch)
@@ -121,24 +135,56 @@ class TestSchedulerPauseGeneration(unittest.TestCase):
scheduler.waiting_queue = [] scheduler.waiting_queue = []
scheduler._add_request_to_queue = MagicMock() scheduler._add_request_to_queue = MagicMock()
retracted = [MagicMock(), MagicMock()]
scheduler.running_batch.retract_all.return_value = retracted
scheduler.running_batch.filter_batch = MagicMock() scheduler.running_batch.filter_batch = MagicMock()
scheduler.server_args = 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,
)
self.assertIsNone(scheduler.chunked_req)
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")) scheduler.pause_generation(PauseGenerationReqInput(mode="retract"))
self.assertTrue(scheduler._engine_paused) self.assertTrue(scheduler._engine_paused)
scheduler.running_batch.retract_all.assert_called_once() self.assertEqual(len(scheduler.waiting_queue), 0)
self.assertEqual(scheduler._add_request_to_queue.call_count, 2) self.assertIs(scheduler.running_batch.reqs, original_reqs)
self.assertIsNone(scheduler.chunked_req)
def test_retract_drains_overlap_queue(self):
"""retract with overlap enabled should drain the result_queue."""
scheduler = self._new_scheduler()
scheduler.enable_overlap = True
mock_batch = MagicMock()
mock_batch.forward_mode.is_extend.return_value = False
scheduler.last_batch = mock_batch
scheduler.result_queue = deque([(MagicMock(), MagicMock())])
scheduler.process_batch_result = MagicMock()
scheduler.pause_generation(PauseGenerationReqInput(mode="retract"))
scheduler.process_batch_result.assert_called_once()
self.assertEqual(len(scheduler.result_queue), 0)
def test_pd_decode_retract_requeues_for_rebootstrap(self): def test_pd_decode_retract_requeues_for_rebootstrap(self):
"""PD decode retract should rebootstrap instead of resuming stale CPU KV.""" """PD decode retract should rebootstrap instead of resuming stale CPU KV."""
scheduler = self._new_scheduler() scheduler = self._new_scheduler()
scheduler.disaggregation_mode = DisaggregationMode.DECODE scheduler.disaggregation_mode = DisaggregationMode.DECODE
scheduler.last_batch = None scheduler.last_batch = None
scheduler.running_batch.reqs = [MagicMock()]
scheduler.running_batch.is_empty.return_value = False scheduler.running_batch.is_empty.return_value = False
scheduler._add_request_to_queue = MagicMock() scheduler._add_request_to_queue = MagicMock()
scheduler.disagg_decode_prealloc_queue = MagicMock() scheduler.disagg_decode_prealloc_queue = MagicMock()
@@ -147,11 +193,12 @@ class TestSchedulerPauseGeneration(unittest.TestCase):
output_ids=[10, 11, 12], output_ids=[10, 11, 12],
time_stats=MagicMock(), time_stats=MagicMock(),
) )
scheduler.running_batch.retract_all.return_value = [req] scheduler.running_batch.reqs = [req]
scheduler.running_batch.filter_batch = MagicMock() scheduler.running_batch.filter_batch = MagicMock()
scheduler.server_args = MagicMock() scheduler.server_args = MagicMock()
scheduler.pause_generation(PauseGenerationReqInput(mode="retract")) with patch("sglang.srt.managers.scheduler.retract_all") as mock_retract_all:
scheduler.pause_generation(PauseGenerationReqInput(mode="retract"))
scheduler._add_request_to_queue.assert_not_called() scheduler._add_request_to_queue.assert_not_called()
scheduler.disagg_decode_prealloc_queue.hold_rebootstrap.assert_called_once_with( scheduler.disagg_decode_prealloc_queue.hold_rebootstrap.assert_called_once_with(
@@ -162,9 +209,8 @@ class TestSchedulerPauseGeneration(unittest.TestCase):
self.assertTrue(req.pd_rebootstrap_in_progress) self.assertTrue(req.pd_rebootstrap_in_progress)
# Rebootstrap recomputes the KV from the prefill, so the retract must skip # Rebootstrap recomputes the KV from the prefill, so the retract must skip
# the device->host KV offload rather than offload-then-delete it. # the device->host KV offload rather than offload-then-delete it.
scheduler.running_batch.retract_all.assert_called_once_with( mock_retract_all.assert_called_once()
scheduler.server_args, offload_kv=False self.assertEqual(mock_retract_all.call_args.kwargs["offload_kv"], False)
)
def test_pd_decode_continue_releases_held_rebootstrap(self): def test_pd_decode_continue_releases_held_rebootstrap(self):
"""continue_generation must enqueue staged rebootstrap reqs on resume.""" """continue_generation must enqueue staged rebootstrap reqs on resume."""
@@ -180,21 +226,6 @@ class TestSchedulerPauseGeneration(unittest.TestCase):
scheduler.disagg_decode_prealloc_queue.enqueue_held_rebootstrap.assert_called_once_with() scheduler.disagg_decode_prealloc_queue.enqueue_held_rebootstrap.assert_called_once_with()
self.assertFalse(scheduler._engine_paused) self.assertFalse(scheduler._engine_paused)
def test_abort_drains_overlap_queue(self):
"""abort with overlap enabled should drain the result_queue."""
scheduler = self._new_scheduler()
scheduler.enable_overlap = True
mock_batch = MagicMock()
mock_batch.forward_mode.is_extend.return_value = False
scheduler.last_batch = mock_batch
scheduler.result_queue = deque([(MagicMock(), MagicMock())])
scheduler.process_batch_result = MagicMock()
scheduler.pause_generation(PauseGenerationReqInput(mode="abort"))
scheduler.process_batch_result.assert_called_once()
self.assertEqual(len(scheduler.result_queue), 0)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()