[Scheduler] Add SGLANG_MAX_NEW_TOKENS_LIMIT to cap per-request max_new_tokens (#22591)
Co-authored-by: AuFlow <AuFlow@users.noreply.github.com>
This commit is contained in:
@@ -167,6 +167,11 @@ SGLang supports various environment variables that can be used to configure its
|
|||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Set the maximum number of requests per poll, with a negative value indicating no limit</td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Set the maximum number of requests per poll, with a negative value indicating no limit</td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`-1`</td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`-1`</td>
|
||||||
</tr>
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_MAX_NEW_TOKENS_LIMIT</code></td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Hard server-side limit for each generation request's `max_new_tokens`; requests asking for more are clipped. Disabled when unset or non-positive.</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Not set</td>
|
||||||
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`SGLANG_DATA_PARALLEL_BUDGET_INTERVAL`</td>
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`SGLANG_DATA_PARALLEL_BUDGET_INTERVAL`</td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Interval for DPBudget updates</td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Interval for DPBudget updates</td>
|
||||||
|
|||||||
@@ -345,6 +345,7 @@ class Envs:
|
|||||||
SGLANG_NEW_TOKEN_RATIO_DECAY_STEPS = EnvInt(600)
|
SGLANG_NEW_TOKEN_RATIO_DECAY_STEPS = EnvInt(600)
|
||||||
SGLANG_RETRACT_DECODE_STEPS = EnvInt(20)
|
SGLANG_RETRACT_DECODE_STEPS = EnvInt(20)
|
||||||
SGLANG_CLIP_MAX_NEW_TOKENS_ESTIMATION = EnvInt(4096)
|
SGLANG_CLIP_MAX_NEW_TOKENS_ESTIMATION = EnvInt(4096)
|
||||||
|
SGLANG_MAX_NEW_TOKENS_LIMIT = EnvInt(None)
|
||||||
|
|
||||||
# Scheduler: recv interval
|
# Scheduler: recv interval
|
||||||
SGLANG_SCHEDULER_RECV_SKIPPER_WEIGHT_DEFAULT = EnvInt(1000)
|
SGLANG_SCHEDULER_RECV_SKIPPER_WEIGHT_DEFAULT = EnvInt(1000)
|
||||||
|
|||||||
@@ -361,6 +361,7 @@ class Scheduler(
|
|||||||
and self.enable_hierarchical_cache
|
and self.enable_hierarchical_cache
|
||||||
)
|
)
|
||||||
self.max_recv_per_poll = envs.SGLANG_SCHEDULER_MAX_RECV_PER_POLL.get()
|
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_hisparse = server_args.enable_hisparse
|
||||||
self.enable_dp_attention = server_args.enable_dp_attention
|
self.enable_dp_attention = server_args.enable_dp_attention
|
||||||
self.enable_unified_memory = server_args.enable_unified_memory
|
self.enable_unified_memory = server_args.enable_unified_memory
|
||||||
@@ -1865,6 +1866,20 @@ class Scheduler(
|
|||||||
|
|
||||||
def init_req_max_new_tokens(self, req):
|
def init_req_max_new_tokens(self, req):
|
||||||
input_len = len(req.origin_input_ids)
|
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:
|
# Keep this bound consistent with PrefillAdder's admission budget:
|
||||||
# ceil_page(input_len) + max_new_tokens + page_size must be strictly
|
# ceil_page(input_len) + max_new_tokens + page_size must be strictly
|
||||||
# smaller than max_total_num_tokens. Otherwise a request can be accepted
|
# 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(
|
req.sampling_params.max_new_tokens = max(
|
||||||
0,
|
0,
|
||||||
min(
|
min(
|
||||||
(
|
max_new_tokens,
|
||||||
req.sampling_params.max_new_tokens
|
|
||||||
if req.sampling_params.max_new_tokens is not None
|
|
||||||
else 1 << 30
|
|
||||||
),
|
|
||||||
self.max_req_len - input_len - 1,
|
self.max_req_len - input_len - 1,
|
||||||
self.max_total_num_tokens - paged_input_len - self.page_size - 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(
|
def _process_and_broadcast_mm_inputs(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user