Rewrite pause_generation retract path as req-level release and requeue for clarity (#30675)

This commit is contained in:
fzyzcjy
2026-07-15 14:32:12 +08:00
committed by GitHub
parent 1967b9ec99
commit 21c62b9830
2 changed files with 282 additions and 97 deletions
+27 -38
View File
@@ -4083,46 +4083,36 @@ class Scheduler(
tmp_batch, tmp_result = self.result_queue.popleft() tmp_batch, tmp_result = self.result_queue.popleft()
self.process_batch_result(tmp_batch, tmp_result) self.process_batch_result(tmp_batch, tmp_result)
if self.last_batch and self.last_batch.forward_mode.is_extend(): retract_reqs = [r for r in self.running_batch.reqs if not r.finished()]
chunked_req_to_exclude = set() if (
self.last_batch.filter_batch( self.last_batch is not None
chunked_req_to_exclude=list(chunked_req_to_exclude) and self.last_batch.forward_mode.is_extend()
)
# Skip merge for disagg prefill: completed prefill requests are # Skip merge for disagg prefill: completed prefill requests are
# already in disagg_prefill_inflight_queue. Merging them into # already in disagg_prefill_inflight_queue. Merging them into
# running_batch leaks them, since the prefill event loop never # running_batch leaks them, since the prefill event loop never
# calls update_running_batch to clean them up. # calls update_running_batch to clean them up.
if ( and self.disaggregation_mode != DisaggregationMode.PREFILL
not self.last_batch.is_empty() ):
and self.disaggregation_mode != DisaggregationMode.PREFILL retract_reqs += [r for r in self.last_batch.reqs if not r.finished()]
):
if self.running_batch.is_empty(): if (
self.running_batch = self.last_batch self.chunked_req is not None
else: and not self.chunked_req.finished()
self.running_batch.merge_batch(self.last_batch) 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.last_batch = None
self.cur_batch_for_debug = None self.cur_batch_for_debug = None
if not self.running_batch.is_empty(): if retract_reqs:
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:
# 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).
retract_all( retract_all(
reqs=retracted_reqs, reqs=retract_reqs,
server_args=self.server_args, server_args=self.server_args,
req_to_token_pool=self.req_to_token_pool, req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
@@ -4130,17 +4120,16 @@ class Scheduler(
hisparse_coordinator=self.hisparse_coordinator, hisparse_coordinator=self.hisparse_coordinator,
offload_kv=False, offload_kv=False,
) )
self.running_batch.reqs = [] self.running_batch.reqs = []
for req in retracted_reqs: for req in retract_reqs:
if self.disaggregation_mode == DisaggregationMode.DECODE: if self.disaggregation_mode == DisaggregationMode.DECODE:
if req.output_ids: if req.output_ids:
req.pd_rebootstrap_forced_output_id = req.output_ids.pop() req.pd_rebootstrap_forced_output_id = req.output_ids.pop()
req.pd_rebootstrap_in_progress = True req.pd_rebootstrap_in_progress = True
req.time_stats.set_retract_time() req.time_stats.set_retract_time()
self.disagg_decode_prealloc_queue.hold_rebootstrap(req) self.disagg_decode_prealloc_queue.hold_rebootstrap(req)
else: else:
self._add_request_to_queue(req) self._add_request_to_queue(req)
self.running_batch.batch_is_full = False self.running_batch.batch_is_full = False
# In disagg-PREFILL, keep a live mid-chunk chunked_req rather than retract it: # 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 # freeing its KV under a live disagg KV-sender crashes pop_bootstrapped or
@@ -1,8 +1,11 @@
import unittest import unittest
from collections import deque from collections import deque
from types import SimpleNamespace from types import SimpleNamespace
from typing import List, Optional
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
import torch
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
@@ -13,8 +16,11 @@ from sglang.srt.managers.io_struct import (
ContinueGenerationReqInput, ContinueGenerationReqInput,
PauseGenerationReqInput, PauseGenerationReqInput,
) )
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
from sglang.srt.managers.scheduler import Scheduler from sglang.srt.managers.scheduler import Scheduler
from sglang.srt.managers.scheduler_components.pool_stats_observer import PoolStats 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=15, suite="base-a-test-cpu")
register_cpu_ci(est_time=9, suite="base-c-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, 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. # pause_generation zeros gen_throughput and flushes KV events.
scheduler.metrics_reporter = MagicMock() scheduler.metrics_reporter = MagicMock()
scheduler.metrics_reporter.current_scheduler_metrics_enabled = False scheduler.metrics_reporter.current_scheduler_metrics_enabled = False
scheduler.kv_events_publisher = MagicMock() scheduler.kv_events_publisher = MagicMock()
return scheduler 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): def test_inplace_only_sets_flag(self):
"""in_place pause should only set _engine_paused and return.""" """in_place pause should only set _engine_paused and return."""
scheduler = self._new_scheduler() scheduler = self._new_scheduler()
@@ -126,82 +184,222 @@ class TestSchedulerPauseGeneration(unittest.TestCase):
self.assertIsNone(scheduler.last_batch) self.assertIsNone(scheduler.last_batch)
self.assertIsNone(scheduler.cur_batch_for_debug) self.assertIsNone(scheduler.cur_batch_for_debug)
def test_retract_clears_running_batch(self): def test_retract_requeues_running_then_last_fold_in(self):
"""retract mode should retract all requests from running_batch.""" """retract requeues running reqs first, then last extend reqs, all released."""
scheduler = self._new_scheduler() scheduler = self._new_scheduler()
scheduler.last_batch = None run_req_a = self._make_req("run-a")
scheduler.running_batch.reqs = [MagicMock(), MagicMock()] run_req_b = self._make_req("run-b")
scheduler.running_batch.__len__ = lambda self: len(self.reqs) last_req = self._make_req("last")
scheduler.running_batch.is_empty.return_value = False scheduler.running_batch = self._make_batch(
scheduler.waiting_queue = [] scheduler, reqs=[run_req_a, run_req_b], with_tensors=True
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,
) )
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) self.assertIsNone(scheduler.chunked_req)
def test_retract_fold_in_releases_via_scheduler_hisparse_coordinator(self): 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.""" """retract of a folded-in last extend batch must release through the scheduler-owned hisparse coordinator."""
scheduler = self._new_scheduler() scheduler = self._new_scheduler()
scheduler.disaggregation_mode = DisaggregationMode.NULL scheduler.hisparse_coordinator = MagicMock()
scheduler.waiting_queue = [] last_req = self._make_req("last")
scheduler._add_request_to_queue = MagicMock() scheduler.running_batch = ScheduleBatch(reqs=[], batch_is_full=True)
scheduler.server_args = MagicMock() scheduler.last_batch = self._make_batch(
scheduler,
req = MagicMock() reqs=[last_req],
req.finished.return_value = False forward_mode=ForwardMode.EXTEND,
req.req_pool_idx = None with_tensors=True,
last_batch = MagicMock() )
last_batch.forward_mode.is_extend.return_value = True requeue_log = self._spy_requeue(scheduler)
last_batch.is_empty.return_value = False
last_batch.reqs = [req]
scheduler.last_batch = last_batch
scheduler.pause_generation(PauseGenerationReqInput(mode="retract")) scheduler.pause_generation(PauseGenerationReqInput(mode="retract"))
scheduler.hisparse_coordinator.retract_req.assert_called_once_with(req) scheduler.hisparse_coordinator.retract_req.assert_called_once_with(last_req)
self.assertEqual( self.assertEqual([entry["req"] for entry in requeue_log], [last_req])
[call.args[0] for call in scheduler._add_request_to_queue.call_args_list],
[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): def test_retract_empty_running_batch_requeues_nothing(self):
"""retract with empty running_batch must not release or requeue any request.""" """retract with empty running_batch must not release or requeue any request."""
scheduler = self._new_scheduler() 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)
self.assertEqual(len(scheduler.waiting_queue), 0) self.assertEqual(len(scheduler.waiting_queue), 0)
self.assertIs(scheduler.running_batch.reqs, original_reqs) self.assertEqual(scheduler.running_batch.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)
def test_retract_disagg_prefill_keeps_live_chunked_req(self): def test_retract_disagg_prefill_keeps_live_chunked_req(self):
"""disagg-PREFILL retract must leave a live mid-chunk chunked_req untouched.""" """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 = self._new_scheduler()
scheduler.disaggregation_mode = DisaggregationMode.DECODE scheduler.disaggregation_mode = DisaggregationMode.DECODE
scheduler.last_batch = None scheduler.last_batch = None
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()
req = SimpleNamespace( req = SimpleNamespace(
finished=lambda: False,
output_ids=[10, 11, 12], output_ids=[10, 11, 12],
time_stats=MagicMock(), time_stats=MagicMock(),
) )
scheduler.running_batch.reqs = [req] 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: with patch("sglang.srt.managers.scheduler.retract_all") as mock_retract_all:
scheduler.pause_generation(PauseGenerationReqInput(mode="retract")) scheduler.pause_generation(PauseGenerationReqInput(mode="retract"))