config: the KV-cache configurator reads the bags (#34096)

This commit is contained in:
Cheng Wan
2026-08-09 14:45:42 -07:00
committed by GitHub
parent e216c2bc59
commit b4284f3eb7
6 changed files with 99 additions and 61 deletions
@@ -6,14 +6,23 @@ from sglang.srt.runtime_context import get_server_args
from sglang.srt.server_args import ServerArgs
def get_alloc_len_per_decode(server_args: ServerArgs) -> int:
def get_alloc_len_per_decode(
server_args: ServerArgs, *, max_draft_tokens: Optional[int] = None
) -> int:
"""``max_draft_tokens`` lets a caller that already resolved the draft-token
bound (the KV-cache configurator reads it off the bags) size with that same
value; the default is the handed instance's own member, never the global."""
if server_args.speculative_algorithm is None:
return 1
# Spec decoding allocates max(topk * num_steps, num_draft_tokens) per decode step.
spec_steps = server_args.speculative_num_steps or 1
spec_topk = server_args.speculative_eagle_topk or 1
spec_tokens = server_args.max_speculative_num_draft_tokens
spec_tokens = (
max_draft_tokens
if max_draft_tokens is not None
else server_args.max_speculative_num_draft_tokens
)
page_size = server_args.page_size
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
@@ -31,7 +40,11 @@ def get_alloc_len_per_decode(server_args: ServerArgs) -> int:
return max(num_new_pages_per_topk * page_size * spec_topk, spec_tokens)
def get_alloc_reserve_per_decode(server_args: Optional[ServerArgs] = None) -> int:
def get_alloc_reserve_per_decode(
server_args: Optional[ServerArgs] = None,
*,
max_draft_tokens: Optional[int] = None,
) -> int:
"""KV length reserved per request at each decode step.
The 2x is a double-buffer that absorbs the kv_committed_len lag in overlap
@@ -43,23 +56,35 @@ def get_alloc_reserve_per_decode(server_args: Optional[ServerArgs] = None) -> in
"""
if server_args is None:
server_args = get_server_args()
return 2 * get_alloc_len_per_decode(server_args)
return 2 * get_alloc_len_per_decode(server_args, max_draft_tokens=max_draft_tokens)
def get_req_to_token_extra_context_len(server_args: ServerArgs) -> int:
def get_req_to_token_extra_context_len(
server_args: ServerArgs, *, max_draft_tokens: Optional[int] = None
) -> int:
"""req_to_token row headroom beyond the model context length.
Sized to hold the decode over-allocation; the spec v2 page>1 topk>1 holey
draft footprint can outgrow the default num_draft_tokens headroom.
``max_draft_tokens`` keeps this row headroom and the caller's other
draft-token-sized buffers on ONE resolved value: the KV-cache configurator
passes the bag-derived bound it also hands the pools, so the two cannot
disagree after a post-publish override. The default stays the handed
instance's member for callers sizing against a specific config object.
"""
if max_draft_tokens is None:
max_draft_tokens = server_args.max_speculative_num_draft_tokens
# FIXME(lsyin): temporary fix for the context length issue under spec decoding
extra = 4 + (server_args.max_speculative_num_draft_tokens or 0)
extra = 4 + (max_draft_tokens or 0)
if server_args.speculative_algorithm is not None and server_args.page_size > 1:
# kv_allocated_len is page-aligned (eagle_prepare_for_decode), so near
# the context limit the aligned reserve can overshoot by page_size - 1;
# without the headroom the row write silently lands in the neighbor row.
extra = max(
extra,
get_alloc_reserve_per_decode(server_args) + server_args.page_size - 1,
get_alloc_reserve_per_decode(server_args, max_draft_tokens=max_draft_tokens)
+ server_args.page_size
- 1,
)
return extra
@@ -64,6 +64,7 @@ from sglang.srt.mem_cache.memory_pool import (
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import (
configured_pp_size,
get_context,
get_disagg,
get_exec,
@@ -71,6 +72,9 @@ from sglang.srt.runtime_context import (
get_parallel,
get_schedule,
get_spec,
mamba_extra_buffer_enabled,
mamba_extra_buffer_lazy_enabled,
max_speculative_num_draft_tokens,
)
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
@@ -404,7 +408,7 @@ class KVCacheConfigurator:
mamba_spec_state_size=sizes.max_running_requests,
cache_params=self.mambaish_config.mamba2_cache_params,
device=self.device,
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
enable_mamba_extra_buffer=mamba_extra_buffer_enabled(),
draft_model_idx=self.draft_model_idx,
speculative_eagle_topk=get_spec().speculative_eagle_topk,
)
@@ -512,7 +516,7 @@ class KVCacheConfigurator:
max_mamba_cache_size=get_schedule().max_mamba_cache_size,
max_num_reqs=max_num_reqs,
enable_memory_saver=get_exec().features.enable_memory_saver,
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
enable_mamba_extra_buffer=mamba_extra_buffer_enabled(),
speculative_num_draft_tokens=get_spec().speculative_num_draft_tokens,
disable_overlap_schedule=get_schedule().disable_overlap_schedule,
need_sort=get_disagg().disaggregation_mode in ("decode", "prefill"),
@@ -637,7 +641,7 @@ class KVCacheConfigurator:
elif current_platform.is_out_of_tree() and not self.mambaish_config:
unsupported_pool_family = "out-of-tree platform KV pool"
elif (
self.server_args.attention_backend == "ascend" and not self.mambaish_config
get_exec().kernel.attention_backend == "ascend" and not self.mambaish_config
):
unsupported_pool_family = "NPU/Ascend KV pool"
elif self.use_mla_backend and is_dsa_model:
@@ -661,7 +665,13 @@ class KVCacheConfigurator:
)
def _build_req_to_token_pool(self, *, max_num_reqs: int) -> ReqToTokenPool:
extra_max_context_len = get_req_to_token_extra_context_len(self.server_args)
# The same bag-derived bound the pools below receive, so the row
# headroom and the speculative buffers cannot disagree after a
# post-publish override.
extra_max_context_len = get_req_to_token_extra_context_len(
self.server_args,
max_draft_tokens=max_speculative_num_draft_tokens(),
)
if get_disagg().disaggregation_mode == "decode":
# Extra slots for pre-allocated requests
@@ -714,9 +724,9 @@ class KVCacheConfigurator:
if self.layer_info.start_layer <= i < self.layer_info.end_layer
]
),
speculative_num_draft_tokens=self.server_args.max_speculative_num_draft_tokens,
speculative_num_draft_tokens=max_speculative_num_draft_tokens(),
speculative_eagle_topk=get_spec().speculative_eagle_topk,
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
enable_mamba_extra_buffer=mamba_extra_buffer_enabled(),
pre_alloc_size=pre_alloc_size,
enable_overlap_schedule=not get_schedule().disable_overlap_schedule,
mamba_size=get_schedule().max_mamba_cache_size,
@@ -768,7 +778,7 @@ class KVCacheConfigurator:
) -> ReqToTokenPool:
# DSPARK/DFLASH commit routes through the backend fold (KDA-only); a
# non-KDA model there would scatter a None intermediate_ssm and crash.
_algo = (self.server_args.speculative_algorithm or "").upper()
_algo = (get_spec().speculative_algorithm or "").upper()
if (
get_exec().mamba.enable_linear_replayssm_spec
and _algo in ("DSPARK", "DFLASH")
@@ -793,9 +803,9 @@ class KVCacheConfigurator:
if self.layer_info.start_layer <= i < self.layer_info.end_layer
]
),
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
enable_mamba_extra_buffer_lazy=self.server_args.enable_mamba_extra_buffer_lazy(),
speculative_num_draft_tokens=self.server_args.max_speculative_num_draft_tokens,
enable_mamba_extra_buffer=mamba_extra_buffer_enabled(),
enable_mamba_extra_buffer_lazy=mamba_extra_buffer_lazy_enabled(),
speculative_num_draft_tokens=max_speculative_num_draft_tokens(),
speculative_eagle_topk=get_spec().speculative_eagle_topk,
enable_overlap_schedule=not get_schedule().disable_overlap_schedule,
start_layer=self.layer_info.start_layer,
@@ -884,7 +894,7 @@ class KVCacheConfigurator:
max_total_num_tokens=sizes.max_total_num_tokens,
)
elif (
self.server_args.attention_backend == "ascend" and not self.mambaish_config
get_exec().kernel.attention_backend == "ascend" and not self.mambaish_config
):
if self.is_hybrid_swa:
token_to_kv_pool = self._build_ascend_swa_kv_pool(
@@ -1033,9 +1043,7 @@ class KVCacheConfigurator:
start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer,
enable_hisparse=get_memory().enable_hisparse,
online_mtp_max_draft_tokens=(
self.server_args.max_speculative_num_draft_tokens or 0
),
online_mtp_max_draft_tokens=(max_speculative_num_draft_tokens() or 0),
)
return token_to_kv_pool
@@ -1496,7 +1504,7 @@ class KVCacheConfigurator:
need_sort=need_sort,
)
elif _is_npu and (
self.server_args.attention_backend == "ascend"
get_exec().kernel.attention_backend == "ascend"
or is_dsv4_model
or self.hybrid_gdn_config is not None
):
@@ -1686,17 +1694,17 @@ class KVCacheConfigurator:
)
additional_ratio = 0
if self.server_args.enable_mamba_extra_buffer():
if mamba_extra_buffer_enabled():
# ping-pong buffer size is 2 when overlap schedule is on, 1 otherwise.
# Lazy mode saves 1 slot (2 → 1) for overlap; non-overlap already uses 1.
if not get_schedule().disable_overlap_schedule:
if self.server_args.enable_mamba_extra_buffer_lazy():
if mamba_extra_buffer_lazy_enabled():
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP_LAZY
else:
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP
else:
assert (
not self.server_args.enable_mamba_extra_buffer_lazy()
not mamba_extra_buffer_lazy_enabled()
), "Lazy extra buffer requires overlap schedule (--disable-overlap-schedule is incompatible)"
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_NO_OVERLAP
elif skip_decode_lock:
@@ -1725,7 +1733,7 @@ class KVCacheConfigurator:
token_capacity = min(token_capacity, user_limit)
# Sync across PP ranks (each may have different layer counts)
if self.server_args.pp_size > 1:
if configured_pp_size() > 1:
tensor = torch.tensor(token_capacity, dtype=torch.int64)
torch.distributed.all_reduce(
tensor,
@@ -1835,7 +1843,6 @@ class KVCacheConfigurator:
def _handle_max_mamba_cache(self, total_rest_memory):
config = self.mambaish_config
server_args = self.server_args
assert config is not None
# mamba_cache_per_req covers every mamba layer, but under PP a rank only
@@ -1873,10 +1880,11 @@ class KVCacheConfigurator:
if replayssm_active:
# GDN sizes the fold window to the draft maximum; the KDA ring
# stays --linear-replayssm-cache-len long (mirrors MambaPool).
max_draft_tokens = max_speculative_num_draft_tokens()
if kimi_linear_config(self.model_config) is not None:
record_len = get_exec().mamba.linear_replayssm_cache_len
elif server_args.max_speculative_num_draft_tokens is not None:
record_len = server_args.max_speculative_num_draft_tokens
elif max_draft_tokens is not None:
record_len = max_draft_tokens
else:
record_len = get_exec().mamba.linear_replayssm_cache_len
replayssm_ring_per_req = (
+15 -10
View File
@@ -969,12 +969,15 @@ class RuntimeContext:
dummy-boundary ``ServerArgs`` carrying ``fields`` and returns it;
``restore()`` (or exiting) reinstates whatever the slot held before.
Transitional — to be deprecated: it exists because production code
still branches on raw ``server_args`` fields at runtime, so forcing a
path needs a full config in the slot. As those readers migrate onto
the named runtime tiers (flags / resources / forward), prefer the
finer-grained overrides; once they cover the branching surface this
override loses its clients and goes away.
This is the sanctioned way for a test to get a published context, and
it stays. The transitional reason it was introduced for — production
code branching on raw ``server_args`` fields at runtime — is gone (the
read ratchet pins business reads at zero), but a test that exercises
bag readers still needs bags, and the bag tree is projected *from an
instance*: something has to publish one. Prefer the finer-grained
scoped overrides (``get_exec().override(...)``, the flag groups'
``override``) on top of a published context when a test only needs to
force one leaf.
"""
return _ServerArgsOverride(self, fields)
@@ -1194,10 +1197,12 @@ ROLE_NAMESPACE_SETS: dict[str, frozenset[str] | None] = {
# Audited (record-mode smokes, plain + DP-attention): the DP controller
# reads only the elastic-EP gate; its module's static read set agrees.
"dp_controller": frozenset({"exec"}),
# Zero bag reads observed (per-instance managers read self.server_args by
# design). Keep full: it may need namespaces (e.g. disagg) once tokenizer
# paths migrate off self.server_args, and restricting on a zero-read
# audit would be guesswork.
# Record-mode audit (2026-08-06, text model, /generate + /get_server_info +
# /v1/models): reads exactly {"serving"} — the per-instance managers read
# self.server_args by design. Still declared full, because that run did not
# exercise the multimodal processors, LoRA/score endpoints, the disagg
# roles, or the gRPC bridge; narrowing needs those shapes audited too, and
# a wrong set fails a request rather than a test.
"tokenizer": None,
# Deployment shapes not exercised locally; audit before restricting.
"encoder": None,