From 7e3d18bbcc441b0e5eef87759535f2b4d4cf7b71 Mon Sep 17 00:00:00 2001 From: cctry Date: Sat, 12 Sep 2026 21:37:10 -0700 Subject: [PATCH] Add a provider hook for prefill-buffer ceilings (#39182) Co-authored-by: cctry <17473714+cctry@users.noreply.github.com> --- python/sglang/srt/arg_groups/overrides.py | 11 ++++- .../srt/arg_groups/prefill_buffer_ceiling.py | 35 +++++++++++++++ python/sglang/srt/runtime_context.py | 14 +++--- test/registered/unit/test_runtime_context.py | 44 +++++++++++++++++++ 4 files changed, 96 insertions(+), 8 deletions(-) create mode 100644 python/sglang/srt/arg_groups/prefill_buffer_ceiling.py diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index c6730c8bf..6d26836de 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -68,6 +68,7 @@ from sglang.srt.arg_groups.model_override_base import ( # noqa: F401 resolving_view, use_mla_backend, ) +from sglang.srt.arg_groups.prefill_buffer_ceiling import prefill_buffer_ceiling_of logger = logging.getLogger(__name__) from sglang.srt.environ import envs @@ -1873,7 +1874,10 @@ def cutedsl_moe_max_num_tokens(server_args: Any) -> int: def max_prefill_buffer_tokens(server_args: Any) -> int: """Prefill-buffer ceiling: chunked_prefill_size, except PP dynamic - chunking can grow chunks toward max_prefill_tokens and probe at 1.25x.""" + chunking can grow chunks toward max_prefill_tokens and probe at 1.25x. + + Records with a registered ceiling provider (see + ``register_prefill_buffer_ceiling``) answer through it.""" cfg = resolving_view(server_args) chunked = ( cfg.chunked_prefill_size @@ -1883,7 +1887,10 @@ def max_prefill_buffer_tokens(server_args: Any) -> int: tokens = chunked if cfg.enable_dynamic_chunking and cfg.pp_size > 1 and chunked: tokens = max(tokens, cfg.max_prefill_tokens or 0, math.ceil(chunked * 1.25)) - return tokens + record = server_args + if isinstance(server_args, (ResolvedView, ResolvingConfig)): + record = record_of(server_args) + return prefill_buffer_ceiling_of(record, tokens) def mamba_cache_chunk_size(server_args: Any) -> int: diff --git a/python/sglang/srt/arg_groups/prefill_buffer_ceiling.py b/python/sglang/srt/arg_groups/prefill_buffer_ceiling.py new file mode 100644 index 000000000..aab022d9a --- /dev/null +++ b/python/sglang/srt/arg_groups/prefill_buffer_ceiling.py @@ -0,0 +1,35 @@ +"""Dependency-free hook shared by pre- and post-publish prefill-buffer sizing.""" + +from typing import Any, Callable, Optional + +_prefill_buffer_ceiling_fn: Optional[Callable[[Any, int], int]] = None + + +def register_prefill_buffer_ceiling( + fn: Callable[[Any, int], int], +) -> Callable[[Any, int], int]: + """Register one provider; repeating the same registration is harmless. + + The provider receives ``(record, default_ceiling)`` and returns the ceiling, + keeping ``default_ceiling`` for records it does not handle. ``record`` is + the original argument record, never a view: its fields remain raw inputs. + Read resolution declarations through ``resolving_view(record)`` from + ``arg_groups.model_override_base``. The provider must not mutate the record + and must use the same sizing policy before and after publication. + + Registering a different provider raises instead of replacing the first. + """ + global _prefill_buffer_ceiling_fn + if _prefill_buffer_ceiling_fn is not None and _prefill_buffer_ceiling_fn is not fn: + raise RuntimeError( + "A different prefill-buffer ceiling provider is already registered" + ) + _prefill_buffer_ceiling_fn = fn + return fn + + +def prefill_buffer_ceiling_of(record: Any, default_ceiling: int) -> int: + """Return the provider's ceiling, or the default when none is registered.""" + if _prefill_buffer_ceiling_fn is None: + return default_ceiling + return _prefill_buffer_ceiling_fn(record, default_ceiling) diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index 559989ee8..a2838bb89 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -57,6 +57,8 @@ from typing import TYPE_CHECKING, Any, Dict, Optional import msgspec +from sglang.srt.arg_groups.prefill_buffer_ceiling import prefill_buffer_ceiling_of + if TYPE_CHECKING: from sglang.srt.model_executor.runner_utils.pool import GraphPoolBorrowState from sglang.srt.server_args import ServerArgs @@ -1792,13 +1794,13 @@ def max_prefill_buffer_tokens() -> int: """The prefill-buffer ceiling: ``chunked_prefill_size``, except PP dynamic chunking can grow chunks toward ``max_prefill_tokens`` and probe at 1.25x. - Every input is a published leaf (``schedule`` plus the configured PP size), - so this derives from the bags and follows a post-publish override; + The default derives from published leaves (``schedule`` plus the configured + PP size), so it follows post-publish overrides; ``overrides.max_prefill_buffer_tokens`` is the pre-publish equivalent and - ``TestDerivedPredicatesAgreeAcrossTiers`` pins the two equal. + ``TestDerivedPredicatesAgreeAcrossTiers`` pins the two equal. Records with + a registered ceiling provider (see ``register_prefill_buffer_ceiling``) + answer through it. """ - import math - schedule = get_schedule() chunked = ( schedule.chunked_prefill_size @@ -1810,7 +1812,7 @@ def max_prefill_buffer_tokens() -> int: tokens = max( tokens, schedule.max_prefill_tokens or 0, math.ceil(chunked * 1.25) ) - return tokens + return prefill_buffer_ceiling_of(get_server_args(), tokens) def pre_capture_activation_reserve_mb(gpu_mem: float | None) -> float: diff --git a/test/registered/unit/test_runtime_context.py b/test/registered/unit/test_runtime_context.py index 49e319c86..16d3a59bc 100644 --- a/test/registered/unit/test_runtime_context.py +++ b/test/registered/unit/test_runtime_context.py @@ -19,7 +19,9 @@ import msgspec.structs import sglang as _sglang import sglang.srt.server_args as server_args_module +from sglang.srt.arg_groups import prefill_buffer_ceiling from sglang.srt.arg_groups.arg_utils import NS, A, Arg +from sglang.srt.arg_groups.model_override_base import resolving_view from sglang.srt.arg_groups.overrides import ( attention_backends_of, ) @@ -44,7 +46,9 @@ from sglang.srt.runtime_context import ( get_exec, get_flags, get_parallel, + get_schedule, get_server_args, + max_prefill_buffer_tokens, max_speculative_num_draft_tokens, publish, publish_role, @@ -1204,6 +1208,46 @@ class TestDerivedPredicatesAgreeAcrossTiers(_IsolatedServerArgs): max_prefill_buffer_tokens(), ) + def test_prefill_buffer_ceiling_hook_honored_across_tiers(self): + args = _FakeResolvedArgs( + chunked_prefill_size=8192, + enable_dynamic_chunking=True, + pp_size=4, + max_prefill_tokens=16384, + ) + + def provider(record, default_ceiling): + self.assertIs(record, args) + return default_ceiling + 5 + + with patch.object(prefill_buffer_ceiling, "_prefill_buffer_ceiling_fn", None): + register = prefill_buffer_ceiling.register_prefill_buffer_ceiling + self.assertEqual(max_prefill_buffer_tokens_of(args), 16384) + self.assertIs(register(provider), provider) + register(provider) + with self.assertRaisesRegex(RuntimeError, "already registered"): + register(lambda record, default_ceiling: default_ceiling) + for record_or_view in (args, resolving_view(args), resolved_view(args)): + self.assertEqual(max_prefill_buffer_tokens_of(record_or_view), 16389) + get_context().set_server_args(args) + self.assertEqual(max_prefill_buffer_tokens(), 16389) + with get_schedule().override(max_prefill_tokens=32768): + self.assertEqual(max_prefill_buffer_tokens(), 32773) + self.assertEqual(args.max_prefill_tokens, 16384) + + def test_prefill_buffer_ceiling_provider_can_preserve_defaults(self): + args = _FakeResolvedArgs(chunked_prefill_size=4096) + + def provider(record, default_ceiling): + return default_ceiling + + with patch.object(prefill_buffer_ceiling, "_prefill_buffer_ceiling_fn", None): + prefill_buffer_ceiling.register_prefill_buffer_ceiling(provider) + for record_or_view in (args, resolving_view(args), resolved_view(args)): + self.assertEqual(max_prefill_buffer_tokens_of(record_or_view), 4096) + get_context().set_server_args(args) + self.assertEqual(max_prefill_buffer_tokens(), 4096) + def test_activation_reserve_matches_the_member(self): from types import SimpleNamespace