diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 0bdd51034..0a2d8b36a 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -656,9 +656,14 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin): return len(req.origin_input_ids) + max(len(req.output_ids) - 1, 0) def _check_if_req_exceed_kv_capacity(self, req: Req) -> bool: + # HiSparse admits up to the host-backed logical capacity. + if self.scheduler.enable_hisparse: + capacity = self.scheduler.tp_worker.model_runner.max_token_pool_size + else: + capacity = self.max_total_num_tokens input_len = self._rebootstrap_prefill_len(req) - if input_len > self.max_total_num_tokens: - message = f"Request {req.rid} exceeds the maximum number of tokens: {input_len} > {self.max_total_num_tokens}" + if input_len > capacity: + message = f"Request {req.rid} exceeds the maximum number of tokens: {input_len} > {capacity}" logger.error(message) prepare_abort(req, message, status_code=HTTPStatus.BAD_REQUEST) self.scheduler.output_streamer.stream_output([req], req.return_logprob) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 8aa3f6551..fc51e70de 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -997,7 +997,7 @@ class ModelRunner: RoutedExpertsCapturer.create( model=self.model, model_config=self.model_config, - num_tokens=self.max_total_num_tokens + self.page_size, + num_tokens=self.max_token_pool_size + self.page_size, max_running_requests=self.max_running_requests, device=self.device, ) @@ -1007,7 +1007,7 @@ class ModelRunner: set_global_indexer_capturer( create_indexer_capturer( model_config=self.model_config, - num_tokens=self.max_total_num_tokens + self.page_size, + num_tokens=self.max_token_pool_size + self.page_size, max_running_requests=self.max_running_requests, device=self.device, ) @@ -1244,6 +1244,16 @@ class ModelRunner: else: return self.max_total_num_tokens + @property + def max_token_pool_size(self): + """Return the max token pool size considering hybrid swa and hisparse settings.""" + if self.enable_hisparse: + # HiSparse uses the host-backed full pool capacity. + size_full = getattr(self.token_to_kv_pool_allocator, "size_full", None) + if size_full is not None: + return size_full + return self.effective_max_total_num_tokens + def _load_format_scope(self, load_format: Optional[str]): """Make this runner's load format the published one while it loads. diff --git a/test/registered/unit/mem_cache/test_hisparse_max_token_pool_size.py b/test/registered/unit/mem_cache/test_hisparse_max_token_pool_size.py new file mode 100644 index 000000000..eb76a1e56 --- /dev/null +++ b/test/registered/unit/mem_cache/test_hisparse_max_token_pool_size.py @@ -0,0 +1,193 @@ +"""Unit tests for HiSparse-aware max token pool sizing. + +Covers the HiSparse host-backed capacity fix: +- `ModelRunner.max_token_pool_size` returns the allocator's `size_full` when + `enable_hisparse` is set (host-backed logical pool), otherwise it delegates + to `effective_max_total_num_tokens`. +- `DecodePreallocQueue._check_if_req_exceed_kv_capacity` uses that ratio-expanded + capacity for admission when HiSparse is enabled, so long-context inputs are + not truncated at the device-only `max_total_num_tokens`. +""" + +import unittest +from types import SimpleNamespace +from unittest.mock import MagicMock + +from sglang.srt.disaggregation.decode import DecodePreallocQueue +from sglang.srt.model_executor.model_runner import ModelRunner +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +def _make_model_runner(**attrs): + """Build a bare ModelRunner (bypass __init__) so that property descriptors + like `max_token_pool_size` and `effective_max_total_num_tokens` resolve via + normal attribute lookup — a plain SimpleNamespace would bypass them and + raise AttributeError on the internal `self.effective_max_total_num_tokens` + read inside `max_token_pool_size`.""" + instance = object.__new__(ModelRunner) + for name, value in attrs.items(): + setattr(instance, name, value) + return instance + + +class TestMaxTokenPoolSize(CustomTestCase): + def test_hisparse_returns_allocator_size_full(self): + """When HiSparse is enabled and the allocator exposes `size_full`, the + host-backed logical capacity (device_pool * host_to_device_ratio) wins + over `effective_max_total_num_tokens`.""" + instance = _make_model_runner( + enable_hisparse=True, + token_to_kv_pool_allocator=SimpleNamespace(size_full=4096), + is_hybrid_swa=False, + max_total_num_tokens=1024, + full_max_total_num_tokens=None, + swa_max_total_num_tokens=None, + ) + self.assertEqual(instance.max_token_pool_size, 4096) + + def test_hisparse_falls_back_when_size_full_missing(self): + """HiSparse-enabled but allocator has no `size_full` attribute + (e.g. non-HiSparse allocator wired at init time). Fall back to the + SWA-aware effective capacity so we never crash on `AttributeError`.""" + instance = _make_model_runner( + enable_hisparse=True, + token_to_kv_pool_allocator=SimpleNamespace(), # no size_full + is_hybrid_swa=False, + max_total_num_tokens=2048, + full_max_total_num_tokens=None, + swa_max_total_num_tokens=None, + ) + self.assertEqual(instance.max_token_pool_size, 2048) + + def test_non_hisparse_uses_effective_max_total_num_tokens(self): + """Non-HiSparse path is unchanged: delegates to + `effective_max_total_num_tokens` (which returns `max_total_num_tokens` + when SWA is not hybrid).""" + instance = _make_model_runner( + enable_hisparse=False, + token_to_kv_pool_allocator=SimpleNamespace(size_full=99999), # ignored + is_hybrid_swa=False, + max_total_num_tokens=1024, + full_max_total_num_tokens=None, + swa_max_total_num_tokens=None, + ) + self.assertEqual(instance.max_token_pool_size, 1024) + + def test_non_hisparse_hybrid_swa_prefers_full_max(self): + instance = _make_model_runner( + enable_hisparse=False, + token_to_kv_pool_allocator=SimpleNamespace(), + is_hybrid_swa=True, + max_total_num_tokens=1024, + full_max_total_num_tokens=3000, + swa_max_total_num_tokens=500, + ) + self.assertEqual(instance.max_token_pool_size, 3000) + self.assertEqual(instance.effective_max_total_num_tokens, 3000) + + +def _make_prealloc_queue( + *, + enable_hisparse: bool, + max_token_pool_size: int, + max_total_num_tokens: int, +): + """Build a minimal DecodePreallocQueue for _check_if_req_exceed_kv_capacity.""" + queue = DecodePreallocQueue.__new__(DecodePreallocQueue) + queue.max_total_num_tokens = max_total_num_tokens + queue.token_to_kv_pool_allocator = SimpleNamespace(size_swa=10**9) + # Disable the SWA-tail branch; this test only exercises the pool-length gate. + queue._uses_swa_tail_prealloc = MagicMock(return_value=False) + + model_runner = SimpleNamespace(max_token_pool_size=max_token_pool_size) + tp_worker = SimpleNamespace(model_runner=model_runner) + queue.scheduler = SimpleNamespace( + enable_hisparse=enable_hisparse, + tp_worker=tp_worker, + output_streamer=MagicMock(), + ) + return queue + + +def _make_req(rid: str, prompt_len: int): + return SimpleNamespace( + rid=rid, + origin_input_ids=[0] * prompt_len, + output_ids=[], + return_logprob=False, + pd_rebootstrap_in_progress=False, + finished_reason=None, + ) + + +class TestCheckIfReqExceedKvCapacity(CustomTestCase): + def test_hisparse_admits_beyond_device_pool_up_to_host_backed_size(self): + """Core regression: request longer than device-only + `max_total_num_tokens` but within HiSparse host-backed + `max_token_pool_size` must NOT be aborted.""" + queue = _make_prealloc_queue( + enable_hisparse=True, + max_token_pool_size=4096, # host-backed logical capacity + max_total_num_tokens=1024, # device pool + ) + req = _make_req("hisparse-long", prompt_len=2048) + + self.assertFalse(queue._check_if_req_exceed_kv_capacity(req)) + queue.scheduler.output_streamer.stream_output.assert_not_called() + + def test_hisparse_rejects_beyond_host_backed_size(self): + """Requests longer than host-backed capacity are still aborted.""" + queue = _make_prealloc_queue( + enable_hisparse=True, + max_token_pool_size=4096, + max_total_num_tokens=1024, + ) + req = _make_req("hisparse-too-long", prompt_len=5000) + + self.assertTrue(queue._check_if_req_exceed_kv_capacity(req)) + queue.scheduler.output_streamer.stream_output.assert_called_once_with( + [req], req.return_logprob + ) + # prepare_abort sets finished_reason to a BAD_REQUEST FINISH_ABORT. + self.assertIsNotNone(req.finished_reason) + + def test_non_hisparse_uses_device_pool_capacity(self): + """Non-HiSparse path must keep using `max_total_num_tokens` — the + HiSparse branch must not bleed into normal decode admission.""" + queue = _make_prealloc_queue( + enable_hisparse=False, + max_token_pool_size=4096, # ignored on non-HiSparse + max_total_num_tokens=1024, + ) + req = _make_req("non-hisparse-too-long", prompt_len=2048) + + self.assertTrue(queue._check_if_req_exceed_kv_capacity(req)) + queue.scheduler.output_streamer.stream_output.assert_called_once_with( + [req], req.return_logprob + ) + + def test_rebootstrap_input_len_used_for_capacity(self): + """Rebootstrap requests carry both prompt and emitted output_ids; the + admission gate must use the rebootstrap-aware length (prompt + output) + rather than just the prompt length.""" + queue = _make_prealloc_queue( + enable_hisparse=True, + max_token_pool_size=100, + max_total_num_tokens=100, + ) + req = SimpleNamespace( + rid="rebootstrap", + origin_input_ids=[0] * 60, + output_ids=[0] * 60, # 60 + 60 = 120 > 100 + return_logprob=False, + pd_rebootstrap_in_progress=True, + finished_reason=None, + ) + self.assertTrue(queue._check_if_req_exceed_kv_capacity(req)) + + +if __name__ == "__main__": + unittest.main()