config: record resolution writes in a declaration stash (#35905)
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5
parent
6218d6ce3f
commit
0e22777572
@@ -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."
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
+601
-183
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user