From aea527afdc1b3196844dfaa9fe3aaceed1d78e77 Mon Sep 17 00:00:00 2001 From: Ke Bao Date: Mon, 4 May 2026 14:52:15 +0800 Subject: [PATCH] Fix swa chunk req deferred (#24318) --- python/sglang/srt/managers/scheduler.py | 21 ++- .../test_scheduler_chunked_req_gate.py | 161 ++++++++++++++++++ 2 files changed, 181 insertions(+), 1 deletion(-) create mode 100644 test/registered/unit/managers/test_scheduler_chunked_req_gate.py diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index dc44902e4..1d4b9fd14 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -1080,6 +1080,15 @@ class Scheduler( elif self.chunked_prefill_size is not None and self.chunked_prefill_size <= 0: self.chunked_prefill_size = None self.chunked_req = None + # Tracks whether the current self.chunked_req was actually scheduled + # into last iteration's batch (i.e., in can_run_list -> got a fresh + # req_pool_idx from prepare_for_extend). Used to gate the + # stash_chunked_request call at the top of get_next_batch_to_run: + # if add_chunked_req early-returned under hybrid-SWA pressure, + # the req_pool_idx was already freed and fill_ids was reset by + # init_next_round_input, so running stash would double-free and + # corrupt prefix_indices. + self._chunked_req_scheduled_last_iter = False self.is_mixed_chunk = ( self.chunked_prefill_size is not None and self.server_args.enable_mixed_chunk @@ -2455,7 +2464,9 @@ class Scheduler( # Move the chunked request out of the batch so that we can merge # only finished requests to running_batch. chunked_req_to_exclude.add(self.chunked_req) - self.stash_chunked_request(self.chunked_req) + + if self._chunked_req_scheduled_last_iter: + self.stash_chunked_request(self.chunked_req) # HiSparse has its own prefill-to-decode transition; skip last_batch merge. if self.enable_hisparse: @@ -2650,6 +2661,11 @@ class Scheduler( if self.chunked_req is not None: self.chunked_req.init_next_round_input() self.chunked_req = adder.add_chunked_req(self.chunked_req) + self._chunked_req_scheduled_last_iter = ( + self.chunked_req in adder.can_run_list + ) + else: + self._chunked_req_scheduled_last_iter = False if self.enable_lora: running_loras = {req.lora_id for req in self.running_batch.reqs} @@ -2741,6 +2757,9 @@ class Scheduler( # Update chunked prefill assert self.chunked_req is None self.chunked_req = adder.new_chunked_req + # new_chunked_req is added to can_run_list by add_one_req, + # so it will be scheduled this iter -> stash is needed next iter. + self._chunked_req_scheduled_last_iter = True if self.chunked_req is not None: self.chunked_req.is_chunked += 1 diff --git a/test/registered/unit/managers/test_scheduler_chunked_req_gate.py b/test/registered/unit/managers/test_scheduler_chunked_req_gate.py new file mode 100644 index 000000000..87a6daf7e --- /dev/null +++ b/test/registered/unit/managers/test_scheduler_chunked_req_gate.py @@ -0,0 +1,161 @@ +"""Regression tests for the SWA chunked-req stash gate (#24252).""" + +import unittest +from types import SimpleNamespace +from unittest.mock import MagicMock + +import torch + +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel + +maybe_stub_sgl_kernel() + +from sglang.srt.managers.schedule_batch import Req +from sglang.srt.managers.scheduler import Scheduler +from sglang.srt.mem_cache.chunk_cache import ChunkCache + +register_cpu_ci(est_time=6, suite="stage-a-test-cpu") + + +def _make_req( + *, + req_pool_idx: int, + fill_ids: list, + prefix_indices: torch.Tensor, + extend_input_len: int, +) -> Req: + req = Req.__new__(Req) + req.rid = "test-req" + req.origin_input_ids = list(fill_ids) + req.output_ids = [] + req.fill_ids = list(fill_ids) + req.prefix_indices = prefix_indices + req.req_pool_idx = req_pool_idx + req.extend_input_len = extend_input_len + req.is_chunked = 0 + req.host_hit_length = 0 + req.cache_protected_len = 0 + req.skip_radix_cache_insert = False + req.last_node = None + req.swa_uuid_for_lock = None + req.session = None + req.return_logprob = False + req.logprob_start_len = -1 + req.positional_embed_overrides = None + req.extra_key = None + req.mamba_pool_idx = None + req.sampling_params = SimpleNamespace(max_new_tokens=128, ignore_eos=False) + return req + + +def _make_req_to_token_pool(num_slots: int, max_context: int) -> SimpleNamespace: + # Slot s contains a recognizable fingerprint [s*1000, s*1000+1, ...] + # so we can tell a corrupted prefix_indices from a healthy one by content. + pool = SimpleNamespace() + pool.req_to_token = ( + torch.arange(max_context, dtype=torch.int32).unsqueeze(0).repeat(num_slots, 1) + + torch.arange(num_slots, dtype=torch.int32).unsqueeze(1) * 1000 + ) + return pool + + +def _make_chunk_cache(req_to_token_pool) -> ChunkCache: + return ChunkCache( + SimpleNamespace( + req_to_token_pool=req_to_token_pool, + token_to_kv_pool_allocator=None, + page_size=1, + ) + ) + + +def _scheduler_for_get_next_batch(*, tree_cache, chunked_req) -> Scheduler: + s = Scheduler.__new__(Scheduler) + s._abort_on_waiting_timeout = MagicMock() + s._abort_on_running_timeout = MagicMock() + s.dllm_config = None + s.dllm_manager = None + s.enable_hisparse = False + s.last_batch = None + s.require_mlp_sync = False + s.spec_algorithm = MagicMock() + s.server_args = MagicMock(speculative_skip_dp_mlp_sync=True) + s.running_batch = MagicMock() + s.running_batch.is_empty.return_value = True + s.running_batch.is_prefill_only = False + s.running_batch.batch_is_full = False + s.running_batch.reqs = [] + s.get_new_batch_prefill = MagicMock(return_value=None) + s.maybe_prepare_mlp_sync_batch = MagicMock(side_effect=lambda batch, **_: batch) + s._maybe_prepare_ngram_embedding = MagicMock(side_effect=lambda batch: batch) + s.update_running_batch = MagicMock(side_effect=lambda batch: batch) + s.tree_cache = tree_cache + s.chunked_req = chunked_req + return s + + +class TestStashGatePreservesPrefixIndices(CustomTestCase): + """Consumer side: real ChunkCache.cache_unfinished_req mutates + req.prefix_indices iff stash actually runs, so prefix_indices content + is the bug-detection signal.""" + + POOL_IDX = 4 + INITIAL_PREFIX_LEN = 8 # what was really cached last iter + POST_RESET_FILL_LEN = 32 # length after init_next_round_input + NUM_SLOTS = 8 + MAX_CONTEXT = 64 + + def _build(self, flag: bool): + pool = _make_req_to_token_pool(self.NUM_SLOTS, self.MAX_CONTEXT) + cache = _make_chunk_cache(pool) + initial_prefix = pool.req_to_token[self.POOL_IDX, : self.INITIAL_PREFIX_LEN].to( + dtype=torch.int64, copy=True + ) + req = _make_req( + req_pool_idx=self.POOL_IDX, + fill_ids=list(range(self.POST_RESET_FILL_LEN)), + prefix_indices=initial_prefix, + extend_input_len=0, + ) + s = _scheduler_for_get_next_batch(tree_cache=cache, chunked_req=req) + s._chunked_req_scheduled_last_iter = flag + return s, req, initial_prefix, pool + + def test_deferred_chunked_req_keeps_real_prefix_indices(self): + # The bug case: a spurious stash on a deferred chunked_req + # would extend prefix_indices to len(fill_ids). + s, req, initial_prefix, _ = self._build(flag=False) + + Scheduler.get_next_batch_to_run(s) + + self.assertEqual(req.prefix_indices.shape[0], self.INITIAL_PREFIX_LEN) + self.assertTrue(torch.equal(req.prefix_indices, initial_prefix)) + + def test_scheduled_chunked_req_advances_prefix_indices_via_real_stash(self): + # Symmetric guard against over-gating: when the chunked_req was + # actually scheduled, stash must run and advance prefix_indices. + s, req, _, pool = self._build(flag=True) + + Scheduler.get_next_batch_to_run(s) + + expected = pool.req_to_token[self.POOL_IDX, : self.POST_RESET_FILL_LEN].to( + dtype=torch.int64 + ) + self.assertEqual(req.prefix_indices.shape[0], self.POST_RESET_FILL_LEN) + self.assertTrue(torch.equal(req.prefix_indices, expected)) + + def test_no_chunked_req_never_mutates_state_even_with_stale_flag(self): + # Retract path clears chunked_req without resetting the flag; + # the outer `if chunked_req is not None` guard must hold. + pool = _make_req_to_token_pool(self.NUM_SLOTS, self.MAX_CONTEXT) + cache = _make_chunk_cache(pool) + s = _scheduler_for_get_next_batch(tree_cache=cache, chunked_req=None) + s._chunked_req_scheduled_last_iter = True + + Scheduler.get_next_batch_to_run(s) + self.assertIsNone(s.chunked_req) + + +if __name__ == "__main__": + unittest.main()