config: record resolution writes in a declaration stash (#35905)

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-08-23 01:17:27 -07:00
committed by GitHub
co-authored by Claude Opus 5
parent 6218d6ce3f
commit 0e22777572
12 changed files with 1449 additions and 286 deletions
@@ -119,6 +119,14 @@ def namespace_of(cls) -> dict:
return out
@functools.lru_cache(maxsize=None)
def field_names(cls) -> frozenset:
"""Names of ``cls`` dataclass fields — what a declaration may name."""
if not dataclasses.is_dataclass(cls):
return frozenset()
return frozenset(field.name for field in dataclasses.fields(cls))
@functools.lru_cache(maxsize=None)
def resolvable_fields(cls) -> frozenset:
"""Names of ``cls`` dataclass fields whose ``Arg`` metadata declares
@@ -3,6 +3,7 @@ from __future__ import annotations
import logging
from typing import TYPE_CHECKING
from sglang.srt.arg_groups.overrides import declare_resolution
from sglang.srt.environ import envs
if TYPE_CHECKING:
@@ -136,7 +137,11 @@ def apply_deepseek_v4_defaults(server_args: ServerArgs, model_arch: str) -> None
run_post_process_pass(server_args, _deepseek_v4_kv_cache_dtype)
if server_args.max_running_requests is None:
server_args.max_running_requests = 256
declare_resolution(
server_args,
"apply_deepseek_v4_defaults",
max_running_requests=256,
)
logger.warning(
f"Setting max_running_requests to {server_args.max_running_requests} for {model_arch}."
)
@@ -163,12 +168,36 @@ def validate_deepseek_v4_cp(server_args: ServerArgs) -> None:
f"got {server_args.cp_strategy}"
)
server_args.enable_dsa_prefill_context_parallel = True
server_args.enable_prefill_context_parallel = False
server_args.dsa_prefill_cp_mode = "round-robin-split"
server_args.enable_dp_attention = True
server_args.moe_dense_tp_size = 1
server_args.attn_cp_size = server_args.tp_size // server_args.dp_size
declare_resolution(
server_args,
"validate_deepseek_v4_cp",
enable_dsa_prefill_context_parallel=True,
)
declare_resolution(
server_args,
"validate_deepseek_v4_cp",
enable_prefill_context_parallel=False,
)
declare_resolution(
server_args,
"validate_deepseek_v4_cp",
dsa_prefill_cp_mode="round-robin-split",
)
declare_resolution(
server_args,
"validate_deepseek_v4_cp",
enable_dp_attention=True,
)
declare_resolution(
server_args,
"validate_deepseek_v4_cp",
moe_dense_tp_size=1,
)
declare_resolution(
server_args,
"validate_deepseek_v4_cp",
attn_cp_size=server_args.tp_size // server_args.dp_size,
)
assert (
server_args.dp_size == 1
), "For round-robin split mode, dp attention is not supported."
+22 -4
View File
@@ -3,6 +3,8 @@ from __future__ import annotations
import logging
from typing import TYPE_CHECKING
from sglang.srt.arg_groups.overrides import declare_resolution
if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs
@@ -20,7 +22,11 @@ def apply_kimi_k3_spec_backend_defaults(server_args: ServerArgs) -> None:
# Decode is left free (its bf16-ssm SM100+ flashinfer default is fine -- the
# target only verifies under spec); the verify backend is pinned directly.
if server_args.linear_attn_verify_backend is None:
server_args.linear_attn_verify_backend = "nv_cutedsl"
declare_resolution(
server_args,
"apply_kimi_k3_spec_backend_defaults",
linear_attn_verify_backend="nv_cutedsl",
)
logger.info(
"Kimi hybrid model with speculative decoding: pinning "
"--linear-attn-verify-backend to nv_cutedsl (uses the fused "
@@ -34,7 +40,11 @@ def apply_kimi_k3_spec_backend_defaults(server_args: ServerArgs) -> None:
and server_args.speculative_draft_attention_backend is None
and is_sm100_supported()
):
server_args.speculative_draft_attention_backend = "trtllm_mha"
declare_resolution(
server_args,
"apply_kimi_k3_spec_backend_defaults",
speculative_draft_attention_backend="trtllm_mha",
)
logger.info(
"Kimi hybrid DSPARK: defaulting "
"--speculative-draft-attention-backend to trtllm_mha."
@@ -72,7 +82,11 @@ def disable_kimi_k3_symm_mem(server_args: ServerArgs) -> None:
and graph.prefill.backend == Backend.DISABLED
):
return
server_args.enable_symm_mem = False
declare_resolution(
server_args,
"disable_kimi_k3_symm_mem",
enable_symm_mem=False,
)
logger.warning(
"Kimi hybrid model: ignoring --enable-symm-mem because CUDA graphs are on. "
"The symmetric-memory pool's per-forward allocations are not valid for the "
@@ -96,7 +110,11 @@ def apply_kimi_k3_linear_attn_defaults(server_args: ServerArgs) -> None:
and server_args.mamba_ssm_dtype == "bfloat16"
and is_sm100_supported()
):
server_args.linear_attn_decode_backend = "triton"
declare_resolution(
server_args,
"apply_kimi_k3_linear_attn_defaults",
linear_attn_decode_backend="triton",
)
logger.info(
"Kimi hybrid model with bf16 SSM state: defaulting "
"--linear-attn-decode-backend to triton."
@@ -7,6 +7,8 @@ from typing import TYPE_CHECKING
if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs
from sglang.srt.arg_groups.overrides import declare_resolution
logger = logging.getLogger(__name__)
@@ -26,8 +28,12 @@ def handle_moe_runner_backend_alias(server_args: ServerArgs) -> None:
"--moe-a2a-backend %s.",
server_args.moe_a2a_backend,
)
server_args.moe_runner_backend = "auto"
server_args.moe_a2a_backend = "megamoe"
declare_resolution(
server_args,
"handle_moe_runner_backend_alias",
moe_runner_backend="auto",
moe_a2a_backend="megamoe",
)
def handle_w4a4_mxfp4_megamoe_env(server_args: ServerArgs) -> None:
+45 -4
View File
@@ -35,7 +35,7 @@ import json
import logging
from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple
from sglang.srt.arg_groups.arg_utils import resolvable_fields
from sglang.srt.arg_groups.arg_utils import field_names, resolvable_fields
from sglang.srt.environ import envs
from sglang.srt.hardware_backend.mlx.runtime import use_mlx
from sglang.srt.model_executor.cuda_graph_config import Backend
@@ -167,9 +167,12 @@ def register_post_process(fn: Callable[..., dict]) -> Callable[..., dict]:
def _declaration_overlay(server_args: Any) -> Dict[str, Any]:
"""Accumulated declared values: declarations never mutate
``server_args``, so mid-resolution readers overlay them from the
declaration stash (last writer wins, like the gate)."""
"""What the declarations say so far, last writer wins.
Passes declare without touching the fields until
``materialize_declarations``, so a mid-resolution reader needs this to see
them; handlers and hooks write as they declare, and for those the overlay
repeats what the field already holds."""
overlay: Dict[str, Any] = {}
for _source, declared in getattr(server_args, "_resolved_overrides", None) or ():
overlay.update(declared)
@@ -219,6 +222,44 @@ def _apply_fields(server_args: Any, fields: Dict[str, Any]) -> None:
object.__setattr__(server_args, "_internal_write", False)
def declare_resolution(server_args: Any, source: str, **fields: Any) -> None:
"""Record a resolution write in the declaration stash, and apply it now.
The stash is what the projection reads, so a resolver that only assigns
the field leaves that write invisible to it. The immediate write keeps the
resolver's successors seeing the value where they read the field directly.
What it does change is which writer wins. A declaration is appended and
replayed last, so a resolver that declares a field a *deferred* writer (a
post-process pass, a registry entry) also decides now beats it, where its
bare assignment used to be overwritten by that writer's declaration. A
resolver that gates on such a field has to read the resolving view rather
than the raw field, or it decides from a value that is already stale.
For resolvers inside ``__post_init__``: the handlers on ``ServerArgs``
(through ``self._declare``) and the ``arg_groups`` hooks and hardware
defaults they call. Resolution that has to wait for the launcher stage
goes through ``declare_late_resolution`` instead.
Names arrive as keyword arguments, which accept anything; a misspelled one
would otherwise become a new attribute that nothing ever reads, so it is
rejected here. This is not the model-override whitelist: that one limits
which fields a *registry entry* may reach, while a resolver writing the
field it owns is the pipeline resolving by construction.
"""
if dataclasses.is_dataclass(type(server_args)):
unknown = sorted(set(fields) - field_names(type(server_args)))
if unknown:
raise AttributeError(f"{source}: {unknown} are not ServerArgs fields")
stash = getattr(server_args, "_resolved_overrides", None)
if stash is None:
stash = []
object.__setattr__(server_args, "_resolved_overrides", stash)
stash.append((source, dict(fields)))
for name, value in fields.items():
setattr(server_args, name, value)
def declare_late_resolution(server_args: Any, source: str, **fields: Any) -> None:
"""Resolve fields on a config that is **not published yet**.
@@ -5,6 +5,7 @@ import logging
import os
from typing import TYPE_CHECKING
from sglang.srt.arg_groups.overrides import declare_resolution
from sglang.srt.environ import envs
if TYPE_CHECKING:
@@ -20,8 +21,16 @@ def handle_pd_disaggregation(server_args: ServerArgs) -> None:
# mooncake, and skip RDMA HCA selection. Must run before backend-name checks.
if server_args.disaggregation_transfer_backend == "mooncake_tcp":
os.environ.setdefault("MC_FORCE_TCP", "1")
server_args.disaggregation_transfer_backend = "mooncake"
server_args.disaggregation_ib_device = None
declare_resolution(
server_args,
"handle_pd_disaggregation",
disaggregation_transfer_backend="mooncake",
)
declare_resolution(
server_args,
"handle_pd_disaggregation",
disaggregation_ib_device=None,
)
logger.info(
"disaggregation transfer backend 'mooncake_tcp' -> mooncake "
"with MC_FORCE_TCP=1 (TCP transport, no RDMA)"
@@ -83,10 +92,18 @@ def handle_pd_disaggregation(server_args: ServerArgs) -> None:
"EXPERIMENTAL: Decode radix cache with DP attention. "
"Requires prefix-aware DP rank routing for optimal cache hits."
)
server_args.disable_radix_cache = False
declare_resolution(
server_args,
"handle_pd_disaggregation",
disable_radix_cache=False,
)
logger.warning("EXPERIMENTAL: Radix cache is enabled for decode server")
else:
server_args.disable_radix_cache = True
declare_resolution(
server_args,
"handle_pd_disaggregation",
disable_radix_cache=True,
)
logger.warning("KV cache is forced as chunk cache for decode server")
# Default the number of *extra* decode req_to_token slots reserved for
@@ -101,7 +118,11 @@ def handle_pd_disaggregation(server_args: ServerArgs) -> None:
)
if per_worker <= 32:
extra_slots = per_worker * 2
server_args.disaggregation_decode_extra_slots = extra_slots
declare_resolution(
server_args,
"handle_pd_disaggregation",
disaggregation_decode_extra_slots=extra_slots,
)
elif server_args.disaggregation_mode == "prefill":
assert (
@@ -153,4 +174,8 @@ def _alias_bootstrap_port_to_api_port(server_args: ServerArgs) -> None:
server_args.disaggregation_bootstrap_port,
server_args.port,
)
server_args.disaggregation_bootstrap_port = server_args.port
declare_resolution(
server_args,
"_alias_bootstrap_port_to_api_port",
disaggregation_bootstrap_port=server_args.port,
)
+210 -51
View File
@@ -5,6 +5,8 @@ import logging
import os
from typing import TYPE_CHECKING, Optional
from sglang.srt.arg_groups.overrides import declare_resolution
if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs
@@ -15,7 +17,11 @@ def _disable_overlap_schedule_for_cpu(server_args: ServerArgs) -> None:
if server_args.device != "cpu" or server_args.disable_overlap_schedule:
return
server_args.disable_overlap_schedule = True
declare_resolution(
server_args,
"_disable_overlap_schedule_for_cpu",
disable_overlap_schedule=True,
)
logger.warning(
"Overlap schedule is not implemented for speculative decoding on CPU."
)
@@ -66,7 +72,11 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None:
server_args.speculative_draft_model_path is not None
and server_args.speculative_draft_model_revision is None
):
server_args.speculative_draft_model_revision = "main"
declare_resolution(
server_args,
"handle_speculative_decoding",
speculative_draft_model_revision="main",
)
# Moved to the resolution pipeline (arg_groups/overrides.py:
# _speculative_moe_runner_default), invoked here at its legacy slot.
@@ -78,7 +88,11 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None:
run_post_process_pass(server_args, _speculative_moe_runner_default)
if server_args.speculative_algorithm is not None:
server_args.speculative_algorithm = server_args.speculative_algorithm.upper()
declare_resolution(
server_args,
"handle_speculative_decoding",
speculative_algorithm=server_args.speculative_algorithm.upper(),
)
# Removal notice for the retired env var; raw os.getenv on purpose -- the
# Envs descriptor is gone. Drop this check after one release.
@@ -95,11 +109,15 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None:
if override_config_file and override_config_file.strip():
kwargs["_configuration_file"] = override_config_file.strip()
server_args.speculative_algorithm = _resolve_speculative_algorithm_alias(
server_args.speculative_algorithm,
server_args.speculative_draft_model_path,
trust_remote_code=server_args.trust_remote_code,
kwargs=kwargs,
declare_resolution(
server_args,
"handle_speculative_decoding",
speculative_algorithm=_resolve_speculative_algorithm_alias(
server_args.speculative_algorithm,
server_args.speculative_draft_model_path,
trust_remote_code=server_args.trust_remote_code,
kwargs=kwargs,
),
)
# Validate --speculative-draft-window-size once, regardless of algorithm.
@@ -110,7 +128,11 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None:
raise ValueError(
f"--speculative-draft-window-size must be positive, got {window_size}."
)
server_args.speculative_draft_window_size = window_size
declare_resolution(
server_args,
"handle_speculative_decoding",
speculative_draft_window_size=window_size,
)
if server_args.speculative_algorithm not in ("EAGLE3", "DFLASH"):
logger.warning(
"--speculative-draft-window-size has no effect with "
@@ -173,22 +195,38 @@ def _handle_dflash(server_args: ServerArgs) -> None:
#
# For DFlash, the natural unit is `block_size` (verify window length).
if server_args.speculative_num_steps is None:
server_args.speculative_num_steps = 1
declare_resolution(
server_args,
"_handle_dflash",
speculative_num_steps=1,
)
elif int(server_args.speculative_num_steps) != 1:
logger.warning(
"DFLASH only supports speculative_num_steps == 1; overriding speculative_num_steps=%s to 1.",
server_args.speculative_num_steps,
)
server_args.speculative_num_steps = 1
declare_resolution(
server_args,
"_handle_dflash",
speculative_num_steps=1,
)
if server_args.speculative_eagle_topk is None:
server_args.speculative_eagle_topk = 1
declare_resolution(
server_args,
"_handle_dflash",
speculative_eagle_topk=1,
)
elif int(server_args.speculative_eagle_topk) != 1:
logger.warning(
"DFLASH only supports speculative_eagle_topk == 1; overriding speculative_eagle_topk=%s to 1.",
server_args.speculative_eagle_topk,
)
server_args.speculative_eagle_topk = 1
declare_resolution(
server_args,
"_handle_dflash",
speculative_eagle_topk=1,
)
if server_args.speculative_dflash_block_size is not None:
if int(server_args.speculative_dflash_block_size) <= 0:
@@ -205,8 +243,10 @@ def _handle_dflash(server_args: ServerArgs) -> None:
f"speculative_num_draft_tokens={server_args.speculative_num_draft_tokens}, "
f"speculative_dflash_block_size={server_args.speculative_dflash_block_size}."
)
server_args.speculative_num_draft_tokens = int(
server_args.speculative_dflash_block_size
declare_resolution(
server_args,
"_handle_dflash",
speculative_num_draft_tokens=int(server_args.speculative_dflash_block_size),
)
if server_args.speculative_num_draft_tokens is None:
@@ -241,7 +281,11 @@ def _handle_dflash(server_args: ServerArgs) -> None:
"speculative_num_draft_tokens is not set; defaulting to %d for DFLASH.",
inferred_block_size,
)
server_args.speculative_num_draft_tokens = inferred_block_size
declare_resolution(
server_args,
"_handle_dflash",
speculative_num_draft_tokens=inferred_block_size,
)
if server_args.speculative_draft_window_size is not None:
draft_tokens = int(server_args.speculative_num_draft_tokens)
@@ -255,13 +299,21 @@ def _handle_dflash(server_args: ServerArgs) -> None:
_resolve_dflash_draft_attention_backend(server_args)
if server_args.max_running_requests is None:
server_args.max_running_requests = 48
declare_resolution(
server_args,
"_handle_dflash",
max_running_requests=48,
)
logger.warning(
"Max running requests is reset to 48 for speculative decoding. You can override this by explicitly setting --max-running-requests."
)
if server_args.enable_mixed_chunk:
server_args.enable_mixed_chunk = False
declare_resolution(
server_args,
"_handle_dflash",
enable_mixed_chunk=False,
)
logger.warning(
"Mixed chunked prefill is disabled because of using dflash speculative decoding."
)
@@ -327,8 +379,16 @@ def _handle_dspark(server_args: ServerArgs) -> None:
if server_args.speculative_draft_model_path is None:
if _target_checkpoint_bundles_dspark_draft(server_args):
server_args.speculative_draft_model_path = server_args.model_path
server_args.speculative_draft_model_revision = server_args.revision
declare_resolution(
server_args,
"_handle_dspark",
speculative_draft_model_path=server_args.model_path,
)
declare_resolution(
server_args,
"_handle_dspark",
speculative_draft_model_revision=server_args.revision,
)
logger.info(
"DSpark draft weights are bundled in the target checkpoint; "
"defaulting --speculative-draft-model-path to --model-path (%s).",
@@ -341,22 +401,38 @@ def _handle_dspark(server_args: ServerArgs) -> None:
)
if server_args.speculative_num_steps is None:
server_args.speculative_num_steps = 1
declare_resolution(
server_args,
"_handle_dspark",
speculative_num_steps=1,
)
elif int(server_args.speculative_num_steps) != 1:
logger.warning(
"DSpark only supports speculative_num_steps == 1; overriding speculative_num_steps=%s to 1.",
server_args.speculative_num_steps,
)
server_args.speculative_num_steps = 1
declare_resolution(
server_args,
"_handle_dspark",
speculative_num_steps=1,
)
if server_args.speculative_eagle_topk is None:
server_args.speculative_eagle_topk = 1
declare_resolution(
server_args,
"_handle_dspark",
speculative_eagle_topk=1,
)
elif int(server_args.speculative_eagle_topk) != 1:
logger.warning(
"DSpark only supports speculative_eagle_topk == 1; overriding speculative_eagle_topk=%s to 1.",
server_args.speculative_eagle_topk,
)
server_args.speculative_eagle_topk = 1
declare_resolution(
server_args,
"_handle_dspark",
speculative_eagle_topk=1,
)
gamma: Optional[int] = None
if server_args.speculative_dspark_block_size is not None:
@@ -398,7 +474,11 @@ def _handle_dspark(server_args: ServerArgs) -> None:
f"(= {verify_window} for gamma={gamma}), but got "
f"speculative_num_draft_tokens={server_args.speculative_num_draft_tokens}."
)
server_args.speculative_num_draft_tokens = verify_window
declare_resolution(
server_args,
"_handle_dspark",
speculative_num_draft_tokens=verify_window,
)
if server_args.speculative_num_draft_tokens is None:
raise ValueError(
@@ -412,13 +492,21 @@ def _handle_dspark(server_args: ServerArgs) -> None:
)
if server_args.max_running_requests is None:
server_args.max_running_requests = 48
declare_resolution(
server_args,
"_handle_dspark",
max_running_requests=48,
)
logger.warning(
"Max running requests is reset to 48 for speculative decoding. You can override this by explicitly setting --max-running-requests."
)
if server_args.enable_mixed_chunk:
server_args.enable_mixed_chunk = False
declare_resolution(
server_args,
"_handle_dspark",
enable_mixed_chunk=False,
)
logger.warning(
"Mixed chunked prefill is disabled because of using dspark speculative decoding."
)
@@ -522,18 +610,30 @@ def _resolve_dflash_draft_attention_backend(server_args: ServerArgs) -> None:
draft_backend = fallback_backend
# FIXME: avoid overriding server args directly; pass the resolved draft
# backend to the draft worker explicitly instead.
server_args.speculative_draft_attention_backend = draft_backend
declare_resolution(
server_args,
"_resolve_dflash_draft_attention_backend",
speculative_draft_attention_backend=draft_backend,
)
def _handle_frozen_kv_mtp(server_args: ServerArgs) -> None:
if server_args.max_running_requests is None:
server_args.max_running_requests = 48
declare_resolution(
server_args,
"_handle_frozen_kv_mtp",
max_running_requests=48,
)
logger.warning(
"Max running requests is reset to 48 for speculative decoding. You can override this by explicitly setting --max-running-requests."
)
if server_args.enable_mixed_chunk:
server_args.enable_mixed_chunk = False
declare_resolution(
server_args,
"_handle_frozen_kv_mtp",
enable_mixed_chunk=False,
)
logger.warning(
"Mixed chunked prefill is disabled because of using "
"Frozen-KV MTP speculative decoding."
@@ -556,7 +656,11 @@ def _handle_eagle_family(server_args: ServerArgs) -> None:
)
if server_args.max_running_requests is None:
server_args.max_running_requests = 48
declare_resolution(
server_args,
"_handle_eagle_family",
max_running_requests=48,
)
logger.warning(
"Max running requests is reset to 48 for speculative decoding. You can override this by explicitly setting --max-running-requests."
)
@@ -570,7 +674,11 @@ def _handle_eagle_family(server_args: ServerArgs) -> None:
)
if server_args.enable_mixed_chunk:
server_args.enable_mixed_chunk = False
declare_resolution(
server_args,
"_handle_eagle_family",
enable_mixed_chunk=False,
)
logger.warning(
"Mixed chunked prefill is disabled because of using "
"eagle speculative decoding."
@@ -592,8 +700,16 @@ def _handle_eagle_family(server_args: ServerArgs) -> None:
"HYV3ForCausalLM",
]:
if server_args.speculative_draft_model_path is None:
server_args.speculative_draft_model_path = server_args.model_path
server_args.speculative_draft_model_revision = server_args.revision
declare_resolution(
server_args,
"_handle_eagle_family",
speculative_draft_model_path=server_args.model_path,
)
declare_resolution(
server_args,
"_handle_eagle_family",
speculative_draft_model_revision=server_args.revision,
)
else:
if model_arch not in [
"MistralLarge3ForCausalLM",
@@ -612,11 +728,16 @@ def _handle_eagle_family(server_args: ServerArgs) -> None:
and server_args.speculative_num_draft_tokens is None
)
(
server_args.speculative_num_steps,
server_args.speculative_eagle_topk,
server_args.speculative_num_draft_tokens,
) = _auto_choose_speculative_params(server_args, model_arch)
steps, topk, draft_tokens = _auto_choose_speculative_params(
server_args, model_arch
)
declare_resolution(
server_args,
"_handle_eagle_family.auto_params",
speculative_num_steps=steps,
speculative_eagle_topk=topk,
speculative_num_draft_tokens=draft_tokens,
)
if "trtllm_mha" in attention_backends_of(resolved_view(server_args)):
if server_args.speculative_eagle_topk > 1:
@@ -680,7 +801,11 @@ def _handle_eagle_family(server_args: ServerArgs) -> None:
logger.warning(
"speculative_num_draft_tokens is adjusted to speculative_num_steps + 1 when speculative_eagle_topk == 1"
)
server_args.speculative_num_draft_tokens = server_args.speculative_num_steps + 1
declare_resolution(
server_args,
"_handle_eagle_family",
speculative_num_draft_tokens=server_args.speculative_num_steps + 1,
)
# topk > 1 + page_size > 1 needs the two-pass cascade draft-decode (shared prefix
# pass + per-branch expand pass with prefix-tail dup). Only these backends implement
@@ -708,23 +833,41 @@ def _handle_ngram(server_args: ServerArgs) -> None:
_disable_overlap_schedule_for_cpu(server_args)
if server_args.max_running_requests is None:
server_args.max_running_requests = 48
declare_resolution(
server_args,
"_handle_ngram",
max_running_requests=48,
)
logger.warning(
"Max running requests is reset to 48 for speculative decoding. You can override this by explicitly setting --max-running-requests."
)
server_args.enable_mixed_chunk = False
server_args.speculative_eagle_topk = server_args.speculative_ngram_max_bfs_breadth
declare_resolution(
server_args,
"_handle_ngram",
enable_mixed_chunk=False,
)
declare_resolution(
server_args,
"_handle_ngram",
speculative_eagle_topk=server_args.speculative_ngram_max_bfs_breadth,
)
if server_args.speculative_num_draft_tokens is None:
server_args.speculative_num_draft_tokens = 12
declare_resolution(
server_args,
"_handle_ngram",
speculative_num_draft_tokens=12,
)
logger.warning(
"speculative_num_draft_tokens is set to 12 by default for ngram speculative decoding. "
"You can override this by explicitly setting --speculative-num-draft-tokens."
)
if server_args.speculative_num_steps is None:
server_args.speculative_num_steps = (
server_args.speculative_num_draft_tokens
// server_args.speculative_eagle_topk
declare_resolution(
server_args,
"_handle_ngram",
speculative_num_steps=server_args.speculative_num_draft_tokens
// server_args.speculative_eagle_topk,
)
if server_args.speculative_ngram_external_corpus_path is not None:
if server_args.speculative_ngram_external_sam_budget <= 0:
@@ -782,7 +925,11 @@ def _maybe_disable_adaptive(server_args: ServerArgs) -> None:
f"speculative_adaptive disabled: {reason}. "
"Falling back to static speculative params."
)
server_args.speculative_adaptive = False
declare_resolution(
server_args,
"_maybe_disable_adaptive",
speculative_adaptive=False,
)
def _init_adaptive_speculative_params(server_args: ServerArgs) -> None:
@@ -795,10 +942,18 @@ def _init_adaptive_speculative_params(server_args: ServerArgs) -> None:
)
if server_args.speculative_eagle_topk is None:
server_args.speculative_eagle_topk = 1
declare_resolution(
server_args,
"_init_adaptive_speculative_params",
speculative_eagle_topk=1,
)
if server_args.speculative_num_steps is None:
server_args.speculative_num_steps = candidate_steps[len(candidate_steps) // 2]
declare_resolution(
server_args,
"_init_adaptive_speculative_params",
speculative_num_steps=candidate_steps[len(candidate_steps) // 2],
)
if server_args.speculative_num_steps not in candidate_steps:
raise ValueError(
@@ -807,7 +962,11 @@ def _init_adaptive_speculative_params(server_args: ServerArgs) -> None:
"Pass one of those values."
)
server_args.speculative_num_draft_tokens = server_args.speculative_num_steps + 1
declare_resolution(
server_args,
"_init_adaptive_speculative_params",
speculative_num_draft_tokens=server_args.speculative_num_steps + 1,
)
def _auto_choose_speculative_params(server_args: ServerArgs, model_arch: str) -> tuple:
+51 -10
View File
@@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, Callable
import torch
from sglang.srt.arg_groups.overrides import declare_resolution
from sglang.srt.environ import envs
from sglang.srt.utils import get_npu_memory_capacity, is_npu
@@ -44,11 +45,27 @@ def set_default_server_args(args: "ServerArgs"):
"""
# NPU only works with "ascend" attention backend for now
args.attention_backend = "ascend"
args.prefill_attention_backend = "ascend"
args.decode_attention_backend = "ascend"
declare_resolution(
args,
"set_default_server_args",
attention_backend="ascend",
)
declare_resolution(
args,
"set_default_server_args",
prefill_attention_backend="ascend",
)
declare_resolution(
args,
"set_default_server_args",
decode_attention_backend="ascend",
)
if args.page_size is None:
args.page_size = 128
declare_resolution(
args,
"set_default_server_args",
page_size=128,
)
# NPU memory settings
decode = args.cuda_graph_config.decode
@@ -57,7 +74,11 @@ def set_default_server_args(args: "ServerArgs"):
# Ascend 910B4,910B4_1
# (chunked_prefill_size 4k, max_bs 16 if tp < 4 else 64)
if args.chunked_prefill_size is None:
args.chunked_prefill_size = 4 * 1024
declare_resolution(
args,
"set_default_server_args",
chunked_prefill_size=4 * 1024,
)
if decode.max_bs is None:
if args.tp_size < 4:
decode.max_bs = 16
@@ -67,7 +88,11 @@ def set_default_server_args(args: "ServerArgs"):
# Ascend 910B1,910B2,910B2C,910B3,910_9391,910_9392,910_9381,910_9382,910_9372,910_9362
# (chunked_prefill_size 8k, max_bs 64 if tp < 4 else 256)
if args.chunked_prefill_size is None:
args.chunked_prefill_size = 8 * 1024
declare_resolution(
args,
"set_default_server_args",
chunked_prefill_size=8 * 1024,
)
if decode.max_bs is None:
if args.tp_size < 4:
decode.max_bs = 64
@@ -75,15 +100,31 @@ def set_default_server_args(args: "ServerArgs"):
decode.max_bs = 256
# NPU does not support CustomAllReduce
args.disable_custom_all_reduce = True
declare_resolution(
args,
"set_default_server_args",
disable_custom_all_reduce=True,
)
# handles hierarchical cache configs
if args.enable_hierarchical_cache:
args.hicache_io_backend = "kernel_ascend"
declare_resolution(
args,
"set_default_server_args",
hicache_io_backend="kernel_ascend",
)
if args.use_mla_backend():
args.hicache_mem_layout = "page_first_kv_split"
declare_resolution(
args,
"set_default_server_args",
hicache_mem_layout="page_first_kv_split",
)
else:
args.hicache_mem_layout = "page_first_direct"
declare_resolution(
args,
"set_default_server_args",
hicache_mem_layout="page_first_direct",
)
@_call_once
File diff suppressed because it is too large Load Diff