diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index c86fa5a37..1da1a9684 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -2602,10 +2602,6 @@ 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 @@ -2699,17 +2695,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): ] if keep_indices is None or len(keep_indices) == 0: - # Filter out all requests + # Filter out all requests. Stale tensors are left as-is: is_empty() + # keys off reqs, so callers drop the batch before a forward reads them. 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): @@ -2767,15 +2755,6 @@ 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/test/registered/8-gpu-models/test_dsa_models_hisparse.py b/test/registered/8-gpu-models/test_dsa_models_hisparse.py deleted file mode 100644 index a6a973cf3..000000000 --- a/test/registered/8-gpu-models/test_dsa_models_hisparse.py +++ /dev/null @@ -1,84 +0,0 @@ -import unittest -from types import SimpleNamespace - -from sglang.srt.utils import kill_process_tree -from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.run_eval import run_eval -from sglang.test.test_utils import ( - DEFAULT_URL_FOR_TEST, - CustomTestCase, - is_in_ci, - popen_launch_server, - write_github_step_summary, -) - -register_cuda_ci(est_time=720, stage="base-c", runner_config="8-gpu-h200", nightly=True) - -GLM5_MODEL_PATH = "zai-org/GLM-5-FP8" - - -class TestGLM5DPHiSparse(CustomTestCase): - @classmethod - def setUpClass(cls): - cls.model = GLM5_MODEL_PATH - cls.base_url = DEFAULT_URL_FOR_TEST - other_args = [ - "--trust-remote-code", - "--tp", - "8", - "--dp", - "8", - "--enable-dp-attention", - "--page-size", - "64", - "--max-running-requests", - "200", - "--mem-fraction-static", - "0.85", - "--disable-radix-cache", - "--kv-cache-dtype", - "bfloat16", - "--dsa-decode-backend", - "flashmla_sparse", - "--enable-hisparse", - "--hisparse-config", - '{"top_k": 2048, "device_buffer_size": 4096, "host_to_device_ratio": 5}', - "--model-loader-extra-config", - '{"enable_multithread_load": true, "num_threads": 64}', - ] - cls.process = popen_launch_server( - cls.model, - cls.base_url, - timeout=7200, - other_args=other_args, - ) - - @classmethod - def tearDownClass(cls): - kill_process_tree(cls.process.pid) - - def test_a_gsm8k( - self, - ): # Append an "a" to make this test run first (alphabetically) to warm up the server - args = SimpleNamespace( - base_url=self.base_url, - model=self.model, - eval_name="gsm8k", - api="completion", - max_tokens=4000, - num_examples=500, - num_threads=100, - num_shots=24, - ) - metrics = run_eval(args) - print(f"{metrics=}") - - if is_in_ci(): - write_github_step_summary( - f"### test_gsm8k (glm-5 hisparse)\n" f'{metrics["score"]=:.3f}\n' - ) - self.assertGreater(metrics["score"], 0.94) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/models_e2e/test_dsa_glm5_hisparse.py b/test/registered/models_e2e/test_dsa_glm5_hisparse.py new file mode 100644 index 000000000..2e5c3c819 --- /dev/null +++ b/test/registered/models_e2e/test_dsa_glm5_hisparse.py @@ -0,0 +1,54 @@ +import unittest + +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kits.eval_accuracy_kit import GSM8KMixin +from sglang.test.server_fixtures.default_fixture import DefaultServerBase + +register_cuda_ci(est_time=720, stage="extra-b", runner_config="8-gpu-h200") + +GLM5_FP8_MODEL_PATH = "zai-org/GLM-5-FP8" + + +class TestGLM5HiSparse(DefaultServerBase, GSM8KMixin): + """GLM-5 FP8 with HiSparse (host-to-device sparse KV offload) on DSA decode. + + HiSparse targets the high-concurrency regime and is not used together with + EAGLE MTP, so this variant runs without speculative decoding (unlike the + DSA-MTP variants in test_dsa_glm5_{dp,tp}_mtp.py). + """ + + model = GLM5_FP8_MODEL_PATH + other_args = [ + "--trust-remote-code", + "--tp", + "8", + "--dp", + "8", + "--enable-dp-attention", + "--page-size", + "64", + "--max-running-requests", + "200", + "--mem-fraction-static", + "0.85", + "--disable-radix-cache", + "--kv-cache-dtype", + "bfloat16", + "--dsa-decode-backend", + "flashmla_sparse", + "--enable-hisparse", + "--hisparse-config", + '{"top_k": 2048, "device_buffer_size": 4096, "host_to_device_ratio": 5}', + "--model-loader-extra-config", + '{"enable_multithread_load": true, "num_threads": 64}', + ] + + # Match the original standalone hisparse eval config. + gsm8k_accuracy_thres = 0.94 + gsm8k_num_questions = 500 + gsm8k_num_threads = 100 + gsm8k_num_shots = 24 + + +if __name__ == "__main__": + unittest.main() 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 index 0b097b7d1..45316e324 100644 --- a/test/registered/unit/managers/test_schedule_batch_req_pool_indices.py +++ b/test/registered/unit/managers/test_schedule_batch_req_pool_indices.py @@ -9,113 +9,85 @@ from sglang.test.test_utils import maybe_stub_sgl_kernel maybe_stub_sgl_kernel() -from sglang.srt.managers.schedule_batch import ScheduleBatch # noqa: E402 +from sglang.srt.managers.hisparse_coordinator import HiSparseCoordinator # noqa: E402 +from sglang.srt.managers.scheduler import Scheduler # 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(), - ) +def _make_req(req_pool_idx, origin_input_ids, output_ids): + return types.SimpleNamespace( + req_pool_idx=req_pool_idx, + origin_input_ids=origin_input_ids, + output_ids=output_ids, + return_logprob=False, + grammar=None, + return_hidden_states=False, + is_prefill_only=False, + ) - 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 - ), - ), + +class TestHisparseDecodeBatchReqPoolCpu(unittest.TestCase): + def test_build_hisparse_decode_batch_populates_req_pool_indices_cpu(self): + # _build_hisparse_decode_batch builds a ScheduleBatch off the normal + # extend path, so it must populate the req_pool_indices_cpu host mirror + # in lockstep with the device tensor. A missing mirror crashes hisparse + # decode bookkeeping (map_last_loc_to_buffer -> _grow_device_buffers + # indexes req_pool_indices_cpu). + scheduler = Scheduler.__new__(Scheduler) + scheduler.device = "cpu" + scheduler.req_to_token_pool = types.SimpleNamespace(device="cpu") + scheduler.token_to_kv_pool_allocator = None + scheduler.tree_cache = None + scheduler.model_config = types.SimpleNamespace( + is_encoder_decoder=False, vocab_size=32 + ) + scheduler.enable_overlap = False + scheduler.spec_algorithm = types.SimpleNamespace(is_none=lambda: True) + scheduler.future_map = MagicMock() + + reqs = [ + _make_req(req_pool_idx=4, origin_input_ids=[1, 2, 3], output_ids=[7]), + _make_req(req_pool_idx=9, origin_input_ids=[1, 2], output_ids=[8]), + ] + + with patch( + "sglang.srt.managers.scheduler.SamplingBatchInfo.from_schedule_batch", + return_value=MagicMock(), ): - 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) + batch = scheduler._build_hisparse_decode_batch(reqs) + # Assert the invariant (cpu mirror == device tensor), not a hardcoded + # copy of the input -- the latter would just restate the builder line. + self.assertIsNotNone(batch.req_pool_indices_cpu) self.assertTrue( - torch.equal(self_batch.req_pool_indices_cpu, torch.tensor([1, 2, 3])) + torch.equal(batch.req_pool_indices_cpu, batch.req_pool_indices.cpu()) ) +class TestHisparseCoordinatorReqPoolCpu(unittest.TestCase): + def test_host_bookkeeping_requires_req_pool_indices_cpu(self): + # Why the mirror must exist: hisparse host bookkeeping indexes + # req_pool_indices_cpu element-wise (int(req_pool_indices_cpu[i]) in + # _eager_backup_previous_token, the first thing map_last_loc_to_buffer + # runs each decode step). A missing mirror (None) raises TypeError there + # -- the exact nightly failure the scheduler-side fix prevents. Asserting + # the crash directly keeps the "mirror is required" contract honest, with + # no mocked attributes (the crash precedes any self access). + coord = HiSparseCoordinator.__new__(HiSparseCoordinator) + seq_lens = torch.tensor([10], dtype=torch.int64) + seq_lens_cpu = torch.tensor([10], dtype=torch.int64) + req_pool_indices = torch.tensor([0], dtype=torch.int64) + out_cache_loc = torch.tensor([0], dtype=torch.int64) + with self.assertRaises(TypeError): + coord.map_last_loc_to_buffer( + seq_lens, + out_cache_loc, + req_pool_indices, + seq_lens_cpu, + req_pool_indices_cpu=None, + ) + + if __name__ == "__main__": unittest.main()