Add a provider hook for prefill-buffer ceilings (#39182)
Co-authored-by: cctry <17473714+cctry@users.noreply.github.com>
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user