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
@@ -0,0 +1,381 @@
|
||||
"""Resolution writes are recorded, not just applied.
|
||||
|
||||
The projection that replaces field materialization reads the declaration stash,
|
||||
so a resolution write that only assigns the field is invisible to it. Every
|
||||
resolver declares now -- the record's handlers through `self._declare`, the
|
||||
hooks and hardware defaults through `declare_resolution` -- and that is pinned
|
||||
two ways: no bare assignment to a field survives anywhere a ServerArgs instance
|
||||
is in reach, and after resolution every declared field agrees with what the
|
||||
stash says. The second check is the one that keeps the transition honest --
|
||||
while a declaration still writes the field immediately, a stash entry and a
|
||||
field can only disagree if something assigned the field behind the stash's
|
||||
back. A third check runs the other way: every field resolution moved has to
|
||||
be explained by the stash, which covers the spellings a source scan cannot
|
||||
see.
|
||||
"""
|
||||
|
||||
import ast
|
||||
import dataclasses
|
||||
import json
|
||||
import os
|
||||
import pathlib
|
||||
import shutil
|
||||
import tempfile
|
||||
import unittest
|
||||
|
||||
import sglang
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||
|
||||
_SRT = pathlib.Path(sglang.__file__).resolve().parent / "srt"
|
||||
|
||||
# Every field of the record: resolution has no bare-assignment writer left, so
|
||||
# the scan states that as a whole rather than a converted-so-far list.
|
||||
_RESOLVED_FIELDS = frozenset(field.name for field in dataclasses.fields(ServerArgs))
|
||||
|
||||
# Shapes the agreement check runs on. Each needs a real config.json:
|
||||
# `model_path="dummy"` takes the pipeline's early return.
|
||||
_MINI_CONFIG = {
|
||||
"architectures": ["LlamaForCausalLM"],
|
||||
"model_type": "llama",
|
||||
"hidden_size": 16,
|
||||
"intermediate_size": 32,
|
||||
"num_attention_heads": 2,
|
||||
"num_key_value_heads": 2,
|
||||
"num_hidden_layers": 2,
|
||||
"vocab_size": 128,
|
||||
"max_position_embeddings": 2048,
|
||||
}
|
||||
|
||||
_SHAPES = (
|
||||
{"tp_size": 2, "dwdp_size": 2},
|
||||
{"random_seed": None},
|
||||
{"enable_deterministic_inference": True},
|
||||
{"enable_return_hidden_states": True},
|
||||
{
|
||||
"speculative_algorithm": "EAGLE",
|
||||
"speculative_num_steps": 3,
|
||||
"speculative_eagle_topk": 1,
|
||||
"speculative_num_draft_tokens": 4,
|
||||
},
|
||||
{"dp_size": 2, "tp_size": 2, "enable_dp_attention": True},
|
||||
{"enable_hierarchical_cache": True},
|
||||
{"disaggregation_mode": "prefill"},
|
||||
{"enable_lora": True, "max_lora_rank": 16},
|
||||
{"kv_cache_dtype": "fp8_e4m3", "page_size": 64},
|
||||
# A pass and a handler both decide this one: waterfill forces `deepep`
|
||||
# and the ascend handler wants `none`. Without this shape nothing in
|
||||
# the set reaches a field two writers disagree about.
|
||||
{"enable_waterfill": True, "moe_a2a_backend": "ascend_tp"},
|
||||
)
|
||||
|
||||
# Which converted fields the shapes above reach; the rest need a device or an
|
||||
# architecture no CPU fixture has, and the source scan covers those. Pinned so
|
||||
# a shape that stops reaching a field fails here. Add to it when adding a shape.
|
||||
_REACHED_BY_SHAPES = frozenset(
|
||||
{
|
||||
"_speculative_draft_quantization_explicitly_set",
|
||||
"allowed_media_domains",
|
||||
"attention_backend",
|
||||
"chunked_prefill_size",
|
||||
"cuda_graph_config",
|
||||
"custom_weight_loader",
|
||||
"device",
|
||||
"disable_cuda_graph",
|
||||
"disaggregation_ib_device",
|
||||
"dp_size",
|
||||
"enable_dp_attention",
|
||||
"enable_dp_attention_local_control_broadcast",
|
||||
"enable_dp_lm_head",
|
||||
"enable_flashinfer_allreduce_fusion",
|
||||
"encoder_transfer_backend",
|
||||
"enforce_disable_flashinfer_allreduce_fusion",
|
||||
"ep_size",
|
||||
"expert_distribution_recorder_buffer_size",
|
||||
"flashinfer_allreduce_fusion_backend",
|
||||
"grammar_backend",
|
||||
"hicache_ratio",
|
||||
"keep_mm_feature_on_device",
|
||||
"load_balance_method",
|
||||
"max_running_requests",
|
||||
"mem_fraction_static",
|
||||
"mm_feature_transport",
|
||||
"mm_process_config",
|
||||
"moe_a2a_backend",
|
||||
"moe_dense_tp_size",
|
||||
"moe_dp_size",
|
||||
"page_size",
|
||||
"random_seed",
|
||||
"return_hidden_states_mode",
|
||||
"sampling_backend",
|
||||
"schedule_conservativeness",
|
||||
"served_model_name",
|
||||
"speculative_algorithm",
|
||||
"speculative_draft_model_quantization",
|
||||
"tokenizer_path",
|
||||
"uses_mamba_radix_cache",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _server_args_writers(tree, path):
|
||||
"""Assignment targets that land on a ServerArgs instance.
|
||||
|
||||
Two mechanisms reach the same instance during resolution: a handler writing
|
||||
`self.<field>`, and a helper elsewhere in the tree writing through a
|
||||
`ServerArgs`-annotated parameter -- `set_default_server_args(args)` is
|
||||
called from the pipeline and writes `args.<field>`. Both bypass the
|
||||
declaration stash, so both have to be scanned; scanning only the handlers
|
||||
would let a field look converted while a second writer still assigns it.
|
||||
"""
|
||||
names = {"self"} if path.name == "server_args.py" else set()
|
||||
# A parameter *named* `server_args` counts with or without the annotation.
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
|
||||
continue
|
||||
args = node.args
|
||||
for arg in args.posonlyargs + args.args + args.kwonlyargs:
|
||||
annotation = arg.annotation
|
||||
if isinstance(annotation, ast.Constant):
|
||||
text = annotation.value
|
||||
elif isinstance(annotation, ast.Name):
|
||||
text = annotation.id
|
||||
elif isinstance(annotation, ast.Attribute):
|
||||
text = annotation.attr
|
||||
else:
|
||||
continue
|
||||
if text == "ServerArgs":
|
||||
names.add(arg.arg)
|
||||
names |= {
|
||||
arg.arg for arg in args.posonlyargs + args.args if arg.arg == "server_args"
|
||||
}
|
||||
return names
|
||||
|
||||
|
||||
def _bare_assignments():
|
||||
"""Assignments to a converted field that never reach the stash."""
|
||||
found = []
|
||||
for path in sorted(_SRT.rglob("*.py")):
|
||||
try:
|
||||
tree = ast.parse(path.read_text(encoding="utf-8-sig"))
|
||||
except SyntaxError:
|
||||
continue
|
||||
names = _server_args_writers(tree, path)
|
||||
if not names:
|
||||
continue
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.Assign):
|
||||
targets = node.targets
|
||||
elif isinstance(node, (ast.AugAssign, ast.AnnAssign)):
|
||||
targets = [node.target]
|
||||
else:
|
||||
continue
|
||||
# Destructured targets count: `(sa.a, sa.b) = f()` writes two
|
||||
# fields and is not an `ast.Attribute` at the top level.
|
||||
flat = []
|
||||
for target in targets:
|
||||
if isinstance(target, (ast.Tuple, ast.List)):
|
||||
flat.extend(target.elts)
|
||||
else:
|
||||
flat.append(target)
|
||||
for target in flat:
|
||||
if (
|
||||
isinstance(target, ast.Attribute)
|
||||
and isinstance(target.value, ast.Name)
|
||||
and target.value.id in names
|
||||
and target.attr in _RESOLVED_FIELDS
|
||||
):
|
||||
found.append(
|
||||
f"{path.relative_to(_SRT)}:{node.lineno} "
|
||||
f"{target.value.id}.{target.attr}"
|
||||
)
|
||||
return sorted(found)
|
||||
|
||||
|
||||
def _stash_overlay(server_args):
|
||||
"""What the declarations say, last writer wins -- the projection's input."""
|
||||
overlay = {}
|
||||
for _source, declared in getattr(server_args, "_resolved_overrides", None) or ():
|
||||
overlay.update(declared)
|
||||
return overlay
|
||||
|
||||
|
||||
class TestResolutionDeclarations(CustomTestCase):
|
||||
def setUp(self):
|
||||
# Resolution writes environment variables, and those outlive the
|
||||
# record that set them.
|
||||
super().setUp()
|
||||
environment = dict(os.environ)
|
||||
|
||||
def restore():
|
||||
os.environ.clear()
|
||||
os.environ.update(environment)
|
||||
|
||||
self.addCleanup(restore)
|
||||
|
||||
def _resolve(self, extra):
|
||||
"""A fully-resolved config: a real config.json, so the pipeline runs
|
||||
past its dummy-model early return."""
|
||||
path = tempfile.mkdtemp(prefix="declarations_")
|
||||
self.addCleanup(shutil.rmtree, path, ignore_errors=True)
|
||||
with open(os.path.join(path, "config.json"), "w") as handle:
|
||||
json.dump(_MINI_CONFIG, handle)
|
||||
fields = {"random_seed": 42}
|
||||
fields.update(extra)
|
||||
return ServerArgs(model_path=path, device="cuda", **fields)
|
||||
|
||||
def test_converted_fields_are_not_assigned_bare(self):
|
||||
bare = _bare_assignments()
|
||||
self.assertEqual(
|
||||
bare,
|
||||
[],
|
||||
"a converted field is assigned directly, so the projection would "
|
||||
"not see this write:\n " + "\n ".join(bare),
|
||||
)
|
||||
|
||||
def test_the_stash_accounts_for_every_change_resolution_made(self):
|
||||
"""The other direction: a field resolution moved is in the stash.
|
||||
|
||||
The source scan states that no *assignment* escapes, which leaves the
|
||||
spellings a source scan cannot see -- a computed name, a write through
|
||||
a helper the scan does not recognize as holding the record. This
|
||||
compares the resolved value against what the caller supplied (or the
|
||||
field's default) and asks the stash to explain every difference, which
|
||||
is what the projection has to be able to do.
|
||||
"""
|
||||
unexplained = []
|
||||
for shape in _SHAPES:
|
||||
supplied = {"random_seed": 42, **shape}
|
||||
server_args = self._resolve(shape)
|
||||
overlay = _stash_overlay(server_args)
|
||||
for field in dataclasses.fields(server_args):
|
||||
if field.name in ("model_path", "device") or field.name in overlay:
|
||||
continue
|
||||
if field.name in supplied:
|
||||
before = supplied[field.name]
|
||||
elif field.default is not dataclasses.MISSING:
|
||||
before = field.default
|
||||
elif field.default_factory is not dataclasses.MISSING:
|
||||
before = field.default_factory()
|
||||
else:
|
||||
continue
|
||||
after = getattr(server_args, field.name, None)
|
||||
if after != before:
|
||||
unexplained.append(
|
||||
f"{shape} -> {field.name}: {before!r} -> {after!r}"
|
||||
)
|
||||
self.assertEqual(
|
||||
unexplained,
|
||||
[],
|
||||
"resolution moved these fields without declaring them, so the "
|
||||
"projection would answer with the unresolved value:\n "
|
||||
+ "\n ".join(unexplained),
|
||||
)
|
||||
|
||||
def test_the_stash_agrees_with_the_fields_it_declared(self):
|
||||
mismatches = []
|
||||
for shape in _SHAPES:
|
||||
server_args = self._resolve(shape)
|
||||
overlay = _stash_overlay(server_args)
|
||||
for field, declared in overlay.items():
|
||||
if field not in _RESOLVED_FIELDS:
|
||||
continue
|
||||
actual = getattr(server_args, field)
|
||||
if actual != declared:
|
||||
mismatches.append(
|
||||
f"{shape} -> {field}: field={actual!r} stash={declared!r}"
|
||||
)
|
||||
self.assertEqual(
|
||||
mismatches,
|
||||
[],
|
||||
"a declared field and its stash entry disagree, so something "
|
||||
"assigned the field behind the declaration:\n " + "\n ".join(mismatches),
|
||||
)
|
||||
|
||||
def test_no_immediate_writer_overrides_a_deferred_one(self):
|
||||
"""A handler must not declare over a value a pass already decided.
|
||||
|
||||
Routing a bare write into the stash changed which writer wins: the
|
||||
appended entry is replayed last, so a handler now beats a pass or a
|
||||
registry entry that ran earlier -- where before, the pass's declaration
|
||||
was applied on top of the handler's bare write. One handler was found
|
||||
that way (it gated on the raw field while its neighbours read the
|
||||
resolving view, so `--enable-waterfill --moe-a2a-backend ascend_tp`
|
||||
silently stopped forcing `deepep`).
|
||||
|
||||
This is the invariant rather than that instance: walk the stash in
|
||||
order and fail when an immediate writer declares a field whose previous
|
||||
entry came from a deferred writer with a *different* value. The
|
||||
deferred sources are derived from the live registries and the constant
|
||||
override table, so a new pass is covered without being listed.
|
||||
"""
|
||||
from sglang.srt.arg_groups import overrides
|
||||
|
||||
deferred = {
|
||||
getattr(fn, "__qualname__", getattr(fn, "__name__", ""))
|
||||
for fn in overrides.POST_PROCESS_PASSES
|
||||
}
|
||||
deferred |= {
|
||||
getattr(fn, "__qualname__", getattr(fn, "__name__", ""))
|
||||
for fns in overrides._MODEL_OVERRIDE_FNS.values()
|
||||
for fn in fns
|
||||
}
|
||||
deferred |= {
|
||||
getattr(fn, "__qualname__", getattr(fn, "__name__", ""))
|
||||
for _predicate, fn in overrides._PREDICATE_OVERRIDE_FNS
|
||||
}
|
||||
# The constant arch -> {field: value} table is a deferred writer too --
|
||||
# it has no callable, and its stash source is spelled by the collector
|
||||
# (`MODEL_OVERRIDES[<arch>]`).
|
||||
deferred |= {f"MODEL_OVERRIDES[{arch!r}]" for arch in overrides.MODEL_OVERRIDES}
|
||||
self.assertGreater(
|
||||
len(deferred), 40, "the deferred-writer set collapsed; nothing to compare"
|
||||
)
|
||||
|
||||
inversions = []
|
||||
for shape in _SHAPES:
|
||||
server_args = self._resolve(shape)
|
||||
decided_by = {}
|
||||
for source, declared in getattr(server_args, "_resolved_overrides", []):
|
||||
for field, value in declared.items():
|
||||
previous = decided_by.get(field)
|
||||
if (
|
||||
previous is not None
|
||||
and previous[0] in deferred
|
||||
and source not in deferred
|
||||
and previous[1] != value
|
||||
):
|
||||
inversions.append(
|
||||
f"{shape} -> {field}: {previous[0]} decided "
|
||||
f"{previous[1]!r}, then {source} declared {value!r}"
|
||||
)
|
||||
decided_by[field] = (source, value)
|
||||
self.assertEqual(
|
||||
inversions,
|
||||
[],
|
||||
"a handler declared over a value a pass or a registry entry had "
|
||||
"already decided; if the handler is meant to win, say so, and if "
|
||||
"it is gating on the field, it has to read the resolving view:\n "
|
||||
+ "\n ".join(inversions),
|
||||
)
|
||||
|
||||
def test_the_shapes_reach_the_fields_they_are_meant_to(self):
|
||||
"""A green agreement check over an empty stash would prove nothing."""
|
||||
declared = set()
|
||||
for shape in _SHAPES:
|
||||
declared |= set(_stash_overlay(self._resolve(shape))) & _RESOLVED_FIELDS
|
||||
missing = sorted(_REACHED_BY_SHAPES - declared)
|
||||
self.assertEqual(
|
||||
missing,
|
||||
[],
|
||||
"the shapes no longer reach these converted fields, so the "
|
||||
"agreement check silently stopped covering them:\n "
|
||||
+ "\n ".join(missing),
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -253,15 +253,19 @@ class TestPublishInstallsSlot(_IsolatedPublish):
|
||||
"""Publish wiring: set_server_args installs the already-resolved object
|
||||
into the context-owned slot (no transformation at publish time)."""
|
||||
|
||||
def test_dummy_fixture_has_empty_stash_and_publishes_cleanly(self):
|
||||
def test_dummy_fixture_publishes_the_object_it_resolved(self):
|
||||
from sglang.srt.server_args import (
|
||||
ServerArgs,
|
||||
set_global_server_args_for_scheduler,
|
||||
)
|
||||
|
||||
sa = ServerArgs(model_path="dummy") # __post_init__ early-returns
|
||||
# The stash is created before the dummy short-circuit and stays empty.
|
||||
self.assertEqual(sa._resolved_overrides, [])
|
||||
# A dummy path short-circuits the pipeline, but the handlers ahead of
|
||||
# that point still declare; whatever they left in the stash is on the
|
||||
# object by the time publish sees it.
|
||||
for source, declared in sa._resolved_overrides:
|
||||
for field, value in declared.items():
|
||||
self.assertEqual(getattr(sa, field), value, f"{source}: {field}")
|
||||
set_global_server_args_for_scheduler(sa)
|
||||
self.assertIs(get_server_args(), sa)
|
||||
|
||||
|
||||
@@ -663,14 +663,22 @@ class TestSuppliedInstanceExposure(CustomTestCase):
|
||||
`speculative_draft_attention_backend`, and no entry ran it) showed
|
||||
that a family nobody listed leaves its readers unpinned. The hook
|
||||
modules under `arg_groups/` are the resolution pipeline's extension
|
||||
points, and their assignment surface (`server_args.field = ...`) is
|
||||
the may-write set, family-blind by construction. Collected like the
|
||||
points -- along with the NPU default helper, which the pipeline calls
|
||||
the same way -- and their write surface is the may-write set,
|
||||
family-blind by
|
||||
construction. A hook writes two ways: `server_args.field = ...`, and
|
||||
`declare_resolution(server_args, source, field=...)`, which records
|
||||
the write in the declaration stash on its way to the field. Counting
|
||||
only the assignment would read a hook's conversion to a declaration as
|
||||
the field having stopped being written. Collected like the
|
||||
late-resolution keywords: statically, failing loudly on an
|
||||
unparsable module. Underscore-prefixed targets are pipeline
|
||||
bookkeeping, not config leaves.
|
||||
"""
|
||||
targets = set()
|
||||
for path in sorted((_PACKAGE_ROOT / "arg_groups").glob("*.py")):
|
||||
modules = sorted((_PACKAGE_ROOT / "arg_groups").glob("*.py"))
|
||||
modules.append(_PACKAGE_ROOT / "hardware_backend/npu/utils.py")
|
||||
for path in modules:
|
||||
try:
|
||||
tree = ast.parse(path.read_text(encoding="utf-8-sig"))
|
||||
except SyntaxError:
|
||||
@@ -680,6 +688,17 @@ class TestSuppliedInstanceExposure(CustomTestCase):
|
||||
tgts = node.targets
|
||||
elif isinstance(node, (ast.AnnAssign, ast.AugAssign)):
|
||||
tgts = [node.target]
|
||||
elif (
|
||||
isinstance(node, ast.Call)
|
||||
and isinstance(node.func, ast.Name)
|
||||
and node.func.id == "declare_resolution"
|
||||
):
|
||||
targets |= {
|
||||
kw.arg
|
||||
for kw in node.keywords
|
||||
if kw.arg and not kw.arg.startswith("_")
|
||||
}
|
||||
continue
|
||||
else:
|
||||
continue
|
||||
for tgt in tgts:
|
||||
@@ -702,6 +721,12 @@ class TestSuppliedInstanceExposure(CustomTestCase):
|
||||
A write site that can never fire is a dead branch to delete upstream,
|
||||
not a census exemption. Only names that are declared dataclass fields
|
||||
count; underscore bookkeeping does not.
|
||||
|
||||
Two spellings write: an assignment, and ``self._declare(source,
|
||||
field=value)``, which records the write in the declaration stash on
|
||||
its way to the field. Counting only assignments would read a handler's
|
||||
conversion to a declaration as the field having stopped being written,
|
||||
which would quietly retire every pinned pair that reads it.
|
||||
"""
|
||||
tree = ast.parse(
|
||||
(_PACKAGE_ROOT / "server_args.py").read_text(encoding="utf-8-sig")
|
||||
@@ -722,6 +747,17 @@ class TestSuppliedInstanceExposure(CustomTestCase):
|
||||
tgts = node.targets
|
||||
elif isinstance(node, (ast.AnnAssign, ast.AugAssign)):
|
||||
tgts = [node.target]
|
||||
elif (
|
||||
isinstance(node, ast.Call)
|
||||
and isinstance(node.func, ast.Attribute)
|
||||
and node.func.attr == "_declare"
|
||||
):
|
||||
targets |= {
|
||||
kw.arg
|
||||
for kw in node.keywords
|
||||
if kw.arg in declared and not kw.arg.startswith("_")
|
||||
}
|
||||
continue
|
||||
else:
|
||||
continue
|
||||
for tgt in tgts:
|
||||
@@ -733,27 +769,24 @@ class TestSuppliedInstanceExposure(CustomTestCase):
|
||||
and not tgt.attr.startswith("_")
|
||||
):
|
||||
targets.add(tgt.attr)
|
||||
# The deprecated-alias normalization loop writes through a *name
|
||||
# tuple* (`for attr in (...): setattr(self, attr, "dsv4")`), which no
|
||||
# assignment scan sees; its field set is pinned here with a drift
|
||||
# guard on the tuple itself.
|
||||
# The deprecated-alias normalization declares through `**renamed`, so
|
||||
# the keyword scan sees no names; its field set is pinned here.
|
||||
alias_fields = {
|
||||
"attention_backend",
|
||||
"decode_attention_backend",
|
||||
"prefill_attention_backend",
|
||||
"speculative_draft_attention_backend",
|
||||
}
|
||||
deprecated = next(
|
||||
node
|
||||
for node in ast.walk(sa_class)
|
||||
if isinstance(node, ast.FunctionDef)
|
||||
and node.name == "_handle_deprecated_args"
|
||||
)
|
||||
found_tuples = [
|
||||
{elt.value for elt in node.iter.elts if isinstance(elt, ast.Constant)}
|
||||
for node in ast.walk(sa_class)
|
||||
if isinstance(node, ast.For)
|
||||
and isinstance(node.iter, ast.Tuple)
|
||||
and any(
|
||||
isinstance(inner, ast.Call)
|
||||
and isinstance(inner.func, ast.Name)
|
||||
and inner.func.id == "setattr"
|
||||
for inner in ast.walk(node)
|
||||
)
|
||||
for node in ast.walk(deprecated)
|
||||
if isinstance(node, ast.For) and isinstance(node.iter, ast.Tuple)
|
||||
]
|
||||
self.assertIn(
|
||||
alias_fields,
|
||||
|
||||
Reference in New Issue
Block a user