From c208a96a7dc06056d8281595199914f250796a57 Mon Sep 17 00:00:00 2001 From: luoroger37 Date: Thu, 18 Jun 2026 10:25:47 +0800 Subject: [PATCH] Fix ScheduleBatch req pool CPU metadata (#28514) --- python/sglang/srt/managers/schedule_batch.py | 22 ++++ python/sglang/srt/managers/scheduler.py | 4 +- .../test_schedule_batch_req_pool_indices.py | 121 ++++++++++++++++++ 3 files changed, 146 insertions(+), 1 deletion(-) create mode 100644 test/registered/unit/managers/test_schedule_batch_req_pool_indices.py diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 02eef3460..7e4fdc55c 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -2596,6 +2596,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): def prepare_for_decode(self): self.forward_mode = ForwardMode.DECODE bs = len(self.reqs) + if self.req_pool_indices_cpu is None and self.req_pool_indices is not None: + self.req_pool_indices_cpu = ( + self.req_pool_indices.detach().cpu().to(dtype=torch.int64) + ) # Decode embeds the last output token via embed_tokens; clear the stale # prefill-time tensor so it doesn't leak into ForwardBatch. self.input_embeds = None @@ -2690,6 +2694,15 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): if keep_indices is None or len(keep_indices) == 0: # Filter out all requests self.reqs = [] + self.req_pool_indices = torch.empty( + 0, dtype=torch.int64, device=self.device + ) + self.req_pool_indices_cpu = torch.empty(0, dtype=torch.int64) + self.seq_lens = torch.empty(0, dtype=torch.int64, device=self.device) + self.seq_lens_cpu = torch.empty(0, dtype=torch.int64) + self.orig_seq_lens = torch.empty(0, dtype=torch.int32, device=self.device) + self.out_cache_loc = None + self.seq_lens_sum = 0 return if len(keep_indices) == len(self.reqs): @@ -2747,6 +2760,15 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): ) def merge_batch(self, other: ScheduleBatch): + if self.req_pool_indices_cpu is None and self.req_pool_indices is not None: + self.req_pool_indices_cpu = ( + self.req_pool_indices.detach().cpu().to(dtype=torch.int64) + ) + if other.req_pool_indices_cpu is None and other.req_pool_indices is not None: + other.req_pool_indices_cpu = ( + other.req_pool_indices.detach().cpu().to(dtype=torch.int64) + ) + # Penalizer orchestrator must be merged before Batch.reqs is merged. This is because # orchestrator.merge() depends on Batch.reqs during preparation of each penalizers, so it # needs to be called with pre-merged Batch.reqs. diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index df9deb1f5..f316aef42 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2438,9 +2438,11 @@ class Scheduler( spec_algorithm=self.spec_algorithm, ) + req_pool_indices = [r.req_pool_idx for r in reqs] batch.req_pool_indices = torch.tensor( - [r.req_pool_idx for r in reqs], dtype=torch.int64, device=device + req_pool_indices, dtype=torch.int64, device=device ) + batch.req_pool_indices_cpu = torch.tensor(req_pool_indices, dtype=torch.int64) seq_lens = [len(r.origin_input_ids) + len(r.output_ids) - 1 for r in reqs] batch.seq_lens = torch.tensor(seq_lens, dtype=torch.int64, device=device) batch.seq_lens_cpu = torch.tensor(seq_lens, dtype=torch.int64) diff --git a/test/registered/unit/managers/test_schedule_batch_req_pool_indices.py b/test/registered/unit/managers/test_schedule_batch_req_pool_indices.py new file mode 100644 index 000000000..0b097b7d1 --- /dev/null +++ b/test/registered/unit/managers/test_schedule_batch_req_pool_indices.py @@ -0,0 +1,121 @@ +import types +import unittest +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 + +maybe_stub_sgl_kernel() + +from sglang.srt.managers.schedule_batch import ScheduleBatch # noqa: E402 + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +class TestScheduleBatchReqPoolIndices(unittest.TestCase): + def test_prepare_for_decode_restores_missing_req_pool_indices_cpu(self): + req = types.SimpleNamespace( + decode_batch_idx=0, + kv_committed_len=10, + kv_allocated_len=10, + ) + batch = ScheduleBatch( + reqs=[req], + model_config=types.SimpleNamespace(is_encoder_decoder=False), + req_pool_indices=torch.tensor([4], dtype=torch.int64), + req_pool_indices_cpu=None, + seq_lens=torch.tensor([10], dtype=torch.int64), + seq_lens_cpu=torch.tensor([10], dtype=torch.int64), + orig_seq_lens=torch.tensor([10], dtype=torch.int32), + seq_lens_sum=10, + sampling_info=types.SimpleNamespace( + penalizer_orchestrator=types.SimpleNamespace(is_required=False) + ), + spec_algorithm=types.SimpleNamespace(is_none=lambda: True), + enable_overlap=False, + device="cpu", + hisparse_coordinator=MagicMock(), + ) + + with ( + patch( + "sglang.srt.managers.schedule_batch.alloc_for_decode", + return_value=torch.tensor([42], dtype=torch.int64), + ), + patch( + "sglang.srt.managers.schedule_batch.get_global_server_args", + return_value=types.SimpleNamespace( + enable_mamba_extra_buffer=lambda: False + ), + ), + ): + batch.prepare_for_decode() + + self.assertTrue(torch.equal(batch.req_pool_indices_cpu, torch.tensor([4]))) + batch.hisparse_coordinator.map_last_loc_to_buffer.assert_called_once() + + def test_filter_batch_to_empty_clears_req_pool_metadata(self): + req = types.SimpleNamespace(finished=lambda: True) + batch = ScheduleBatch( + reqs=[req], + model_config=types.SimpleNamespace(is_encoder_decoder=False), + req_pool_indices=torch.tensor([4], dtype=torch.int64), + req_pool_indices_cpu=torch.tensor([4], dtype=torch.int64), + seq_lens=torch.tensor([10], dtype=torch.int64), + seq_lens_cpu=torch.tensor([10], dtype=torch.int64), + orig_seq_lens=torch.tensor([10], dtype=torch.int32), + seq_lens_sum=10, + device="cpu", + ) + + batch.filter_batch() + + self.assertEqual(batch.req_pool_indices.numel(), 0) + self.assertEqual(batch.req_pool_indices_cpu.numel(), 0) + self.assertEqual(batch.seq_lens.numel(), 0) + self.assertEqual(batch.seq_lens_cpu.numel(), 0) + self.assertEqual(batch.seq_lens_sum, 0) + + def test_merge_batch_restores_missing_req_pool_indices_cpu(self): + self_batch = ScheduleBatch( + reqs=[object(), object()], + model_config=types.SimpleNamespace(is_encoder_decoder=False), + req_pool_indices=torch.tensor([1, 2], dtype=torch.int64), + req_pool_indices_cpu=None, + seq_lens=torch.tensor([10, 20], dtype=torch.int64), + seq_lens_cpu=torch.tensor([10, 20], dtype=torch.int64), + orig_seq_lens=torch.tensor([10, 20], dtype=torch.int32), + seq_lens_sum=30, + sampling_info=MagicMock(), + return_logprob=False, + has_grammar=False, + return_hidden_states=False, + is_prefill_only=False, + ) + other_batch = ScheduleBatch( + reqs=[object()], + model_config=types.SimpleNamespace(is_encoder_decoder=False), + req_pool_indices=torch.tensor([3], dtype=torch.int64), + req_pool_indices_cpu=None, + seq_lens=torch.tensor([30], dtype=torch.int64), + seq_lens_cpu=torch.tensor([30], dtype=torch.int64), + orig_seq_lens=torch.tensor([30], dtype=torch.int32), + seq_lens_sum=30, + sampling_info=MagicMock(), + return_logprob=False, + has_grammar=False, + return_hidden_states=False, + is_prefill_only=False, + ) + + self_batch.merge_batch(other_batch) + + self.assertTrue( + torch.equal(self_batch.req_pool_indices_cpu, torch.tensor([1, 2, 3])) + ) + + +if __name__ == "__main__": + unittest.main()