From bf417440e90eae60878308919794cf57dafeb01c Mon Sep 17 00:00:00 2001 From: AuFlow <73925903+AuFlow@users.noreply.github.com> Date: Fri, 17 Jul 2026 12:34:10 +0800 Subject: [PATCH] [Scheduler] Add `SGLANG_MAX_NEW_TOKENS_LIMIT` to cap per-request `max_new_tokens` (#22591) Co-authored-by: AuFlow --- .../docs/references/environment_variables.mdx | 5 + python/sglang/srt/environ.py | 1 + python/sglang/srt/managers/scheduler.py | 25 ++- .../test_scheduler_init_req_max_new_tokens.py | 176 ++++++++++++++++++ 4 files changed, 202 insertions(+), 5 deletions(-) create mode 100644 test/registered/unit/managers/test_scheduler_init_req_max_new_tokens.py diff --git a/docs_new/docs/references/environment_variables.mdx b/docs_new/docs/references/environment_variables.mdx index 3434b2f73..275c32638 100644 --- a/docs_new/docs/references/environment_variables.mdx +++ b/docs_new/docs/references/environment_variables.mdx @@ -167,6 +167,11 @@ SGLang supports various environment variables that can be used to configure its Set the maximum number of requests per poll, with a negative value indicating no limit `-1` + + SGLANG_MAX_NEW_TOKENS_LIMIT + Hard server-side limit for each generation request's `max_new_tokens`; requests asking for more are clipped. Disabled when unset or non-positive. + Not set + `SGLANG_DATA_PARALLEL_BUDGET_INTERVAL` Interval for DPBudget updates diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 07b3adcb0..d5beef75b 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -345,6 +345,7 @@ class Envs: SGLANG_NEW_TOKEN_RATIO_DECAY_STEPS = EnvInt(600) SGLANG_RETRACT_DECODE_STEPS = EnvInt(20) SGLANG_CLIP_MAX_NEW_TOKENS_ESTIMATION = EnvInt(4096) + SGLANG_MAX_NEW_TOKENS_LIMIT = EnvInt(None) # Scheduler: recv interval SGLANG_SCHEDULER_RECV_SKIPPER_WEIGHT_DEFAULT = EnvInt(1000) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 8d025d71e..db44aad33 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -361,6 +361,7 @@ class Scheduler( and self.enable_hierarchical_cache ) self.max_recv_per_poll = envs.SGLANG_SCHEDULER_MAX_RECV_PER_POLL.get() + self.max_new_tokens_limit = envs.SGLANG_MAX_NEW_TOKENS_LIMIT.get() self.enable_hisparse = server_args.enable_hisparse self.enable_dp_attention = server_args.enable_dp_attention self.enable_unified_memory = server_args.enable_unified_memory @@ -1865,6 +1866,20 @@ class Scheduler( def init_req_max_new_tokens(self, req): input_len = len(req.origin_input_ids) + max_new_tokens = ( + req.sampling_params.max_new_tokens + if req.sampling_params.max_new_tokens is not None + else 1 << 30 + ) + if self.max_new_tokens_limit is not None and self.max_new_tokens_limit > 0: + if max_new_tokens > self.max_new_tokens_limit: + logger.warning( + f"Capping max_new_tokens of request {req.rid} to " + f"SGLANG_MAX_NEW_TOKENS_LIMIT={self.max_new_tokens_limit} " + f"(requested: {req.sampling_params.max_new_tokens})." + ) + max_new_tokens = min(max_new_tokens, self.max_new_tokens_limit) + # Keep this bound consistent with PrefillAdder's admission budget: # ceil_page(input_len) + max_new_tokens + page_size must be strictly # smaller than max_total_num_tokens. Otherwise a request can be accepted @@ -1874,15 +1889,15 @@ class Scheduler( req.sampling_params.max_new_tokens = max( 0, min( - ( - req.sampling_params.max_new_tokens - if req.sampling_params.max_new_tokens is not None - else 1 << 30 - ), + max_new_tokens, self.max_req_len - input_len - 1, self.max_total_num_tokens - paged_input_len - self.page_size - 1, ), ) + # Clipping above can push max_new_tokens below min_new_tokens, which + # would suppress EOS for the whole generation. Restore the invariant. + if req.sampling_params.min_new_tokens > req.sampling_params.max_new_tokens: + req.sampling_params.min_new_tokens = req.sampling_params.max_new_tokens def _process_and_broadcast_mm_inputs( self, diff --git a/test/registered/unit/managers/test_scheduler_init_req_max_new_tokens.py b/test/registered/unit/managers/test_scheduler_init_req_max_new_tokens.py new file mode 100644 index 000000000..c9029f186 --- /dev/null +++ b/test/registered/unit/managers/test_scheduler_init_req_max_new_tokens.py @@ -0,0 +1,176 @@ +import logging +import unittest +from types import SimpleNamespace + +from sglang.srt.environ import envs +from sglang.srt.managers.scheduler import Scheduler +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=1, suite="stage-a-test-cpu") + + +class TestSchedulerInitReqMaxNewTokens(unittest.TestCase): + """Property tests for Scheduler.init_req_max_new_tokens. + + Rules enforced when clipping a request's max_new_tokens: + 1. context: input_len + max_new_tokens < max_req_len + 2. admission budget (PrefillAdder): + ceil_page(input_len) + max_new_tokens + page_size < max_total_num_tokens + 3. env limit: <= SGLANG_MAX_NEW_TOKENS_LIMIT when set and positive + 4. never above the requested value + 5. min_new_tokens <= max_new_tokens afterwards + + Each case asserts all rules hold and the result is tight: one more token + would violate a rule or exceed the request. Over-long inputs degenerate to + max_new_tokens = 0 and are rejected by later admission checks. + """ + + @classmethod + def setUpClass(cls): + # Silence the per-request capping warning; the sweep triggers it a lot. + cls._scheduler_logger = logging.getLogger("sglang.srt.managers.scheduler") + cls._old_level = cls._scheduler_logger.level + cls._scheduler_logger.setLevel(logging.ERROR) + + @classmethod + def tearDownClass(cls): + cls._scheduler_logger.setLevel(cls._old_level) + + def _new_scheduler( + self, + max_req_len: int = 128, + max_total_num_tokens: int = 1024, + page_size: int = 1, + ) -> Scheduler: + scheduler = Scheduler.__new__(Scheduler) + scheduler.max_req_len = max_req_len + scheduler.max_total_num_tokens = max_total_num_tokens + scheduler.page_size = page_size + scheduler.max_new_tokens_limit = envs.SGLANG_MAX_NEW_TOKENS_LIMIT.get() + return scheduler + + def _new_req(self, max_new_tokens, input_len: int = 8, min_new_tokens: int = 0): + return SimpleNamespace( + rid="test-req", + origin_input_ids=[0] * input_len, + sampling_params=SimpleNamespace( + max_new_tokens=max_new_tokens, min_new_tokens=min_new_tokens + ), + ) + + def _init_and_check(self, scheduler, req) -> int: + """Run init_req_max_new_tokens, then assert all admission rules hold + and the result is tight. Returns the resulting max_new_tokens.""" + requested = req.sampling_params.max_new_tokens + scheduler.init_req_max_new_tokens(req) + max_new_tokens = req.sampling_params.max_new_tokens + + input_len = len(req.origin_input_ids) + page_size = scheduler.page_size + paged_input_len = -(-input_len // page_size) * page_size + limit = scheduler.max_new_tokens_limit + limit_active = limit is not None and limit > 0 + + def satisfies_rules(candidate: int) -> bool: + context_ok = input_len + candidate < scheduler.max_req_len + budget_ok = ( + paged_input_len + candidate + page_size < scheduler.max_total_num_tokens + ) + limit_ok = not limit_active or candidate <= limit + requested_ok = requested is None or candidate <= requested + return context_ok and budget_ok and limit_ok and requested_ok + + self.assertGreaterEqual(max_new_tokens, 0) + if max_new_tokens > 0: + self.assertTrue(satisfies_rules(max_new_tokens)) + self.assertFalse(satisfies_rules(max_new_tokens + 1)) + self.assertLessEqual(req.sampling_params.min_new_tokens, max_new_tokens) + return max_new_tokens + + def test_limit_disabled_by_default(self): + with envs.SGLANG_MAX_NEW_TOKENS_LIMIT.override(None): + scheduler = self._new_scheduler() + req = self._new_req(max_new_tokens=64) + self.assertEqual(self._init_and_check(scheduler, req), 64) + + def test_limit_clips_explicit_request(self): + with envs.SGLANG_MAX_NEW_TOKENS_LIMIT.override(16): + scheduler = self._new_scheduler() + req = self._new_req(max_new_tokens=64) + self.assertEqual(self._init_and_check(scheduler, req), 16) + + def test_limit_applies_when_request_unset(self): + with envs.SGLANG_MAX_NEW_TOKENS_LIMIT.override(16): + scheduler = self._new_scheduler() + req = self._new_req(max_new_tokens=None) + self.assertEqual(self._init_and_check(scheduler, req), 16) + + def test_non_positive_limit_is_ignored(self): + for limit in (0, -1): + with self.subTest(limit=limit): + with envs.SGLANG_MAX_NEW_TOKENS_LIMIT.override(limit): + scheduler = self._new_scheduler() + req = self._new_req(max_new_tokens=64) + self.assertEqual(self._init_and_check(scheduler, req), 64) + + def test_context_rule_binds_tighter_than_limit(self): + max_req_len, input_len = 32, 20 + with envs.SGLANG_MAX_NEW_TOKENS_LIMIT.override(16): + scheduler = self._new_scheduler(max_req_len=max_req_len) + req = self._new_req(max_new_tokens=64, input_len=input_len) + self.assertEqual( + self._init_and_check(scheduler, req), max_req_len - input_len - 1 + ) + + def test_budget_rule_binds_tighter_than_limit(self): + max_total_num_tokens, page_size, input_len = 24, 4, 8 + with envs.SGLANG_MAX_NEW_TOKENS_LIMIT.override(32): + scheduler = self._new_scheduler( + max_total_num_tokens=max_total_num_tokens, page_size=page_size + ) + req = self._new_req(max_new_tokens=64, input_len=input_len) + paged_input_len = -(-input_len // page_size) * page_size + self.assertEqual( + self._init_and_check(scheduler, req), + max_total_num_tokens - paged_input_len - page_size - 1, + ) + + def test_min_new_tokens_clamped_to_limit(self): + with envs.SGLANG_MAX_NEW_TOKENS_LIMIT.override(16): + scheduler = self._new_scheduler() + req = self._new_req(max_new_tokens=64, min_new_tokens=32) + self.assertEqual(self._init_and_check(scheduler, req), 16) + self.assertEqual(req.sampling_params.min_new_tokens, 16) + + def test_admission_rules_sweep(self): + for page_size in (1, 4, 16): + for input_len in (1, 8, 100): + for requested in (None, 0, 5, 64, 1 << 20): + for limit in (None, 0, 16, 1 << 20): + for max_req_len, max_total_num_tokens in ( + (128, 1024), + (32, 24), + (128, 24), + ): + with self.subTest( + page_size=page_size, + input_len=input_len, + requested=requested, + limit=limit, + max_req_len=max_req_len, + max_total_num_tokens=max_total_num_tokens, + ): + with envs.SGLANG_MAX_NEW_TOKENS_LIMIT.override(limit): + scheduler = self._new_scheduler( + max_req_len=max_req_len, + max_total_num_tokens=max_total_num_tokens, + page_size=page_size, + ) + req = self._new_req( + max_new_tokens=requested, input_len=input_len + ) + self._init_and_check(scheduler, req) + + +if __name__ == "__main__": + unittest.main()