config: resolution reads the declarations, not the fields (#36253)
This commit is contained in:
@@ -3,7 +3,10 @@ from __future__ import annotations
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sglang.srt.arg_groups.overrides import declare_resolution
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
declare_resolution,
|
||||
resolving_view,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -16,18 +19,16 @@ def validate_deepseek_v4_mega_moe_token_budget(
|
||||
server_args: ServerArgs,
|
||||
) -> None:
|
||||
"""Ensure the DSV4 prefill budget fits MegaMoE's per-rank buffer."""
|
||||
mega_moe_enabled = server_args.moe_a2a_backend == "megamoe"
|
||||
if not mega_moe_enabled or server_args.disaggregation_mode == "decode":
|
||||
cfg = resolving_view(server_args)
|
||||
mega_moe_enabled = cfg.moe_a2a_backend == "megamoe"
|
||||
if not mega_moe_enabled or cfg.disaggregation_mode == "decode":
|
||||
# decode node will skip the check because decode bs is not relevant with --chunk-prefill-size
|
||||
return
|
||||
|
||||
if server_args.pp_size > 1 and server_args.enable_dynamic_chunking:
|
||||
if cfg.pp_size > 1 and cfg.enable_dynamic_chunking:
|
||||
return
|
||||
|
||||
if (
|
||||
server_args.chunked_prefill_size is None
|
||||
or server_args.chunked_prefill_size <= 0
|
||||
):
|
||||
if cfg.chunked_prefill_size is None or cfg.chunked_prefill_size <= 0:
|
||||
raise ValueError(
|
||||
"DeepSeekV4 with MegaMoE requires chunked prefill to be enabled. "
|
||||
"Set --chunked-prefill-size to a positive value; "
|
||||
@@ -35,40 +36,38 @@ def validate_deepseek_v4_mega_moe_token_budget(
|
||||
"token requirement would not have a strict prefill-forward bound."
|
||||
)
|
||||
|
||||
if server_args.enable_prefill_cp:
|
||||
token_partition_size = server_args.attn_cp_size
|
||||
if cfg.enable_prefill_cp:
|
||||
token_partition_size = cfg.attn_cp_size
|
||||
token_partition_name = "attn_cp_size"
|
||||
token_alignment = 1
|
||||
local_chunked_prefill_size = (
|
||||
server_args.chunked_prefill_size + token_partition_size - 1
|
||||
cfg.chunked_prefill_size + token_partition_size - 1
|
||||
) // token_partition_size
|
||||
elif server_args.enable_dp_attention:
|
||||
token_partition_size = server_args.dp_size
|
||||
elif cfg.enable_dp_attention:
|
||||
token_partition_size = cfg.dp_size
|
||||
token_partition_name = "dp_size"
|
||||
token_alignment = max(
|
||||
server_args.tp_size // server_args.dp_size // server_args.attn_cp_size,
|
||||
cfg.tp_size // cfg.dp_size // cfg.attn_cp_size,
|
||||
1,
|
||||
)
|
||||
local_chunked_prefill_size = (
|
||||
server_args.chunked_prefill_size // token_partition_size
|
||||
)
|
||||
local_chunked_prefill_size = cfg.chunked_prefill_size // token_partition_size
|
||||
else:
|
||||
# Pure TP and PP with static chunking are handled here.
|
||||
token_partition_size = 1
|
||||
token_partition_name = "none"
|
||||
# global_num_tokens will ceil_align to attn_tp_size so the validation needs to do alignment as well
|
||||
token_alignment = max(
|
||||
server_args.tp_size // token_partition_size // server_args.attn_cp_size,
|
||||
cfg.tp_size // token_partition_size // cfg.attn_cp_size,
|
||||
1,
|
||||
)
|
||||
local_chunked_prefill_size = server_args.chunked_prefill_size
|
||||
local_chunked_prefill_size = cfg.chunked_prefill_size
|
||||
|
||||
if local_chunked_prefill_size <= 0:
|
||||
raise ValueError(
|
||||
"DeepSeekV4 with MegaMoE requires a positive effective per-rank "
|
||||
"chunked prefill size. "
|
||||
f"Current values: chunked_prefill_size="
|
||||
f"{server_args.chunked_prefill_size}, "
|
||||
f"{cfg.chunked_prefill_size}, "
|
||||
f"token_partition={token_partition_name}, "
|
||||
f"token_partition_size={token_partition_size}."
|
||||
)
|
||||
@@ -87,7 +86,7 @@ def validate_deepseek_v4_mega_moe_token_budget(
|
||||
"SGLANG_OPT_DEEPGEMM_MEGA_MOE_NUM_MAX_TOKENS_PER_RANK to "
|
||||
"cover each rank's effective prefill token budget. "
|
||||
f"Current values: chunked_prefill_size="
|
||||
f"{server_args.chunked_prefill_size}, "
|
||||
f"{cfg.chunked_prefill_size}, "
|
||||
f"token_partition={token_partition_name}, "
|
||||
f"token_partition_size={token_partition_size}, "
|
||||
f"token_alignment={token_alignment}, "
|
||||
@@ -112,6 +111,7 @@ def apply_deepseek_v4_defaults(server_args: ServerArgs, model_arch: str) -> None
|
||||
max_running_requests fill (the speculative hook is a later writer of
|
||||
that field) and the validations.
|
||||
"""
|
||||
cfg = resolving_view(server_args)
|
||||
from sglang.srt.utils import is_hip
|
||||
|
||||
# FlashMLA sparse prefill (SGLANG_OPT_FLASHMLA_SPARSE_PREFILL, default on)
|
||||
@@ -136,36 +136,36 @@ 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:
|
||||
if cfg.max_running_requests is None:
|
||||
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}."
|
||||
f"Setting max_running_requests to {cfg.max_running_requests} for {model_arch}."
|
||||
)
|
||||
|
||||
if server_args.speculative_algorithm is not None:
|
||||
assert server_args.speculative_algorithm in (
|
||||
if cfg.speculative_algorithm is not None:
|
||||
assert cfg.speculative_algorithm in (
|
||||
"EAGLE",
|
||||
"DSPARK",
|
||||
), f"Only EAGLE and DSPARK speculative algorithms are supported for {model_arch}"
|
||||
if server_args.speculative_algorithm == "EAGLE":
|
||||
if cfg.speculative_algorithm == "EAGLE":
|
||||
assert (
|
||||
server_args.speculative_eagle_topk == 1
|
||||
cfg.speculative_eagle_topk == 1
|
||||
), f"Only EAGLE speculative algorithm with topk == 1 is supported for {model_arch}"
|
||||
|
||||
|
||||
def validate_deepseek_v4_cp(server_args: ServerArgs) -> None:
|
||||
"""Validate DeepSeek V4 context-parallel configuration."""
|
||||
if not server_args.enable_prefill_cp:
|
||||
cfg = resolving_view(server_args)
|
||||
if not cfg.enable_prefill_cp:
|
||||
return
|
||||
|
||||
if server_args.cp_strategy != "interleave":
|
||||
if cfg.cp_strategy != "interleave":
|
||||
raise ValueError(
|
||||
"DeepSeekV4 only supports interleave CP strategy, "
|
||||
f"got {server_args.cp_strategy}"
|
||||
"DeepSeekV4 only supports interleave CP strategy, " f"got {cfg.cp_strategy}"
|
||||
)
|
||||
|
||||
declare_resolution(
|
||||
@@ -196,19 +196,19 @@ def validate_deepseek_v4_cp(server_args: ServerArgs) -> None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"validate_deepseek_v4_cp",
|
||||
attn_cp_size=server_args.tp_size // server_args.dp_size,
|
||||
attn_cp_size=cfg.tp_size // cfg.dp_size,
|
||||
)
|
||||
assert (
|
||||
server_args.dp_size == 1
|
||||
cfg.dp_size == 1
|
||||
), "For round-robin split mode, dp attention is not supported."
|
||||
assert (
|
||||
server_args.tp_size <= 8
|
||||
cfg.tp_size <= 8
|
||||
), "Context parallel only supports single machine (tp_size <= 8). Cross-machine CP has precision issues."
|
||||
if server_args.moe_a2a_backend not in ("none", "deepep", "megamoe"):
|
||||
if cfg.moe_a2a_backend not in ("none", "deepep", "megamoe"):
|
||||
raise ValueError(
|
||||
"DeepSeekV4 CP supports moe_a2a_backend in "
|
||||
"('none', 'deepep', 'megamoe'), "
|
||||
f"got {server_args.moe_a2a_backend!r}."
|
||||
f"got {cfg.moe_a2a_backend!r}."
|
||||
)
|
||||
logger.warning(
|
||||
"Disabling SGLANG_OPT_FLASHMLA_SPARSE_PREFILL because DeepSeekV4 "
|
||||
@@ -217,6 +217,6 @@ def validate_deepseek_v4_cp(server_args: ServerArgs) -> None:
|
||||
envs.SGLANG_OPT_FLASHMLA_SPARSE_PREFILL.set(False)
|
||||
logger.warning(
|
||||
f"Enable Context Parallel for DeepSeekV4, "
|
||||
f"dp_size={server_args.dp_size}, moe_dense_tp_size={server_args.moe_dense_tp_size}, "
|
||||
f"attn_cp_size={server_args.attn_cp_size}, ep_size={server_args.ep_size}, tp_size={server_args.tp_size}"
|
||||
f"dp_size={cfg.dp_size}, moe_dense_tp_size={cfg.moe_dense_tp_size}, "
|
||||
f"attn_cp_size={cfg.attn_cp_size}, ep_size={cfg.ep_size}, tp_size={cfg.tp_size}"
|
||||
)
|
||||
|
||||
@@ -9,7 +9,7 @@ import os
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from sglang.srt.arg_groups.overrides import declare_resolution
|
||||
from sglang.srt.arg_groups.overrides import declare_resolution, resolving_view
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.model_executor.cuda_graph_config import (
|
||||
Backend,
|
||||
@@ -27,31 +27,32 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
def handle_expert_pack(server_args: Any) -> None:
|
||||
"""Normalize expert-pack settings and report all startup errors together."""
|
||||
if server_args.load_format != "expert_pack":
|
||||
cfg = resolving_view(server_args)
|
||||
if cfg.load_format != "expert_pack":
|
||||
return
|
||||
|
||||
errors = []
|
||||
parallelism = (
|
||||
("tensor", "--tp-size", server_args.tp_size),
|
||||
("data", "--dp-size", server_args.dp_size),
|
||||
("expert", "--ep-size", server_args.ep_size),
|
||||
("tensor", "--tp-size", cfg.tp_size),
|
||||
("data", "--dp-size", cfg.dp_size),
|
||||
("expert", "--ep-size", cfg.ep_size),
|
||||
)
|
||||
for label, option, size in parallelism:
|
||||
if size != 1:
|
||||
errors.append(f"{label} parallelism ({option}) must be 1, got {size}")
|
||||
|
||||
if server_args.enforce_shared_experts_fusion:
|
||||
if cfg.enforce_shared_experts_fusion:
|
||||
errors.append(
|
||||
"--enforce-shared-experts-fusion is incompatible with expert_pack"
|
||||
)
|
||||
if server_args.enable_waterfill:
|
||||
if cfg.enable_waterfill:
|
||||
errors.append("--enable-waterfill is incompatible with expert_pack")
|
||||
|
||||
explicit_cuda_graph_backends = {
|
||||
Phase.DECODE: server_args.cuda_graph_backend_decode,
|
||||
Phase.PREFILL: server_args.cuda_graph_backend_prefill,
|
||||
Phase.DECODE: cfg.cuda_graph_backend_decode,
|
||||
Phase.PREFILL: cfg.cuda_graph_backend_prefill,
|
||||
}
|
||||
raw_cuda_graph_config = server_args.cuda_graph_config
|
||||
raw_cuda_graph_config = cfg.cuda_graph_config
|
||||
if isinstance(raw_cuda_graph_config, CudaGraphConfig):
|
||||
raw_cuda_graph_config = raw_cuda_graph_config.to_dict()
|
||||
for phase in Phase.ALL:
|
||||
@@ -69,7 +70,7 @@ def handle_expert_pack(server_args: Any) -> None:
|
||||
f"disabled, got {explicit_backend!r}"
|
||||
)
|
||||
|
||||
loader_config = server_args.model_loader_extra_config or {}
|
||||
loader_config = cfg.model_loader_extra_config or {}
|
||||
if isinstance(loader_config, str):
|
||||
try:
|
||||
loader_config = json.loads(loader_config)
|
||||
@@ -82,7 +83,7 @@ def handle_expert_pack(server_args: Any) -> None:
|
||||
|
||||
# A raw GGUF path is the public input form. Preparation is performed once
|
||||
# here, before model-config parsing and before the loader is constructed.
|
||||
raw_model_path = Path(server_args.model_path).expanduser()
|
||||
raw_model_path = Path(cfg.model_path).expanduser()
|
||||
raw_preparation_failed = False
|
||||
if not errors and raw_model_path.is_file():
|
||||
try:
|
||||
@@ -136,12 +137,12 @@ def handle_expert_pack(server_args: Any) -> None:
|
||||
errors.append(f"expert-pack file does not exist: {pack_path}")
|
||||
|
||||
model_kind = None
|
||||
model_path = parse_path("--model-path", server_args.model_path)
|
||||
model_path = parse_path("--model-path", cfg.model_path)
|
||||
if not raw_preparation_failed:
|
||||
if model_path is None or not model_path.is_dir():
|
||||
errors.append(
|
||||
"--model-path must be a local GGUF shard or tokenizer/config "
|
||||
f"directory for expert_pack, got {server_args.model_path!r}"
|
||||
f"directory for expert_pack, got {cfg.model_path!r}"
|
||||
)
|
||||
else:
|
||||
try:
|
||||
|
||||
@@ -3,6 +3,8 @@ from __future__ import annotations
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
@@ -79,7 +81,8 @@ def validate_hisparse_kv_cache_dtype(server_args: ServerArgs) -> None:
|
||||
|
||||
def validate_hisparse(server_args: ServerArgs) -> None:
|
||||
"""Validate --enable-hisparse constraints (model class, radix cache, DSA backend)."""
|
||||
if not server_args.enable_hisparse:
|
||||
cfg = resolving_view(server_args)
|
||||
if not cfg.enable_hisparse:
|
||||
return
|
||||
|
||||
from sglang.srt.configs.model_config import (
|
||||
@@ -96,7 +99,7 @@ def validate_hisparse(server_args: ServerArgs) -> None:
|
||||
)
|
||||
|
||||
assert (
|
||||
server_args.disable_radix_cache
|
||||
cfg.disable_radix_cache
|
||||
), "Hierarchical sparse attention currently requires --disable-radix-cache."
|
||||
|
||||
# DSv4 hisparse handles its own dtype/backend pairing elsewhere; the dtype-
|
||||
|
||||
@@ -3,7 +3,10 @@ from __future__ import annotations
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sglang.srt.arg_groups.overrides import declare_resolution
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
declare_resolution,
|
||||
resolving_view,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
@@ -13,15 +16,16 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
def apply_kimi_k3_spec_backend_defaults(server_args: ServerArgs) -> None:
|
||||
"""Apply speculative backend defaults for Kimi hybrid models."""
|
||||
cfg = resolving_view(server_args)
|
||||
from sglang.srt.utils import is_sm100_supported
|
||||
|
||||
if server_args.speculative_algorithm is None:
|
||||
if cfg.speculative_algorithm is None:
|
||||
return
|
||||
|
||||
# Use the fused Kimi-K3/DSPARK CuTeDSL kernel for KDA target verification.
|
||||
# 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:
|
||||
if cfg.linear_attn_verify_backend is None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"apply_kimi_k3_spec_backend_defaults",
|
||||
@@ -36,8 +40,8 @@ def apply_kimi_k3_spec_backend_defaults(server_args: ServerArgs) -> None:
|
||||
# dspark's draft is dense MQA; trtllm_mha avoids flashinfer's blocking
|
||||
# per-step host plan. DSPARK-only: other spec algos use MLA-family drafts.
|
||||
if (
|
||||
server_args.speculative_algorithm == "DSPARK"
|
||||
and server_args.speculative_draft_attention_backend is None
|
||||
cfg.speculative_algorithm == "DSPARK"
|
||||
and cfg.speculative_draft_attention_backend is None
|
||||
and is_sm100_supported()
|
||||
):
|
||||
declare_resolution(
|
||||
@@ -63,20 +67,21 @@ def disable_kimi_k3_symm_mem(server_args: ServerArgs) -> None:
|
||||
Gates on the arch itself: this runs from cuda-graph resolution, which is earlier
|
||||
than the model-specific hook block.
|
||||
"""
|
||||
cfg = resolving_view(server_args)
|
||||
from sglang.srt.connector import ConnectorType
|
||||
from sglang.srt.model_executor.cuda_graph_config import Backend
|
||||
from sglang.srt.utils import parse_connector_type
|
||||
|
||||
if not server_args.enable_symm_mem:
|
||||
if not cfg.enable_symm_mem:
|
||||
return
|
||||
if parse_connector_type(server_args.model_path) == ConnectorType.INSTANCE:
|
||||
if parse_connector_type(cfg.model_path) == ConnectorType.INSTANCE:
|
||||
return
|
||||
if server_args.get_model_config().hf_config.architectures[0] not in (
|
||||
"KimiLinearForCausalLM",
|
||||
"KimiK3ForConditionalGeneration",
|
||||
):
|
||||
return
|
||||
graph = server_args.cuda_graph_config
|
||||
graph = cfg.cuda_graph_config
|
||||
if (
|
||||
graph.decode.backend == Backend.DISABLED
|
||||
and graph.prefill.backend == Backend.DISABLED
|
||||
@@ -100,14 +105,15 @@ def disable_kimi_k3_symm_mem(server_args: ServerArgs) -> None:
|
||||
|
||||
def apply_kimi_k3_linear_attn_defaults(server_args: ServerArgs) -> None:
|
||||
"""KDA decode-fallback default for Kimi hybrid models (spec-independent)."""
|
||||
cfg = resolving_view(server_args)
|
||||
from sglang.srt.utils import is_sm100_supported
|
||||
|
||||
# Preempts the generic SM100+bf16 flashinfer switch (a GDN default): on
|
||||
# KDA shapes the triton packed decode measures ~35% faster than
|
||||
# recurrent_kda across bs 1-256, and ReplaySSM requires triton.
|
||||
if (
|
||||
server_args.linear_attn_decode_backend is None
|
||||
and server_args.mamba_ssm_dtype == "bfloat16"
|
||||
cfg.linear_attn_decode_backend is None
|
||||
and cfg.mamba_ssm_dtype == "bfloat16"
|
||||
and is_sm100_supported()
|
||||
):
|
||||
declare_resolution(
|
||||
|
||||
@@ -7,7 +7,10 @@ from typing import TYPE_CHECKING
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
from sglang.srt.arg_groups.overrides import declare_resolution
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
declare_resolution,
|
||||
resolving_view,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -18,15 +21,16 @@ def handle_mega_moe(server_args: ServerArgs) -> None:
|
||||
|
||||
|
||||
def handle_moe_runner_backend_alias(server_args: ServerArgs) -> None:
|
||||
if server_args.moe_runner_backend != "megamoe":
|
||||
cfg = resolving_view(server_args)
|
||||
if cfg.moe_runner_backend != "megamoe":
|
||||
return
|
||||
|
||||
if server_args.moe_a2a_backend not in ("none", "megamoe"):
|
||||
if cfg.moe_a2a_backend not in ("none", "megamoe"):
|
||||
logger.warning(
|
||||
"--moe-runner-backend megamoe is an alias for "
|
||||
"--moe-a2a-backend megamoe; overriding "
|
||||
"--moe-a2a-backend %s.",
|
||||
server_args.moe_a2a_backend,
|
||||
cfg.moe_a2a_backend,
|
||||
)
|
||||
declare_resolution(
|
||||
server_args,
|
||||
@@ -37,7 +41,8 @@ def handle_moe_runner_backend_alias(server_args: ServerArgs) -> None:
|
||||
|
||||
|
||||
def handle_w4a4_mxfp4_megamoe_env(server_args: ServerArgs) -> None:
|
||||
if not server_args.enable_w4a4_mxfp4_megamoe:
|
||||
cfg = resolving_view(server_args)
|
||||
if not cfg.enable_w4a4_mxfp4_megamoe:
|
||||
return
|
||||
|
||||
os.environ["DG_USE_FP4_ACTS"] = "1"
|
||||
|
||||
@@ -150,6 +150,42 @@ class ResolvedView:
|
||||
)
|
||||
|
||||
|
||||
class ResolvingConfig:
|
||||
"""Live read view of the resolution result: the declaration stash over the
|
||||
record's fields, looked up per read.
|
||||
|
||||
``ResolvedView`` snapshots the overlay when it is built, which is what a
|
||||
post-process pass wants -- it reads the state at its slot. A resolver that
|
||||
reads *after* declaring, or after calling something that declares, needs the
|
||||
current answer instead, so this one walks the stash on every read. It falls
|
||||
through to the field, which is where the raw input lives.
|
||||
"""
|
||||
|
||||
__slots__ = ("_server_args",)
|
||||
|
||||
def __init__(self, server_args: Any):
|
||||
object.__setattr__(self, "_server_args", server_args)
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
server_args = object.__getattribute__(self, "_server_args")
|
||||
for _source, declared in reversed(
|
||||
getattr(server_args, "_resolved_overrides", None) or ()
|
||||
):
|
||||
if name in declared:
|
||||
return declared[name]
|
||||
return getattr(server_args, name)
|
||||
|
||||
def __setattr__(self, name: str, value: Any) -> None:
|
||||
raise AttributeError(
|
||||
"ResolvingConfig is read-only; resolution writes through declarations"
|
||||
)
|
||||
|
||||
|
||||
def resolving_view(server_args: Any) -> ResolvingConfig:
|
||||
"""A live read view of what resolution has decided so far."""
|
||||
return ResolvingConfig(server_args)
|
||||
|
||||
|
||||
# Ordered post-process passes (the normalization stage). List order is the
|
||||
# end-state execution order and mirrors today's handler call sequence in
|
||||
# __post_init__; during the transition each pass is invoked from its legacy
|
||||
@@ -581,16 +617,17 @@ def _require_kimi_k3_cutedsl_dcp_support() -> None:
|
||||
|
||||
@_register_for("KimiK3ForConditionalGeneration")
|
||||
def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
if server_args.dcp_size > 1:
|
||||
cfg = resolving_view(server_args)
|
||||
if cfg.dcp_size > 1:
|
||||
overrides = {}
|
||||
if server_args.enable_symm_mem:
|
||||
if cfg.enable_symm_mem:
|
||||
logger.warning(
|
||||
"Kimi-K3 DCP disables --enable-symm-mem due to decode CUDA "
|
||||
"graph correctness issues."
|
||||
)
|
||||
overrides["enable_symm_mem"] = False
|
||||
|
||||
if server_args.speculative_algorithm == "DSPARK":
|
||||
if cfg.speculative_algorithm == "DSPARK":
|
||||
from sglang.srt.speculative.ragged_verify import (
|
||||
RaggedVerifyMode,
|
||||
read_ragged_verify_mode,
|
||||
@@ -611,7 +648,7 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
# lacks that DCP path (TypeError: unexpected kwarg 'causal_seqs').
|
||||
overrides["speculative_attention_mode"] = "decode"
|
||||
|
||||
prefill_backend, decode_backend = attention_backends_of(server_args)
|
||||
prefill_backend, decode_backend = attention_backends_of(cfg)
|
||||
if decode_backend == "cutedsl_mla" or decode_backend is None:
|
||||
_require_kimi_k3_cutedsl_dcp_support()
|
||||
logger.info(
|
||||
@@ -630,7 +667,7 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
)
|
||||
logger.info(
|
||||
"Kimi-K3 DCP with tokenspeed mla backend overrides KV cache dtype: "
|
||||
f"{server_args.kv_cache_dtype!r} -> 'fp8_e4m3'."
|
||||
f"{cfg.kv_cache_dtype!r} -> 'fp8_e4m3'."
|
||||
)
|
||||
overrides.update(
|
||||
prefill_attention_backend="tokenspeed_mla",
|
||||
@@ -642,7 +679,7 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
f"Decode attention backend for Kimi-K3 DCP must be 'cutedsl_mla' or 'tokenspeed_mla', got {decode_backend!r}."
|
||||
)
|
||||
|
||||
if server_args.dcp_replicate_q_proj is None:
|
||||
if cfg.dcp_replicate_q_proj is None:
|
||||
logger.info("Kimi-K3 DCP enables replicated Q projection by default.")
|
||||
overrides["dcp_replicate_q_proj"] = True
|
||||
|
||||
@@ -650,7 +687,7 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
dcp_comm_backend = "fi_a2a" if is_mnnvl_fabric_device() else "a2a"
|
||||
logger.info(
|
||||
"Kimi-K3 DCP selects communication backend on "
|
||||
f"{device_name!r}: {server_args.dcp_comm_backend!r} -> "
|
||||
f"{device_name!r}: {cfg.dcp_comm_backend!r} -> "
|
||||
f"{dcp_comm_backend!r}."
|
||||
)
|
||||
overrides["dcp_comm_backend"] = dcp_comm_backend
|
||||
@@ -659,7 +696,7 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
if not (is_sm100_supported() and get_device_sm() in (100, 103)):
|
||||
return {}
|
||||
backends_unset = server_args.is_attention_backend_not_set()
|
||||
if server_args.speculative_algorithm != "DSPARK":
|
||||
if cfg.speculative_algorithm != "DSPARK":
|
||||
if not backends_unset:
|
||||
return {}
|
||||
logger.info(
|
||||
@@ -673,9 +710,9 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
# DSPARK: verify runs on the decode backend (mode=decode below), so this
|
||||
# picks the verify kernel -- mode=prefill routes it to flashinfer, which is
|
||||
# slow and syncs, while plain decode is cold under dspark.
|
||||
q_len = server_args.speculative_num_draft_tokens or (
|
||||
server_args.speculative_dspark_block_size + 1
|
||||
if server_args.speculative_dspark_block_size is not None
|
||||
q_len = cfg.speculative_num_draft_tokens or (
|
||||
cfg.speculative_dspark_block_size + 1
|
||||
if cfg.speculative_dspark_block_size is not None
|
||||
# Checkpoint auto-infer happens after overrides; K3 draft uses block 7.
|
||||
else 8
|
||||
)
|
||||
@@ -688,8 +725,8 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
# Explicit backend knobs keep priority, but the mode is a separate knob
|
||||
# that still needs declaring -- else verify stays on the prefill backend,
|
||||
# whose host-side plan (flashinfer by default) forces a per-step D2H.
|
||||
_, backend = attention_backends_of(server_args)
|
||||
if _dspark_verify_on_decode_backend(backend, q_len, server_args.kv_cache_dtype):
|
||||
_, backend = attention_backends_of(cfg)
|
||||
if _dspark_verify_on_decode_backend(backend, q_len, cfg.kv_cache_dtype):
|
||||
overrides["speculative_attention_mode"] = "decode"
|
||||
logger.info(
|
||||
"Kimi-K3 DSPARK on SM100/SM103: decode/verify attention backend "
|
||||
@@ -727,7 +764,8 @@ def _kimi_k3_moe_runner_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
# (M=bs) and the target-verify (M=bs*(gamma+1)) regimes on SM100/SM103.
|
||||
# SM107 uses the same packed-MXFP4 runner; leaving auto unresolved falls
|
||||
# back to BF16 weight materialization during model loading.
|
||||
if server_args.moe_runner_backend != "auto":
|
||||
cfg = resolving_view(server_args)
|
||||
if cfg.moe_runner_backend != "auto":
|
||||
return {}
|
||||
if not (is_sm100_supported() and get_device_sm() in (100, 103, 107)):
|
||||
return {}
|
||||
@@ -757,6 +795,7 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
writers), the kv-cache/split-backend defaults, the quant/moe block (read
|
||||
before it by _set_default_dsa_kv_cache_dtype) and the env writes stay in
|
||||
the branch."""
|
||||
cfg = resolving_view(server_args)
|
||||
from sglang.srt.configs.model_config import is_deepseek_dsa
|
||||
|
||||
overrides: Dict[str, Any] = {}
|
||||
@@ -767,39 +806,39 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
overrides["attention_backend"] = "dsa"
|
||||
logger.info("Use dsa attention backend for DeepSeek with DSA.")
|
||||
if not is_npu() and not is_xpu(): # CUDA or ROCm GPU
|
||||
if server_args.enable_prefill_cp:
|
||||
if cfg.enable_prefill_cp:
|
||||
logger.warning(
|
||||
"Context parallel feature is still under experiment. It has only been verified on Hopper platform."
|
||||
)
|
||||
overrides["enable_dp_attention"] = True
|
||||
overrides["moe_dense_tp_size"] = 1
|
||||
if server_args.cp_strategy == "zigzag":
|
||||
if cfg.cp_strategy == "zigzag":
|
||||
overrides["moe_a2a_backend"] = "deepep"
|
||||
overrides["ep_size"] = server_args.tp_size
|
||||
overrides["ep_size"] = cfg.tp_size
|
||||
logger.warning(
|
||||
"zigzag DSA CP requires moe_dense_tp_size=1, "
|
||||
"moe_a2a_backend=deepep, ep_size=tp_size, batch_size=1."
|
||||
)
|
||||
else:
|
||||
assert (
|
||||
server_args.dp_size == 1
|
||||
cfg.dp_size == 1
|
||||
), "interleave DSA CP does not support DP attention."
|
||||
assert (
|
||||
server_args.tp_size <= 8
|
||||
cfg.tp_size <= 8
|
||||
), "Context parallel only supports single machine (tp_size <= 8). Cross-machine CP has precision issues."
|
||||
# Note(kpham-sgl): Keep attn_tp_size == 1 under DSA CP.
|
||||
# DSACPLayerCommunicator does not all-reduce attention-TP
|
||||
# partial o_proj outputs before replicated dense FFNs.
|
||||
attn_cp_size = server_args.tp_size // server_args.dp_size
|
||||
attn_cp_size = cfg.tp_size // cfg.dp_size
|
||||
overrides["attn_cp_size"] = attn_cp_size
|
||||
logger.warning(
|
||||
"Enabled DSA context parallel: "
|
||||
f"strategy={server_args.cp_strategy}, dp_size={server_args.dp_size}, "
|
||||
f"strategy={cfg.cp_strategy}, dp_size={cfg.dp_size}, "
|
||||
f"moe_dense_tp_size={overrides['moe_dense_tp_size']}, "
|
||||
f"ep_size={overrides.get('ep_size', server_args.ep_size)}, tp_size={server_args.tp_size}, "
|
||||
f"ep_size={overrides.get('ep_size', cfg.ep_size)}, tp_size={cfg.tp_size}, "
|
||||
f"attn_cp_size={attn_cp_size}, "
|
||||
f"kv_cache_dtype={server_args.kv_cache_dtype}, "
|
||||
f"moe_a2a_backend={overrides.get('moe_a2a_backend', server_args.moe_a2a_backend)}, "
|
||||
f"kv_cache_dtype={cfg.kv_cache_dtype}, "
|
||||
f"moe_a2a_backend={overrides.get('moe_a2a_backend', cfg.moe_a2a_backend)}, "
|
||||
f"cuda_graph_config[prefill].backend=disabled"
|
||||
)
|
||||
|
||||
@@ -826,9 +865,9 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
# DeepSeek V3/R1/V3.1
|
||||
if is_sm100_supported():
|
||||
if (
|
||||
server_args.attention_backend is None
|
||||
and server_args.prefill_attention_backend is None
|
||||
and server_args.decode_attention_backend is None
|
||||
cfg.attention_backend is None
|
||||
and cfg.prefill_attention_backend is None
|
||||
and cfg.decode_attention_backend is None
|
||||
):
|
||||
overrides["attention_backend"] = "trtllm_mla"
|
||||
logger.info(
|
||||
@@ -836,7 +875,7 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
)
|
||||
# MLA prefill CP auto-config. Mirrors the NSA CP block above
|
||||
# (minus the in-seq/round-robin mode split, which MLA CP does not support)
|
||||
if server_args.enable_prefill_cp and server_args.use_mla_backend():
|
||||
if cfg.enable_prefill_cp and server_args.use_mla_backend():
|
||||
logger.warning(
|
||||
"MLA prefill context parallel is still experimental. "
|
||||
"Verified on Hopper with the fa3 backend."
|
||||
@@ -845,22 +884,22 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
# TODO(kpham-sgl) Supports moe_dense_tp_size != 1.
|
||||
overrides["moe_dense_tp_size"] = 1
|
||||
overrides["moe_a2a_backend"] = "deepep"
|
||||
overrides["ep_size"] = server_args.tp_size
|
||||
overrides["ep_size"] = cfg.tp_size
|
||||
logger.warning(
|
||||
"For MLA CP, we have the following restrictions: moe_dense_tp_size == 1, moe_a2a_backend == deepep, ep_size == tp_size, batch_size == 1"
|
||||
)
|
||||
# FIXME(kpham-sgl): Keep attn_tp_size == 1 under MLA CP.
|
||||
# DSACPLayerCommunicator does not all-reduce attention-TP
|
||||
# partial o_proj outputs before replicated dense FFNs.
|
||||
attn_cp_size = server_args.tp_size // server_args.dp_size
|
||||
attn_cp_size = cfg.tp_size // cfg.dp_size
|
||||
overrides["attn_cp_size"] = attn_cp_size
|
||||
logger.warning(
|
||||
f"Enable Context Parallel opt for MLA, "
|
||||
f"Setting dp_size == {server_args.dp_size} and "
|
||||
f"Setting dp_size == {cfg.dp_size} and "
|
||||
f"attn_cp_size == {attn_cp_size}, "
|
||||
f"moe_dense_tp_size == {overrides['moe_dense_tp_size']}, "
|
||||
f"ep_size == {overrides['ep_size']}, "
|
||||
f"tp_size == {server_args.tp_size}, "
|
||||
f"tp_size == {cfg.tp_size}, "
|
||||
f"moe_a2a_backend {overrides['moe_a2a_backend']}, "
|
||||
f"cuda_graph_config[prefill].backend=disabled"
|
||||
)
|
||||
@@ -870,8 +909,9 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
# Keep in sync with MIMO_V2_MODEL_ARCHS (server_args.py / configs/hf_config.py).
|
||||
@_register_for("MiMoV2ForCausalLM", "MiMoV2FlashForCausalLM")
|
||||
def _mimo_v2_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
cfg = resolving_view(server_args)
|
||||
overrides: Dict[str, Any] = {}
|
||||
if server_args.speculative_algorithm == "EAGLE":
|
||||
if cfg.speculative_algorithm == "EAGLE":
|
||||
logger.info("Enable multi-layer EAGLE speculative decoding for MiMoV2 model.")
|
||||
overrides["enable_multi_layer_eagle"] = True
|
||||
|
||||
@@ -879,7 +919,7 @@ def _mimo_v2_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
# slower at bs=1 decode. FP4 checkpoints use flashinfer_mxfp4 instead.
|
||||
if (
|
||||
is_sm100_supported()
|
||||
and server_args.moe_runner_backend == "auto"
|
||||
and cfg.moe_runner_backend == "auto"
|
||||
and get_quantization_config(hf_config) == "fp8"
|
||||
):
|
||||
overrides["moe_runner_backend"] = "flashinfer_trtllm"
|
||||
@@ -889,13 +929,14 @@ def _mimo_v2_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
|
||||
@_register_for("MiniMaxM2ForCausalLM")
|
||||
def _minimax_m2_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
cfg = resolving_view(server_args)
|
||||
overrides = {"enable_tf32_matmul": True}
|
||||
logger.info(
|
||||
"Enable TF32 matmul for MiniMaxM2ForCausalLM model to improve gate gemm performance."
|
||||
)
|
||||
if (
|
||||
is_sm100_supported()
|
||||
and server_args.moe_runner_backend == "auto"
|
||||
and cfg.moe_runner_backend == "auto"
|
||||
and server_args.get_model_config().quantization == "modelopt_fp4"
|
||||
):
|
||||
overrides["moe_runner_backend"] = "flashinfer_trtllm_routed"
|
||||
@@ -909,10 +950,11 @@ def _minimax_m2_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
@_register_for("MiniMaxM3SparseForCausalLM", "MiniMaxM3SparseForConditionalGeneration")
|
||||
def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
|
||||
cfg = resolving_view(server_args)
|
||||
overrides: Dict[str, Any] = {}
|
||||
|
||||
quant_method = get_quantization_config(hf_config)
|
||||
quant_resolved = server_args.quantization
|
||||
quant_resolved = cfg.quantization
|
||||
if (
|
||||
quant_resolved is None
|
||||
and not server_args._quantization_explicitly_unset
|
||||
@@ -924,16 +966,12 @@ def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
if is_hip():
|
||||
if server_args.is_attention_backend_not_set():
|
||||
overrides["attention_backend"] = "triton"
|
||||
if server_args.moe_runner_backend == "auto" and quant_resolved == "mxfp8":
|
||||
if cfg.moe_runner_backend == "auto" and quant_resolved == "mxfp8":
|
||||
overrides["moe_runner_backend"] = "triton"
|
||||
if not envs.USE_ROCM_AITER_ROPE_BACKEND.is_set():
|
||||
envs.USE_ROCM_AITER_ROPE_BACKEND.set("0")
|
||||
aiter_fusion_resolved = server_args.enable_aiter_allreduce_fusion
|
||||
if (
|
||||
server_args.ep_size > 1
|
||||
and server_args.moe_a2a_backend == "none"
|
||||
and aiter_fusion_resolved
|
||||
):
|
||||
aiter_fusion_resolved = cfg.enable_aiter_allreduce_fusion
|
||||
if cfg.ep_size > 1 and cfg.moe_a2a_backend == "none" and aiter_fusion_resolved:
|
||||
logger.warning(
|
||||
"Disable --enable-aiter-allreduce-fusion for MiniMax-M3 "
|
||||
"standard EP on ROCm because the deferred fused all-reduce "
|
||||
@@ -951,7 +989,7 @@ def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
elif is_sm100_supported():
|
||||
if server_args.is_attention_backend_not_set():
|
||||
if (
|
||||
server_args.kv_cache_dtype == "fp8_e4m3"
|
||||
cfg.kv_cache_dtype == "fp8_e4m3"
|
||||
and not envs.SGLANG_DISABLE_M3_FP8_ATTN_GEMM.get()
|
||||
):
|
||||
# fp8 attention GEMMs activate whenever possible
|
||||
@@ -962,42 +1000,36 @@ def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
overrides["attention_backend"] = "trtllm_mha"
|
||||
else:
|
||||
overrides["attention_backend"] = "fa4"
|
||||
backend_resolved = overrides.get(
|
||||
"attention_backend", server_args.attention_backend
|
||||
)
|
||||
page_resolved = server_args.page_size
|
||||
backend_resolved = overrides.get("attention_backend", cfg.attention_backend)
|
||||
page_resolved = cfg.page_size
|
||||
# fa4 (fmha_sm100) and trtllm_mha both allow the page_size == 128
|
||||
# sparse block MSA needs (trtllm_mha via trtllm-gen's dynamic
|
||||
# tokens-per-page kernels).
|
||||
if page_resolved is None and backend_resolved in ("fa4", "trtllm_mha"):
|
||||
overrides["page_size"] = 128
|
||||
page_resolved = 128
|
||||
if server_args.moe_runner_backend == "auto" and quant_resolved == "mxfp8":
|
||||
if cfg.moe_runner_backend == "auto" and quant_resolved == "mxfp8":
|
||||
overrides["moe_runner_backend"] = "deep_gemm"
|
||||
elif (
|
||||
server_args.moe_runner_backend == "auto"
|
||||
and quant_resolved == "modelopt_mixed"
|
||||
):
|
||||
elif cfg.moe_runner_backend == "auto" and quant_resolved == "modelopt_mixed":
|
||||
overrides["moe_runner_backend"] = "flashinfer_trtllm_routed"
|
||||
logger.info(
|
||||
"MiniMax-M3 on SM100: attention_backend="
|
||||
f"{overrides.get('attention_backend', server_args.attention_backend)}, page_size={page_resolved}, "
|
||||
f"moe_runner_backend={overrides.get('moe_runner_backend', server_args.moe_runner_backend)}."
|
||||
f"{overrides.get('attention_backend', cfg.attention_backend)}, page_size={page_resolved}, "
|
||||
f"moe_runner_backend={overrides.get('moe_runner_backend', cfg.moe_runner_backend)}."
|
||||
)
|
||||
elif is_sm90_supported():
|
||||
if server_args.is_attention_backend_not_set():
|
||||
overrides["attention_backend"] = "fa3"
|
||||
page_resolved = server_args.page_size
|
||||
page_resolved = cfg.page_size
|
||||
if (
|
||||
page_resolved is None
|
||||
and overrides.get("attention_backend", server_args.attention_backend)
|
||||
== "fa3"
|
||||
and overrides.get("attention_backend", cfg.attention_backend) == "fa3"
|
||||
):
|
||||
overrides["page_size"] = 128
|
||||
page_resolved = 128
|
||||
logger.info(
|
||||
"MiniMax-M3 on Hopper: attention_backend="
|
||||
f"{overrides.get('attention_backend', server_args.attention_backend)}, page_size={page_resolved} "
|
||||
f"{overrides.get('attention_backend', cfg.attention_backend)}, page_size={page_resolved} "
|
||||
"(MSA is SM100-only; sparse attention runs on the Triton path)."
|
||||
)
|
||||
|
||||
@@ -1008,7 +1040,7 @@ def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
# silently dispatch the e4m3 kernel, so e5m2 stays on the widening Triton
|
||||
# path), log when the fp8 GEMM mode is active, and log when the
|
||||
# SGLANG_DISABLE_M3_FP8_ATTN_GEMM kill switch suppresses it.
|
||||
if server_args.kv_cache_dtype == "fp8_e5m2":
|
||||
if cfg.kv_cache_dtype == "fp8_e5m2":
|
||||
logger.warning(
|
||||
"MiniMax-M3 with kv_cache_dtype fp8_e5m2: fp8 attention GEMMs stay "
|
||||
"DISABLED (fmha_sm100's variant lookup would silently dispatch the "
|
||||
@@ -1016,9 +1048,8 @@ def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
"Triton path. Use --kv-cache-dtype fp8_e4m3 for fp8 attention GEMMs."
|
||||
)
|
||||
elif (
|
||||
server_args.kv_cache_dtype == "fp8_e4m3"
|
||||
and overrides.get("attention_backend", server_args.attention_backend)
|
||||
== "trtllm_mha"
|
||||
cfg.kv_cache_dtype == "fp8_e4m3"
|
||||
and overrides.get("attention_backend", cfg.attention_backend) == "trtllm_mha"
|
||||
and is_sm100_supported()
|
||||
):
|
||||
if envs.SGLANG_DISABLE_M3_FP8_ATTN_GEMM.get():
|
||||
@@ -1036,9 +1067,7 @@ def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
"force the pre-fp8 numerics."
|
||||
)
|
||||
|
||||
moe_runner_resolved = overrides.get(
|
||||
"moe_runner_backend", server_args.moe_runner_backend
|
||||
)
|
||||
moe_runner_resolved = overrides.get("moe_runner_backend", cfg.moe_runner_backend)
|
||||
if quant_resolved is None and moe_runner_resolved in ("auto", "deep_gemm"):
|
||||
if moe_runner_resolved == "deep_gemm":
|
||||
logger.warning(
|
||||
@@ -1078,6 +1107,7 @@ def _exaone_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
|
||||
@_register_for("GptOssForCausalLM")
|
||||
def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
cfg = resolving_view(server_args)
|
||||
overrides: Dict[str, Any] = {}
|
||||
# Set attention backend for GPT-OSS
|
||||
if server_args.is_attention_backend_not_set():
|
||||
@@ -1100,14 +1130,14 @@ def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
# Check for bf16 dtype on Intel XPU. Reads the pristine dtype request,
|
||||
# which equals the legacy mid-branch read: dtype had no earlier writer
|
||||
# for this arch.
|
||||
if server_args.dtype == "auto":
|
||||
if cfg.dtype == "auto":
|
||||
logger.warning(
|
||||
"GptOssForCausalLM on Intel XPU currently supports bfloat16 dtype only"
|
||||
)
|
||||
elif server_args.dtype not in ["bfloat16"]:
|
||||
elif cfg.dtype not in ["bfloat16"]:
|
||||
raise NotImplementedError(
|
||||
f"GptOssForCausalLM on Intel XPU only supports bfloat16 dtype, "
|
||||
f"but got '{server_args.dtype}'. Please use --dtype bfloat16 or remove --dtype to use auto."
|
||||
f"but got '{cfg.dtype}'. Please use --dtype bfloat16 or remove --dtype to use auto."
|
||||
)
|
||||
quantization_config = getattr(hf_config, "quantization_config", None)
|
||||
is_mxfp4_quant_format = (
|
||||
@@ -1117,7 +1147,7 @@ def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
if is_mxfp4_quant_format:
|
||||
# use bf16 for mxfp4 triton kernels
|
||||
overrides["dtype"] = "bfloat16"
|
||||
if server_args.moe_runner_backend == "auto":
|
||||
if cfg.moe_runner_backend == "auto":
|
||||
|
||||
if is_sm100_supported() and is_mxfp4_quant_format:
|
||||
overrides["moe_runner_backend"] = "flashinfer_mxfp4"
|
||||
@@ -1155,9 +1185,9 @@ def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
"Detected MUSA with SGLANG_DEEPEP_BF16_DISPATCH for bf16 model, using deep_gemm kernel."
|
||||
)
|
||||
elif (
|
||||
server_args.ep_size == 1
|
||||
cfg.ep_size == 1
|
||||
and is_triton_kernels_available()
|
||||
and server_args.quantization is None
|
||||
and cfg.quantization is None
|
||||
and not (is_cpu() and cpu_has_amx_support())
|
||||
):
|
||||
# The triton_kernels package segfaults on Blackwell (B200)
|
||||
@@ -1179,18 +1209,19 @@ def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
# Keep in sync with LLAMA4_MODEL_ARCHS (server_args.py).
|
||||
@_register_for("Llama4ForConditionalGeneration", "Llama4ForCausalLM")
|
||||
def _llama4_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
if server_args.device == "cpu":
|
||||
cfg = resolving_view(server_args)
|
||||
if cfg.device == "cpu":
|
||||
return {}
|
||||
overrides: Dict[str, Any] = {}
|
||||
# Auto-select attention backend for Llama4 if not specified
|
||||
if server_args.attention_backend is None:
|
||||
if cfg.attention_backend is None:
|
||||
if is_sm100_supported():
|
||||
backend, platform = "trtllm_mha", "sm100"
|
||||
elif is_sm90_supported():
|
||||
backend, platform = "fa3", "sm90"
|
||||
elif is_hip():
|
||||
backend, platform = "aiter", "hip"
|
||||
elif server_args.device == "xpu":
|
||||
elif cfg.device == "xpu":
|
||||
backend, platform = "intel_xpu", "xpu"
|
||||
else:
|
||||
backend, platform = "triton", "other platforms"
|
||||
@@ -1198,8 +1229,8 @@ def _llama4_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
f"Use {backend} as attention backend on {platform} for Llama4 model"
|
||||
)
|
||||
overrides["attention_backend"] = backend
|
||||
if is_sm100_supported() and server_args.moe_runner_backend == "auto":
|
||||
if server_args.quantization in {"fp8", "modelopt_fp8"}:
|
||||
if is_sm100_supported() and cfg.moe_runner_backend == "auto":
|
||||
if cfg.quantization in {"fp8", "modelopt_fp8"}:
|
||||
overrides["moe_runner_backend"] = "flashinfer_trtllm"
|
||||
logger.info(
|
||||
"Use flashinfer_trtllm as MoE runner backend on SM100 for Llama4"
|
||||
@@ -1213,6 +1244,7 @@ def _llama4_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
"Gemma4UnifiedForConditionalGeneration",
|
||||
)
|
||||
def _gemma4_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
cfg = resolving_view(server_args)
|
||||
overrides: Dict[str, Any] = {}
|
||||
default_attention_backend = "trtllm_mha" if is_sm100_supported() else "triton"
|
||||
if server_args.is_attention_backend_not_set():
|
||||
@@ -1223,9 +1255,9 @@ def _gemma4_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
# If only one split backend is set, keep the other side on a
|
||||
# Gemma4-compatible fallback instead of letting generic backend selection
|
||||
# choose an unsupported backend later.
|
||||
elif server_args.attention_backend is None:
|
||||
elif cfg.attention_backend is None:
|
||||
overrides["attention_backend"] = default_attention_backend
|
||||
if is_sm100_supported() and server_args.moe_runner_backend == "auto":
|
||||
if is_sm100_supported() and cfg.moe_runner_backend == "auto":
|
||||
if server_args.get_model_config().quantization == "modelopt_fp4":
|
||||
overrides["quantization"] = "modelopt_fp4"
|
||||
overrides["moe_runner_backend"] = "flashinfer_trtllm"
|
||||
@@ -1255,7 +1287,8 @@ def _moss_vl_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
|
||||
@_register_for("MiniCPMForCausalLM", "MiniCPMSALAForCausalLM")
|
||||
def _minicpm_sala_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
if server_args.enable_dp_attention:
|
||||
cfg = resolving_view(server_args)
|
||||
if cfg.enable_dp_attention:
|
||||
raise ValueError("MiniCPM does not support DP attention")
|
||||
has_sparse_attention = getattr(hf_config, "has_minicpm_sparse_attention", False)
|
||||
has_hybrid_attention = has_sparse_attention or getattr(
|
||||
@@ -1263,7 +1296,7 @@ def _minicpm_sala_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
)
|
||||
overrides: Dict[str, Any] = {}
|
||||
if has_hybrid_attention:
|
||||
if server_args.enable_hierarchical_cache:
|
||||
if cfg.enable_hierarchical_cache:
|
||||
raise ValueError("MiniCPM SALA does not support hierarchical cache")
|
||||
overrides["disable_radix_cache"] = True
|
||||
if envs.SGLANG_MINICPM_FORCE_DENSE.get():
|
||||
@@ -1273,29 +1306,29 @@ def _minicpm_sala_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
}
|
||||
# Literal keys keep the written-field set statically derivable; a loop
|
||||
# variable hides it from the census in test_chain_read_ratchet.py.
|
||||
dense_attention = dense_backends.get(server_args.attention_backend)
|
||||
dense_attention = dense_backends.get(cfg.attention_backend)
|
||||
if dense_attention is not None:
|
||||
overrides["attention_backend"] = dense_attention
|
||||
dense_prefill = dense_backends.get(server_args.prefill_attention_backend)
|
||||
dense_prefill = dense_backends.get(cfg.prefill_attention_backend)
|
||||
if dense_prefill is not None:
|
||||
overrides["prefill_attention_backend"] = dense_prefill
|
||||
dense_decode = dense_backends.get(server_args.decode_attention_backend)
|
||||
dense_decode = dense_backends.get(cfg.decode_attention_backend)
|
||||
if dense_decode is not None:
|
||||
overrides["decode_attention_backend"] = dense_decode
|
||||
elif has_sparse_attention:
|
||||
uses_sparse_backend = server_args.is_attention_backend_not_set() or any(
|
||||
uses_sparse_backend = cfg.is_attention_backend_not_set() or any(
|
||||
backend in ("minicpm_flashattn", "minicpm_flashinfer")
|
||||
for backend in (
|
||||
server_args.attention_backend,
|
||||
server_args.prefill_attention_backend,
|
||||
server_args.decode_attention_backend,
|
||||
cfg.attention_backend,
|
||||
cfg.prefill_attention_backend,
|
||||
cfg.decode_attention_backend,
|
||||
)
|
||||
)
|
||||
if uses_sparse_backend and server_args.disaggregation_mode != "null":
|
||||
if uses_sparse_backend and cfg.disaggregation_mode != "null":
|
||||
raise ValueError(
|
||||
"MiniCPM sparse attention does not support PD disaggregation"
|
||||
)
|
||||
if server_args.is_attention_backend_not_set():
|
||||
if cfg.is_attention_backend_not_set():
|
||||
overrides["attention_backend"] = (
|
||||
"minicpm_flashinfer"
|
||||
if is_blackwell_supported()
|
||||
@@ -1306,7 +1339,8 @@ def _minicpm_sala_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
|
||||
@_register_for("MiniCPMV4_6ForConditionalGeneration")
|
||||
def _minicpm_v4_6_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
if is_sm100_supported() and server_args.attention_backend is None:
|
||||
cfg = resolving_view(server_args)
|
||||
if is_sm100_supported() and cfg.attention_backend is None:
|
||||
return {"attention_backend": "triton"}
|
||||
return {}
|
||||
|
||||
@@ -1315,24 +1349,27 @@ def _minicpm_v4_6_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
"FalconH1ForCausalLM", "JetNemotronForCausalLM", "JetVLMForConditionalGeneration"
|
||||
)
|
||||
def _falcon_h1_jet_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
if is_sm100_supported() and server_args.attention_backend is None:
|
||||
cfg = resolving_view(server_args)
|
||||
if is_sm100_supported() and cfg.attention_backend is None:
|
||||
return {"attention_backend": "triton"}
|
||||
return {}
|
||||
|
||||
|
||||
@_register_for("GraniteMoeHybridForCausalLM")
|
||||
def _granite_moe_hybrid_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
cfg = resolving_view(server_args)
|
||||
has_mamba = any(
|
||||
layer_type == "mamba" for layer_type in getattr(hf_config, "layer_types", [])
|
||||
)
|
||||
if has_mamba and is_sm100_supported() and server_args.attention_backend is None:
|
||||
if has_mamba and is_sm100_supported() and cfg.attention_backend is None:
|
||||
return {"attention_backend": "flashinfer"}
|
||||
return {}
|
||||
|
||||
|
||||
@_register_for("Lfm2ForCausalLM", "Lfm2MoeForCausalLM")
|
||||
def _lfm2_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
if is_sm100_supported() and server_args.attention_backend is None:
|
||||
cfg = resolving_view(server_args)
|
||||
if is_sm100_supported() and cfg.attention_backend is None:
|
||||
return {"attention_backend": "flashinfer"}
|
||||
return {}
|
||||
|
||||
@@ -1343,13 +1380,14 @@ def _deepseek_v4_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
arg_groups/deepseek_v4_hook.py). The kv-cache dtype and NPU split-backend
|
||||
writes, the max_running_requests fill and the validations stay in the
|
||||
hook at its legacy slot."""
|
||||
cfg = resolving_view(server_args)
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
model_arch = hf_config.architectures[0]
|
||||
overrides: Dict[str, Any] = {"attention_backend": "dsv4"}
|
||||
|
||||
page_size = 256
|
||||
if server_args.device == "npu":
|
||||
if cfg.device == "npu":
|
||||
# NPU keeps the device-aware "dsv4" backend (the registry routes it to
|
||||
# the Ascend V4 subclass); only the pool geometry / dtype differ.
|
||||
# set_default_server_args() pins all three backends to "ascend" for
|
||||
@@ -1363,11 +1401,11 @@ def _deepseek_v4_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
f"Use dsv4 attention backend for {model_arch}, setting page_size to {page_size}."
|
||||
)
|
||||
|
||||
if server_args.swa_full_tokens_ratio == ServerArgs.swa_full_tokens_ratio:
|
||||
if cfg.swa_full_tokens_ratio == ServerArgs.swa_full_tokens_ratio:
|
||||
overrides["swa_full_tokens_ratio"] = 0.1
|
||||
logger.info(f"Setting swa_full_tokens_ratio to 0.1 for {model_arch}.")
|
||||
|
||||
if server_args.moe_runner_backend == "auto":
|
||||
if cfg.moe_runner_backend == "auto":
|
||||
model_config = server_args.get_model_config()
|
||||
# nvidia/DeepSeek-V4-Pro-NVFP4 uses the routed TRT-LLM runner.
|
||||
if model_config.nvfp4_moe_meta is not None:
|
||||
@@ -1377,9 +1415,9 @@ def _deepseek_v4_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
f"{model_arch} hybrid FP8+NVFP4 checkpoint."
|
||||
)
|
||||
elif (
|
||||
server_args.device == "cuda"
|
||||
cfg.device == "cuda"
|
||||
and not is_hip()
|
||||
and server_args.moe_a2a_backend == "none"
|
||||
and cfg.moe_a2a_backend == "none"
|
||||
and not envs.SGLANG_DSV4_FP4_DEQUANT.get()
|
||||
and model_config.is_fp4_experts
|
||||
and (is_sm90_supported() or is_sm100_supported() or is_sm120_supported())
|
||||
@@ -1407,6 +1445,7 @@ def _inkling_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
prefill.backend, and an explicit --cuda-graph-backend-prefill /
|
||||
--disable-prefill-cuda-graph still wins. The unified-radix env write follows
|
||||
the MiniMax-M3 handler precedent (env is not a resolvable server-arg)."""
|
||||
cfg = resolving_view(server_args)
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
overrides: Dict[str, Any] = {}
|
||||
@@ -1415,14 +1454,14 @@ def _inkling_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
# cuda_graph_backend_prefill declared here lands too late (the breakable
|
||||
# default would already have been auto-disabled for this multimodal arch).
|
||||
# It is set inline before _handle_cuda_graph_config instead.
|
||||
if server_args.swa_full_tokens_ratio == ServerArgs.swa_full_tokens_ratio:
|
||||
if cfg.swa_full_tokens_ratio == ServerArgs.swa_full_tokens_ratio:
|
||||
overrides["swa_full_tokens_ratio"] = 0.1
|
||||
if server_args.mamba_full_memory_ratio == ServerArgs.mamba_full_memory_ratio:
|
||||
if cfg.mamba_full_memory_ratio == ServerArgs.mamba_full_memory_ratio:
|
||||
overrides["mamba_full_memory_ratio"] = 0.1
|
||||
# Inkling requires the extra-buffer mamba strategy (inkling.py asserts
|
||||
# enable_mamba_extra_buffer()); the generic "auto" resolution does not cover
|
||||
# Inkling, so pin it here. Yields to an explicit --mamba-scheduler-strategy.
|
||||
if server_args.mamba_radix_cache_strategy == ServerArgs.mamba_radix_cache_strategy:
|
||||
if cfg.mamba_radix_cache_strategy == ServerArgs.mamba_radix_cache_strategy:
|
||||
overrides["mamba_radix_cache_strategy"] = "extra_buffer"
|
||||
# Inkling attention runs only on the fa4 (Blackwell) or triton backends --
|
||||
# models/inkling_common/attn.py asserts attention_backend in {fa4, triton}.
|
||||
@@ -1447,6 +1486,7 @@ def _nemotron_h_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
"""NemotronH quantization / MoE runner / attention backend defaults
|
||||
(absorbed from the retired arg_groups/nemotron_h_hook.py; the mamba radix
|
||||
cache handling and the triton-backend assert stay in the arch branch)."""
|
||||
cfg = resolving_view(server_args)
|
||||
model_arch = hf_config.architectures[0]
|
||||
model_config = server_args.get_model_config()
|
||||
overrides: Dict[str, Any] = {}
|
||||
@@ -1457,7 +1497,7 @@ def _nemotron_h_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
"modelopt_fp4",
|
||||
"modelopt_mixed",
|
||||
]
|
||||
quantization = server_args.quantization
|
||||
quantization = cfg.quantization
|
||||
if is_modelopt:
|
||||
assert model_config.hf_config.mlp_hidden_act == "relu2"
|
||||
if model_config.quantization == "modelopt":
|
||||
@@ -1482,22 +1522,22 @@ def _nemotron_h_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
)
|
||||
|
||||
if has_w4a16_moe_layers:
|
||||
if server_args.moe_a2a_backend != "none":
|
||||
if cfg.moe_a2a_backend != "none":
|
||||
raise ValueError("W4A16_NVFP4 MoE layers require --moe-a2a-backend=none.")
|
||||
if server_args.moe_runner_backend not in ("auto", "marlin"):
|
||||
if cfg.moe_runner_backend not in ("auto", "marlin"):
|
||||
raise ValueError(
|
||||
"W4A16_NVFP4 MoE layers require --moe-runner-backend=marlin."
|
||||
)
|
||||
if server_args.moe_runner_backend == "auto":
|
||||
if cfg.moe_runner_backend == "auto":
|
||||
overrides["moe_runner_backend"] = "marlin"
|
||||
logger.info(
|
||||
"Use marlin as MoE runner backend for "
|
||||
f"{model_arch} with W4A16_NVFP4 MoE layers"
|
||||
)
|
||||
elif (is_modelopt or model_config.quantization is None) and (
|
||||
server_args.moe_runner_backend == "auto"
|
||||
cfg.moe_runner_backend == "auto"
|
||||
):
|
||||
if is_sm100_supported() and server_args.moe_a2a_backend == "none":
|
||||
if is_sm100_supported() and cfg.moe_a2a_backend == "none":
|
||||
overrides["moe_runner_backend"] = "flashinfer_trtllm"
|
||||
logger.info(
|
||||
f"Use flashinfer_trtllm as MoE runner backend on sm100 for {model_arch}"
|
||||
@@ -1518,27 +1558,27 @@ def _nemotron_h_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
else:
|
||||
overrides["moe_runner_backend"] = "flashinfer_cutlass"
|
||||
|
||||
if is_blackwell_supported() and server_args.is_attention_backend_not_set():
|
||||
if server_args.speculative_algorithm is not None:
|
||||
speculative_algorithm = server_args.speculative_algorithm.upper()
|
||||
if is_sm100_supported() and server_args.speculative_eagle_topk in (
|
||||
if is_blackwell_supported() and cfg.is_attention_backend_not_set():
|
||||
if cfg.speculative_algorithm is not None:
|
||||
speculative_algorithm = cfg.speculative_algorithm.upper()
|
||||
if is_sm100_supported() and cfg.speculative_eagle_topk in (
|
||||
None,
|
||||
1,
|
||||
):
|
||||
overrides["attention_backend"] = "trtllm_mha"
|
||||
if server_args.page_size is None:
|
||||
if cfg.page_size is None:
|
||||
overrides["page_size"] = 64
|
||||
if server_args.mamba_radix_cache_strategy == "auto":
|
||||
if cfg.mamba_radix_cache_strategy == "auto":
|
||||
overrides["mamba_radix_cache_strategy"] = "extra_buffer"
|
||||
if (
|
||||
server_args.speculative_draft_attention_backend is None
|
||||
cfg.speculative_draft_attention_backend is None
|
||||
and speculative_algorithm in ("EAGLE", "NEXTN", "DSPARK")
|
||||
):
|
||||
overrides["speculative_draft_attention_backend"] = "trtllm_mha"
|
||||
else:
|
||||
overrides["attention_backend"] = "triton"
|
||||
if (
|
||||
server_args.speculative_draft_attention_backend is None
|
||||
cfg.speculative_draft_attention_backend is None
|
||||
and speculative_algorithm in ("EAGLE", "NEXTN", "DFLASH", "DSPARK")
|
||||
):
|
||||
overrides["speculative_draft_attention_backend"] = "flashinfer"
|
||||
@@ -1555,7 +1595,8 @@ def _nemotron_h_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
"Qwen3_5ForConditionalGeneration",
|
||||
)
|
||||
def _qwen3_5_hybrid_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
if not is_sm100_supported() or server_args.attention_backend is not None:
|
||||
cfg = resolving_view(server_args)
|
||||
if not is_sm100_supported() or cfg.attention_backend is not None:
|
||||
return {}
|
||||
sm100_default_attn_backend = "triton"
|
||||
# trtllm_mha requires speculative_eagle_topk == 1 and page_size > 1.
|
||||
@@ -1572,8 +1613,8 @@ def _qwen3_5_hybrid_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
# already-written field here).
|
||||
if default_attn_backend == "trtllm_mha" and not (
|
||||
not mamba_extra_buffer_of(resolved_view(server_args))
|
||||
and not server_args.disable_radix_cache
|
||||
and server_args.speculative_algorithm is None
|
||||
and not cfg.disable_radix_cache
|
||||
and cfg.speculative_algorithm is None
|
||||
):
|
||||
sm100_default_attn_backend = "trtllm_mha"
|
||||
return {
|
||||
@@ -1585,7 +1626,8 @@ def _qwen3_5_hybrid_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
@_register_for("InternS2MobiusForConditionalGeneration")
|
||||
def _interns2_mobius_baseline_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
"""Select the only MoE runner validated for the 2,560-expert baseline."""
|
||||
if server_args.moe_runner_backend == "auto":
|
||||
cfg = resolving_view(server_args)
|
||||
if cfg.moe_runner_backend == "auto":
|
||||
return {"moe_runner_backend": "triton_kernel"}
|
||||
return {}
|
||||
|
||||
@@ -1593,11 +1635,8 @@ def _interns2_mobius_baseline_overrides(server_args: Any, hf_config: Any) -> dic
|
||||
@_register_for("Qwen3VLForConditionalGeneration")
|
||||
def _qwen3vl_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
|
||||
if (
|
||||
is_hip()
|
||||
and envs.SGLANG_USE_AITER_UNIFIED_ATTN.get()
|
||||
and server_args.page_size is None
|
||||
):
|
||||
cfg = resolving_view(server_args)
|
||||
if is_hip() and envs.SGLANG_USE_AITER_UNIFIED_ATTN.get() and cfg.page_size is None:
|
||||
logger.info(
|
||||
"Setting page_size=16 for aiter unified attention on Qwen3VLForConditionalGeneration."
|
||||
)
|
||||
@@ -1614,10 +1653,11 @@ def _qwen3vl_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
"Qwen3_5ForConditionalGeneration",
|
||||
)
|
||||
def _qwen3_moe_family_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
cfg = resolving_view(server_args)
|
||||
overrides: Dict[str, Any] = {}
|
||||
if is_sm100_supported():
|
||||
quant_method = get_quantization_config(hf_config)
|
||||
quantization = server_args.quantization
|
||||
quantization = cfg.quantization
|
||||
if (
|
||||
quantization is None
|
||||
and not server_args._quantization_explicitly_unset
|
||||
@@ -1627,8 +1667,8 @@ def _qwen3_moe_family_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
quantization = quant_method
|
||||
if (
|
||||
(quantization in ("fp8", "modelopt_fp4") or quantization is None)
|
||||
and server_args.moe_a2a_backend == "none"
|
||||
and server_args.moe_runner_backend == "auto"
|
||||
and cfg.moe_a2a_backend == "none"
|
||||
and cfg.moe_runner_backend == "auto"
|
||||
):
|
||||
overrides["moe_runner_backend"] = "flashinfer_trtllm"
|
||||
logger.info(
|
||||
@@ -1640,6 +1680,7 @@ def _qwen3_moe_family_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
|
||||
@_register_for("Glm4MoeForCausalLM")
|
||||
def _glm4_moe_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
cfg = resolving_view(server_args)
|
||||
overrides: Dict[str, Any] = {}
|
||||
if is_sm100_supported():
|
||||
quantization_config = getattr(hf_config, "quantization_config", None)
|
||||
@@ -1648,7 +1689,7 @@ def _glm4_moe_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
if quantization_config is not None
|
||||
else None
|
||||
)
|
||||
quantization = server_args.quantization
|
||||
quantization = cfg.quantization
|
||||
if (
|
||||
quantization is None
|
||||
and not server_args._quantization_explicitly_unset
|
||||
@@ -1658,8 +1699,8 @@ def _glm4_moe_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
quantization = quant_method
|
||||
if (
|
||||
quantization in {"modelopt_fp4", None}
|
||||
and server_args.moe_a2a_backend == "none"
|
||||
and server_args.moe_runner_backend == "auto"
|
||||
and cfg.moe_a2a_backend == "none"
|
||||
and cfg.moe_runner_backend == "auto"
|
||||
):
|
||||
overrides["moe_runner_backend"] = "flashinfer_trtllm"
|
||||
logger.info(
|
||||
@@ -1674,13 +1715,14 @@ def _glm4_moe_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
|
||||
@_register_for("Olmo2ForCausalLM")
|
||||
def _olmo2_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
cfg = resolving_view(server_args)
|
||||
overrides: Dict[str, Any] = {}
|
||||
# FIXME: https://github.com/sgl-project/sglang/pull/7367 is not compatible with Olmo3 model.
|
||||
logger.warning(
|
||||
f"Disabling hybrid SWA memory for {hf_config.architectures[0]} as it is not yet supported."
|
||||
)
|
||||
overrides["disable_hybrid_swa_memory"] = True
|
||||
if server_args.attention_backend is None:
|
||||
if cfg.attention_backend is None:
|
||||
if is_cuda() and is_sm100_supported():
|
||||
overrides["attention_backend"] = "trtllm_mha"
|
||||
elif is_cuda() and get_device_sm() >= 80:
|
||||
@@ -1695,6 +1737,7 @@ def _olmo2_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
or "Step3p7ForConditionalGeneration" in arch
|
||||
)
|
||||
def _step3p_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
cfg = resolving_view(server_args)
|
||||
overrides: Dict[str, Any] = {}
|
||||
if server_args.is_attention_backend_not_set():
|
||||
if is_blackwell_supported():
|
||||
@@ -1703,12 +1746,12 @@ def _step3p_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
elif is_sm90_supported():
|
||||
logger.info("Auto-select fa3 attention backend for Step3p7 on Hopper.")
|
||||
overrides["attention_backend"] = "fa3"
|
||||
if server_args.speculative_algorithm == "EAGLE":
|
||||
if cfg.speculative_algorithm == "EAGLE":
|
||||
logger.info(
|
||||
"Enable multi-layer EAGLE speculative decoding for Step3p5ForCausalLM model."
|
||||
)
|
||||
overrides["enable_multi_layer_eagle"] = True
|
||||
if server_args.enable_hierarchical_cache:
|
||||
if cfg.enable_hierarchical_cache:
|
||||
logger.warning(
|
||||
"Reset swa_full_tokens_ratio to 1.0 for Step3p5ForCausalLM model with hierarchical cache"
|
||||
)
|
||||
@@ -2170,7 +2213,8 @@ def _deepseek_v4_kv_cache_dtype(view: Any) -> dict:
|
||||
|
||||
@_register_for("MuseGlimmerForConditionalGeneration", "MuseGlimmerForCausalLM")
|
||||
def _muse_glimmer_fp4_gemm_runner_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
if is_sm120_supported() and server_args.fp4_gemm_runner_backend == "auto":
|
||||
cfg = resolving_view(server_args)
|
||||
if is_sm120_supported() and cfg.fp4_gemm_runner_backend == "auto":
|
||||
logger.info("Use marlin as FP4 GEMM runner backend on SM120 for Muse Glimmer")
|
||||
return {"fp4_gemm_runner_backend": "marlin"}
|
||||
return {}
|
||||
|
||||
@@ -5,7 +5,10 @@ import logging
|
||||
import os
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sglang.srt.arg_groups.overrides import declare_resolution
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
declare_resolution,
|
||||
resolving_view,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -16,10 +19,11 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
def handle_pd_disaggregation(server_args: ServerArgs) -> None:
|
||||
"""Validate and normalize PD-disaggregation server args."""
|
||||
cfg = resolving_view(server_args)
|
||||
# "mooncake_tcp" is mooncake with the TCP transport forced: set MC_FORCE_TCP
|
||||
# so mooncake installs TcpTransport instead of RDMA, rewrite the backend to
|
||||
# mooncake, and skip RDMA HCA selection. Must run before backend-name checks.
|
||||
if server_args.disaggregation_transfer_backend == "mooncake_tcp":
|
||||
if cfg.disaggregation_transfer_backend == "mooncake_tcp":
|
||||
os.environ.setdefault("MC_FORCE_TCP", "1")
|
||||
declare_resolution(
|
||||
server_args,
|
||||
@@ -36,17 +40,17 @@ def handle_pd_disaggregation(server_args: ServerArgs) -> None:
|
||||
"with MC_FORCE_TCP=1 (TCP transport, no RDMA)"
|
||||
)
|
||||
|
||||
if server_args.disaggregation_mode == "prefill" and server_args.dcp_size > 1:
|
||||
if cfg.disaggregation_mode == "prefill" and cfg.dcp_size > 1:
|
||||
logger.warning(
|
||||
"DCP on a PD prefill server is supported when prefill and decode "
|
||||
"use the same DCP layout, but it usually adds communication "
|
||||
"overhead without improving prefill performance."
|
||||
)
|
||||
|
||||
if server_args.disaggregation_mode == "decode" and server_args.dcp_size > 1:
|
||||
if cfg.disaggregation_mode == "decode" and cfg.dcp_size > 1:
|
||||
# Fake transfer moves no KV and is only used for synthetic decode
|
||||
# benchmarks, so it does not need the DCP relayout from Mooncake/NIXL.
|
||||
if server_args.disaggregation_transfer_backend not in (
|
||||
if cfg.disaggregation_transfer_backend not in (
|
||||
"mooncake",
|
||||
"nixl",
|
||||
"fake",
|
||||
@@ -54,36 +58,36 @@ def handle_pd_disaggregation(server_args: ServerArgs) -> None:
|
||||
raise ValueError(
|
||||
"PD decode DCP requires --disaggregation-transfer-backend "
|
||||
"mooncake, nixl, or fake for synthetic benchmarking, got "
|
||||
f"{server_args.disaggregation_transfer_backend!r}."
|
||||
f"{cfg.disaggregation_transfer_backend!r}."
|
||||
)
|
||||
if server_args.disaggregation_decode_enable_radix_cache:
|
||||
if cfg.disaggregation_decode_enable_radix_cache:
|
||||
raise ValueError(
|
||||
"PD decode DCP currently requires chunk cache; "
|
||||
"--disaggregation-decode-enable-radix-cache is not supported."
|
||||
)
|
||||
if server_args.enable_hierarchical_cache:
|
||||
if cfg.enable_hierarchical_cache:
|
||||
raise ValueError(
|
||||
"PD decode DCP currently requires chunk cache; "
|
||||
"--enable-hierarchical-cache is not supported."
|
||||
)
|
||||
|
||||
if server_args.disaggregation_mode == "decode":
|
||||
if server_args.disaggregation_decode_enable_radix_cache:
|
||||
if server_args.enable_hisparse:
|
||||
if cfg.disaggregation_mode == "decode":
|
||||
if cfg.disaggregation_decode_enable_radix_cache:
|
||||
if cfg.enable_hisparse:
|
||||
raise ValueError(
|
||||
"--disaggregation-decode-enable-radix-cache is incompatible "
|
||||
"with --enable-hisparse"
|
||||
)
|
||||
if server_args.disaggregation_transfer_backend == "fake":
|
||||
if cfg.disaggregation_transfer_backend == "fake":
|
||||
raise ValueError(
|
||||
"--disaggregation-decode-enable-radix-cache is incompatible "
|
||||
"with --disaggregation-transfer-backend fake"
|
||||
)
|
||||
if server_args.speculative_algorithm is not None:
|
||||
if cfg.speculative_algorithm is not None:
|
||||
raise ValueError(
|
||||
"--disaggregation-decode-enable-radix-cache is incompatible "
|
||||
"with speculative decoding "
|
||||
f"(--speculative-algorithm {server_args.speculative_algorithm})"
|
||||
f"(--speculative-algorithm {cfg.speculative_algorithm})"
|
||||
)
|
||||
from sglang.srt.arg_groups.overrides import resolved_view
|
||||
|
||||
@@ -110,12 +114,10 @@ def handle_pd_disaggregation(server_args: ServerArgs) -> None:
|
||||
# in-transfer (being-received-from-prefill) requests, on top of the
|
||||
# max_running_requests-derived pool. Large batches get none; small
|
||||
# per-worker batches reserve 2x the batch as cheap overlap headroom.
|
||||
if server_args.disaggregation_decode_extra_slots is None:
|
||||
if cfg.disaggregation_decode_extra_slots is None:
|
||||
extra_slots = 0
|
||||
if server_args.max_running_requests is not None:
|
||||
per_worker = server_args.max_running_requests // max(
|
||||
1, server_args.dp_size
|
||||
)
|
||||
if cfg.max_running_requests is not None:
|
||||
per_worker = cfg.max_running_requests // max(1, cfg.dp_size)
|
||||
if per_worker <= 32:
|
||||
extra_slots = per_worker * 2
|
||||
declare_resolution(
|
||||
@@ -124,23 +126,23 @@ def handle_pd_disaggregation(server_args: ServerArgs) -> None:
|
||||
disaggregation_decode_extra_slots=extra_slots,
|
||||
)
|
||||
|
||||
elif server_args.disaggregation_mode == "prefill":
|
||||
elif cfg.disaggregation_mode == "prefill":
|
||||
assert (
|
||||
server_args.disaggregation_transfer_backend != "fake"
|
||||
cfg.disaggregation_transfer_backend != "fake"
|
||||
), "Prefill server does not support 'fake' as the transfer backend"
|
||||
|
||||
if envs.SGLANG_RUST_SERVER.get():
|
||||
_alias_bootstrap_port_to_api_port(server_args)
|
||||
|
||||
if server_args.disaggregation_mode in ("prefill", "decode"):
|
||||
if cfg.disaggregation_mode in ("prefill", "decode"):
|
||||
if (
|
||||
envs.SGLANG_DISAGG_STAGING_BUFFER.get()
|
||||
and server_args.disaggregation_transfer_backend not in ("mooncake", "nixl")
|
||||
and cfg.disaggregation_transfer_backend not in ("mooncake", "nixl")
|
||||
):
|
||||
raise ValueError(
|
||||
f"SGLANG_DISAGG_STAGING_BUFFER requires "
|
||||
f"disaggregation_transfer_backend='mooncake' or 'nixl', "
|
||||
f"got '{server_args.disaggregation_transfer_backend}'."
|
||||
f"got '{cfg.disaggregation_transfer_backend}'."
|
||||
)
|
||||
|
||||
|
||||
@@ -151,31 +153,32 @@ def _alias_bootstrap_port_to_api_port(server_args: ServerArgs) -> None:
|
||||
field and agrees automatically. Decode is untouched: there the field names
|
||||
the PREFILL side's bootstrap port and must stay as the operator set it.
|
||||
"""
|
||||
cfg = resolving_view(server_args)
|
||||
default_port = next(
|
||||
f.default
|
||||
for f in dataclasses.fields(server_args)
|
||||
if f.name == "disaggregation_bootstrap_port"
|
||||
)
|
||||
if server_args.disaggregation_bootstrap_port not in (
|
||||
if cfg.disaggregation_bootstrap_port not in (
|
||||
default_port,
|
||||
server_args.port,
|
||||
cfg.port,
|
||||
):
|
||||
raise ValueError(
|
||||
"SGLANG_RUST_SERVER serves the PD KV bootstrap registry on the api "
|
||||
"port itself; --disaggregation-bootstrap-port "
|
||||
f"{server_args.disaggregation_bootstrap_port} conflicts with --port "
|
||||
f"{server_args.port}. Drop --disaggregation-bootstrap-port (decode "
|
||||
f"{cfg.disaggregation_bootstrap_port} conflicts with --port "
|
||||
f"{cfg.port}. Drop --disaggregation-bootstrap-port (decode "
|
||||
"nodes and the PD router must then target the prefill api port)."
|
||||
)
|
||||
if server_args.disaggregation_bootstrap_port != server_args.port:
|
||||
if cfg.disaggregation_bootstrap_port != cfg.port:
|
||||
logger.info(
|
||||
"SGLANG_RUST_SERVER: KV bootstrap registry is served on the api "
|
||||
"port; disaggregation_bootstrap_port %d -> %d",
|
||||
server_args.disaggregation_bootstrap_port,
|
||||
server_args.port,
|
||||
cfg.disaggregation_bootstrap_port,
|
||||
cfg.port,
|
||||
)
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_alias_bootstrap_port_to_api_port",
|
||||
disaggregation_bootstrap_port=server_args.port,
|
||||
disaggregation_bootstrap_port=cfg.port,
|
||||
)
|
||||
|
||||
@@ -8,6 +8,7 @@ from typing import TYPE_CHECKING, Optional
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
declare_direct_writes,
|
||||
declare_resolution,
|
||||
resolving_view,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -17,7 +18,8 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _disable_overlap_schedule_for_cpu(server_args: ServerArgs) -> None:
|
||||
if server_args.device != "cpu" or server_args.disable_overlap_schedule:
|
||||
cfg = resolving_view(server_args)
|
||||
if cfg.device != "cpu" or cfg.disable_overlap_schedule:
|
||||
return
|
||||
|
||||
declare_resolution(
|
||||
@@ -71,9 +73,10 @@ def _resolve_speculative_algorithm_alias(
|
||||
|
||||
|
||||
def handle_speculative_decoding(server_args: ServerArgs) -> None:
|
||||
cfg = resolving_view(server_args)
|
||||
if (
|
||||
server_args.speculative_draft_model_path is not None
|
||||
and server_args.speculative_draft_model_revision is None
|
||||
cfg.speculative_draft_model_path is not None
|
||||
and cfg.speculative_draft_model_revision is None
|
||||
):
|
||||
declare_resolution(
|
||||
server_args,
|
||||
@@ -90,11 +93,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:
|
||||
if cfg.speculative_algorithm is not None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"handle_speculative_decoding",
|
||||
speculative_algorithm=server_args.speculative_algorithm.upper(),
|
||||
speculative_algorithm=cfg.speculative_algorithm.upper(),
|
||||
)
|
||||
|
||||
# Removal notice for the retired env var; raw os.getenv on purpose -- the
|
||||
@@ -108,7 +111,7 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None:
|
||||
|
||||
kwargs = {}
|
||||
|
||||
override_config_file = server_args.decrypted_draft_config_file
|
||||
override_config_file = cfg.decrypted_draft_config_file
|
||||
if override_config_file and override_config_file.strip():
|
||||
kwargs["_configuration_file"] = override_config_file.strip()
|
||||
|
||||
@@ -116,17 +119,17 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None:
|
||||
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,
|
||||
cfg.speculative_algorithm,
|
||||
cfg.speculative_draft_model_path,
|
||||
trust_remote_code=cfg.trust_remote_code,
|
||||
kwargs=kwargs,
|
||||
),
|
||||
)
|
||||
|
||||
# Validate --speculative-draft-window-size once, regardless of algorithm.
|
||||
# Consumed by DFLASH (compact draft KV cache) and Llama EAGLE-3 (drafter attention SWA).
|
||||
if server_args.speculative_draft_window_size is not None:
|
||||
window_size = int(server_args.speculative_draft_window_size)
|
||||
if cfg.speculative_draft_window_size is not None:
|
||||
window_size = int(cfg.speculative_draft_window_size)
|
||||
if window_size <= 0:
|
||||
raise ValueError(
|
||||
f"--speculative-draft-window-size must be positive, got {window_size}."
|
||||
@@ -136,19 +139,19 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None:
|
||||
"handle_speculative_decoding",
|
||||
speculative_draft_window_size=window_size,
|
||||
)
|
||||
if server_args.speculative_algorithm not in ("EAGLE3", "DFLASH"):
|
||||
if cfg.speculative_algorithm not in ("EAGLE3", "DFLASH"):
|
||||
logger.warning(
|
||||
"--speculative-draft-window-size has no effect with "
|
||||
"speculative_algorithm=%s (honored by Llama EAGLE-3 and DFLASH only).",
|
||||
server_args.speculative_algorithm,
|
||||
cfg.speculative_algorithm,
|
||||
)
|
||||
|
||||
algo = None
|
||||
if server_args.speculative_algorithm is not None:
|
||||
if cfg.speculative_algorithm is not None:
|
||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||
from sglang.srt.speculative.spec_registry import CustomSpecAlgo
|
||||
|
||||
algo = SpeculativeAlgorithm.from_string(server_args.speculative_algorithm)
|
||||
algo = SpeculativeAlgorithm.from_string(cfg.speculative_algorithm)
|
||||
|
||||
# TODO: move the per-algorithm validation below into spec module hooks.
|
||||
if isinstance(algo, CustomSpecAlgo) and algo.validate_server_args is not None:
|
||||
@@ -158,15 +161,15 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None:
|
||||
algo.validate_server_args,
|
||||
)
|
||||
|
||||
if server_args.speculative_skip_dp_mlp_sync:
|
||||
assert server_args.speculative_algorithm == "EAGLE", (
|
||||
if cfg.speculative_skip_dp_mlp_sync:
|
||||
assert cfg.speculative_algorithm == "EAGLE", (
|
||||
"--speculative-skip-dp-mlp-sync is only supported with "
|
||||
f"speculative_algorithm == EAGLE, got {server_args.speculative_algorithm}."
|
||||
f"speculative_algorithm == EAGLE, got {cfg.speculative_algorithm}."
|
||||
)
|
||||
|
||||
if server_args.speculative_adaptive:
|
||||
if cfg.speculative_adaptive:
|
||||
_maybe_disable_adaptive(server_args)
|
||||
if server_args.speculative_adaptive:
|
||||
if cfg.speculative_adaptive:
|
||||
_init_adaptive_speculative_params(server_args)
|
||||
|
||||
if algo is not None:
|
||||
@@ -180,9 +183,10 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None:
|
||||
|
||||
|
||||
def _handle_dflash(server_args: ServerArgs) -> None:
|
||||
cfg = resolving_view(server_args)
|
||||
from sglang.srt.arg_groups.overrides import resolved_view
|
||||
|
||||
if not (server_args.device.startswith("cuda") or server_args.device == "npu"):
|
||||
if not (cfg.device.startswith("cuda") or cfg.device == "npu"):
|
||||
raise ValueError(
|
||||
"DFLASH speculative decoding only supports CUDA and NPU devices."
|
||||
)
|
||||
@@ -192,12 +196,12 @@ def _handle_dflash(server_args: ServerArgs) -> None:
|
||||
"Currently DFLASH speculative decoding does not support dp attention."
|
||||
)
|
||||
|
||||
if server_args.pp_size != 1:
|
||||
if cfg.pp_size != 1:
|
||||
raise ValueError(
|
||||
"Currently DFLASH speculative decoding only supports pp_size == 1."
|
||||
)
|
||||
|
||||
if server_args.speculative_draft_model_path is None:
|
||||
if cfg.speculative_draft_model_path is None:
|
||||
raise ValueError(
|
||||
"DFLASH speculative decoding requires setting --speculative-draft-model-path."
|
||||
)
|
||||
@@ -207,16 +211,16 @@ def _handle_dflash(server_args: ServerArgs) -> None:
|
||||
# RoPE reservation). Force them to 1 to avoid surprising memory behavior.
|
||||
#
|
||||
# For DFlash, the natural unit is `block_size` (verify window length).
|
||||
if server_args.speculative_num_steps is None:
|
||||
if cfg.speculative_num_steps is None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_dflash",
|
||||
speculative_num_steps=1,
|
||||
)
|
||||
elif int(server_args.speculative_num_steps) != 1:
|
||||
elif int(cfg.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,
|
||||
cfg.speculative_num_steps,
|
||||
)
|
||||
declare_resolution(
|
||||
server_args,
|
||||
@@ -224,16 +228,16 @@ def _handle_dflash(server_args: ServerArgs) -> None:
|
||||
speculative_num_steps=1,
|
||||
)
|
||||
|
||||
if server_args.speculative_eagle_topk is None:
|
||||
if cfg.speculative_eagle_topk is None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_dflash",
|
||||
speculative_eagle_topk=1,
|
||||
)
|
||||
elif int(server_args.speculative_eagle_topk) != 1:
|
||||
elif int(cfg.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,
|
||||
cfg.speculative_eagle_topk,
|
||||
)
|
||||
declare_resolution(
|
||||
server_args,
|
||||
@@ -241,41 +245,41 @@ def _handle_dflash(server_args: ServerArgs) -> None:
|
||||
speculative_eagle_topk=1,
|
||||
)
|
||||
|
||||
if server_args.speculative_dflash_block_size is not None:
|
||||
if int(server_args.speculative_dflash_block_size) <= 0:
|
||||
if cfg.speculative_dflash_block_size is not None:
|
||||
if int(cfg.speculative_dflash_block_size) <= 0:
|
||||
raise ValueError(
|
||||
"DFLASH requires --speculative-dflash-block-size to be positive, "
|
||||
f"got {server_args.speculative_dflash_block_size}."
|
||||
f"got {cfg.speculative_dflash_block_size}."
|
||||
)
|
||||
if server_args.speculative_num_draft_tokens is not None and int(
|
||||
server_args.speculative_num_draft_tokens
|
||||
) != int(server_args.speculative_dflash_block_size):
|
||||
if cfg.speculative_num_draft_tokens is not None and int(
|
||||
cfg.speculative_num_draft_tokens
|
||||
) != int(cfg.speculative_dflash_block_size):
|
||||
raise ValueError(
|
||||
"Both --speculative-num-draft-tokens and --speculative-dflash-block-size are set "
|
||||
"but they differ. For DFLASH they must match. "
|
||||
f"speculative_num_draft_tokens={server_args.speculative_num_draft_tokens}, "
|
||||
f"speculative_dflash_block_size={server_args.speculative_dflash_block_size}."
|
||||
f"speculative_num_draft_tokens={cfg.speculative_num_draft_tokens}, "
|
||||
f"speculative_dflash_block_size={cfg.speculative_dflash_block_size}."
|
||||
)
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_dflash",
|
||||
speculative_num_draft_tokens=int(server_args.speculative_dflash_block_size),
|
||||
speculative_num_draft_tokens=int(cfg.speculative_dflash_block_size),
|
||||
)
|
||||
|
||||
if server_args.speculative_num_draft_tokens is None:
|
||||
if cfg.speculative_num_draft_tokens is None:
|
||||
from sglang.srt.speculative.dflash_utils import (
|
||||
parse_dflash_draft_config,
|
||||
)
|
||||
|
||||
model_override_args = json.loads(server_args.json_model_override_args)
|
||||
model_override_args = json.loads(cfg.json_model_override_args)
|
||||
inferred_block_size = None
|
||||
try:
|
||||
from sglang.srt.utils.hf_transformers_utils import get_config
|
||||
|
||||
draft_hf_config = get_config(
|
||||
server_args.speculative_draft_model_path,
|
||||
trust_remote_code=server_args.trust_remote_code,
|
||||
revision=server_args.speculative_draft_model_revision,
|
||||
cfg.speculative_draft_model_path,
|
||||
trust_remote_code=cfg.trust_remote_code,
|
||||
revision=cfg.speculative_draft_model_revision,
|
||||
model_override_args=model_override_args,
|
||||
)
|
||||
inferred_block_size = parse_dflash_draft_config(
|
||||
@@ -300,18 +304,18 @@ def _handle_dflash(server_args: ServerArgs) -> None:
|
||||
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)
|
||||
if server_args.speculative_draft_window_size < draft_tokens:
|
||||
if cfg.speculative_draft_window_size is not None:
|
||||
draft_tokens = int(cfg.speculative_num_draft_tokens)
|
||||
if cfg.speculative_draft_window_size < draft_tokens:
|
||||
raise ValueError(
|
||||
"--speculative-draft-window-size must be >= "
|
||||
"--speculative-num-draft-tokens (block_size). "
|
||||
f"window_size={server_args.speculative_draft_window_size}, block_size={draft_tokens}."
|
||||
f"window_size={cfg.speculative_draft_window_size}, block_size={draft_tokens}."
|
||||
)
|
||||
|
||||
_resolve_dflash_draft_attention_backend(server_args)
|
||||
|
||||
if server_args.max_running_requests is None:
|
||||
if cfg.max_running_requests is None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_dflash",
|
||||
@@ -321,7 +325,7 @@ def _handle_dflash(server_args: ServerArgs) -> None:
|
||||
"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:
|
||||
if cfg.enable_mixed_chunk:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_dflash",
|
||||
@@ -341,23 +345,24 @@ def _target_checkpoint_bundles_dspark_draft(server_args: ServerArgs) -> bool:
|
||||
|
||||
|
||||
def _handle_dspark(server_args: ServerArgs) -> None:
|
||||
_is_npu = server_args.device.startswith("npu")
|
||||
if not server_args.device.startswith(("cuda", "npu")):
|
||||
cfg = resolving_view(server_args)
|
||||
_is_npu = cfg.device.startswith("npu")
|
||||
if not cfg.device.startswith(("cuda", "npu")):
|
||||
raise ValueError(
|
||||
"DSpark speculative decoding only supports CUDA or NPU device."
|
||||
)
|
||||
|
||||
# dp_size==1 with dp_attention is a degenerate flag under DSV4 CP; skip DP-only checks.
|
||||
if server_args.enable_dp_attention and server_args.dp_size > 1:
|
||||
if not server_args.enable_dp_lm_head:
|
||||
if cfg.enable_dp_attention and cfg.dp_size > 1:
|
||||
if not cfg.enable_dp_lm_head:
|
||||
raise ValueError("DSpark with dp attention requires --enable-dp-lm-head.")
|
||||
if not _is_npu and server_args.moe_a2a_backend not in ("none", "megamoe"):
|
||||
if not _is_npu and cfg.moe_a2a_backend not in ("none", "megamoe"):
|
||||
raise ValueError(
|
||||
"DSpark with dp attention supports moe_a2a_backend 'none' "
|
||||
"(built-in TP MoE) or 'megamoe', got "
|
||||
f"{server_args.moe_a2a_backend!r}."
|
||||
f"{cfg.moe_a2a_backend!r}."
|
||||
)
|
||||
if not _is_npu and server_args.moe_a2a_backend != "none":
|
||||
if not _is_npu and cfg.moe_a2a_backend != "none":
|
||||
from sglang.srt.speculative.ragged_verify import (
|
||||
RaggedVerifyMode,
|
||||
read_ragged_verify_mode,
|
||||
@@ -366,46 +371,46 @@ def _handle_dspark(server_args: ServerArgs) -> None:
|
||||
if read_ragged_verify_mode() is not RaggedVerifyMode.STATIC:
|
||||
raise ValueError(
|
||||
"DSpark with dp attention + "
|
||||
f"moe_a2a_backend={server_args.moe_a2a_backend!r} requires "
|
||||
f"moe_a2a_backend={cfg.moe_a2a_backend!r} requires "
|
||||
"SGLANG_RAGGED_VERIFY_MODE=static."
|
||||
)
|
||||
if server_args.attn_cp_size > 1:
|
||||
if cfg.attn_cp_size > 1:
|
||||
raise ValueError(
|
||||
"DSpark with dp attention does not support context parallel "
|
||||
f"(attn_cp_size={server_args.attn_cp_size})."
|
||||
f"(attn_cp_size={cfg.attn_cp_size})."
|
||||
)
|
||||
if (
|
||||
not _is_npu
|
||||
and server_args.speculative_moe_a2a_backend is not None
|
||||
and server_args.speculative_moe_a2a_backend != server_args.moe_a2a_backend
|
||||
and cfg.speculative_moe_a2a_backend is not None
|
||||
and cfg.speculative_moe_a2a_backend != cfg.moe_a2a_backend
|
||||
):
|
||||
raise ValueError(
|
||||
"DSpark ignores --speculative-moe-a2a-backend; with dp attention it "
|
||||
f"must match the target moe_a2a_backend={server_args.moe_a2a_backend!r} "
|
||||
f"(got {server_args.speculative_moe_a2a_backend!r})."
|
||||
f"must match the target moe_a2a_backend={cfg.moe_a2a_backend!r} "
|
||||
f"(got {cfg.speculative_moe_a2a_backend!r})."
|
||||
)
|
||||
|
||||
if server_args.pp_size != 1:
|
||||
if cfg.pp_size != 1:
|
||||
raise ValueError(
|
||||
"Currently DSpark speculative decoding only supports pp_size == 1."
|
||||
)
|
||||
|
||||
if server_args.speculative_draft_model_path is None:
|
||||
if cfg.speculative_draft_model_path is None:
|
||||
if _target_checkpoint_bundles_dspark_draft(server_args):
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_dspark",
|
||||
speculative_draft_model_path=server_args.model_path,
|
||||
speculative_draft_model_path=cfg.model_path,
|
||||
)
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_dspark",
|
||||
speculative_draft_model_revision=server_args.revision,
|
||||
speculative_draft_model_revision=cfg.revision,
|
||||
)
|
||||
logger.info(
|
||||
"DSpark draft weights are bundled in the target checkpoint; "
|
||||
"defaulting --speculative-draft-model-path to --model-path (%s).",
|
||||
server_args.model_path,
|
||||
cfg.model_path,
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
@@ -413,16 +418,16 @@ def _handle_dspark(server_args: ServerArgs) -> None:
|
||||
"--speculative-draft-model-path."
|
||||
)
|
||||
|
||||
if server_args.speculative_num_steps is None:
|
||||
if cfg.speculative_num_steps is None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_dspark",
|
||||
speculative_num_steps=1,
|
||||
)
|
||||
elif int(server_args.speculative_num_steps) != 1:
|
||||
elif int(cfg.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,
|
||||
cfg.speculative_num_steps,
|
||||
)
|
||||
declare_resolution(
|
||||
server_args,
|
||||
@@ -430,16 +435,16 @@ def _handle_dspark(server_args: ServerArgs) -> None:
|
||||
speculative_num_steps=1,
|
||||
)
|
||||
|
||||
if server_args.speculative_eagle_topk is None:
|
||||
if cfg.speculative_eagle_topk is None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_dspark",
|
||||
speculative_eagle_topk=1,
|
||||
)
|
||||
elif int(server_args.speculative_eagle_topk) != 1:
|
||||
elif int(cfg.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,
|
||||
cfg.speculative_eagle_topk,
|
||||
)
|
||||
declare_resolution(
|
||||
server_args,
|
||||
@@ -463,17 +468,17 @@ def _handle_dspark(server_args: ServerArgs) -> None:
|
||||
)
|
||||
|
||||
gamma: Optional[int] = None
|
||||
if server_args.speculative_dspark_block_size is not None:
|
||||
if int(server_args.speculative_dspark_block_size) <= 0:
|
||||
if cfg.speculative_dspark_block_size is not None:
|
||||
if int(cfg.speculative_dspark_block_size) <= 0:
|
||||
raise ValueError(
|
||||
"DSpark requires --speculative-dspark-block-size to be positive, "
|
||||
f"got {server_args.speculative_dspark_block_size}."
|
||||
f"got {cfg.speculative_dspark_block_size}."
|
||||
)
|
||||
gamma = int(server_args.speculative_dspark_block_size)
|
||||
gamma = int(cfg.speculative_dspark_block_size)
|
||||
else:
|
||||
if draft_config is not None:
|
||||
gamma = draft_config.resolve_gamma(default=None)
|
||||
if gamma is None and server_args.speculative_num_draft_tokens is None:
|
||||
if gamma is None and cfg.speculative_num_draft_tokens is None:
|
||||
gamma = DEFAULT_DSPARK_GAMMA
|
||||
logger.warning(
|
||||
"DSpark gamma is not set; defaulting to %d.",
|
||||
@@ -483,13 +488,13 @@ def _handle_dspark(server_args: ServerArgs) -> None:
|
||||
if gamma is not None:
|
||||
verify_window = int(gamma) + 1
|
||||
if (
|
||||
server_args.speculative_num_draft_tokens is not None
|
||||
and int(server_args.speculative_num_draft_tokens) != verify_window
|
||||
cfg.speculative_num_draft_tokens is not None
|
||||
and int(cfg.speculative_num_draft_tokens) != verify_window
|
||||
):
|
||||
raise ValueError(
|
||||
"DSpark speculative_num_draft_tokens must equal gamma + 1 "
|
||||
f"(= {verify_window} for gamma={gamma}), but got "
|
||||
f"speculative_num_draft_tokens={server_args.speculative_num_draft_tokens}."
|
||||
f"speculative_num_draft_tokens={cfg.speculative_num_draft_tokens}."
|
||||
)
|
||||
declare_resolution(
|
||||
server_args,
|
||||
@@ -497,18 +502,18 @@ def _handle_dspark(server_args: ServerArgs) -> None:
|
||||
speculative_num_draft_tokens=verify_window,
|
||||
)
|
||||
|
||||
if server_args.speculative_num_draft_tokens is None:
|
||||
if cfg.speculative_num_draft_tokens is None:
|
||||
raise ValueError(
|
||||
"DSpark could not resolve speculative_num_draft_tokens; set "
|
||||
"--speculative-dspark-block-size (= gamma)."
|
||||
)
|
||||
if int(server_args.speculative_num_draft_tokens) < 2:
|
||||
if int(cfg.speculative_num_draft_tokens) < 2:
|
||||
raise ValueError(
|
||||
"DSpark speculative_num_draft_tokens must be >= 2 (= gamma + 1), "
|
||||
f"got {server_args.speculative_num_draft_tokens}."
|
||||
f"got {cfg.speculative_num_draft_tokens}."
|
||||
)
|
||||
|
||||
if server_args.max_running_requests is None:
|
||||
if cfg.max_running_requests is None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_dspark",
|
||||
@@ -518,7 +523,7 @@ def _handle_dspark(server_args: ServerArgs) -> None:
|
||||
"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:
|
||||
if cfg.enable_mixed_chunk:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_dspark",
|
||||
@@ -535,7 +540,7 @@ def _handle_dspark(server_args: ServerArgs) -> None:
|
||||
|
||||
ragged_mode = read_ragged_verify_mode()
|
||||
if (
|
||||
server_args.speculative_dspark_align_verify_tokens_to_graph_tier
|
||||
cfg.speculative_dspark_align_verify_tokens_to_graph_tier
|
||||
and ragged_mode is not RaggedVerifyMode.COMPACT
|
||||
):
|
||||
logger.warning(
|
||||
@@ -544,10 +549,7 @@ def _handle_dspark(server_args: ServerArgs) -> None:
|
||||
"a no-op.",
|
||||
ragged_mode.value,
|
||||
)
|
||||
if (
|
||||
server_args.speculative_dspark_sps_table_path
|
||||
and ragged_mode is RaggedVerifyMode.STATIC
|
||||
):
|
||||
if cfg.speculative_dspark_sps_table_path and ragged_mode is RaggedVerifyMode.STATIC:
|
||||
logger.warning(
|
||||
"--speculative-dspark-sps-table-path feeds the ragged-verify budget "
|
||||
"scheduler, which is off under SGLANG_RAGGED_VERIFY_MODE=static; it "
|
||||
@@ -561,6 +563,7 @@ def _resolve_dflash_draft_attention_backend(server_args: ServerArgs) -> None:
|
||||
Consumed by ModelRunner's `is_draft_worker` override (one backend for all
|
||||
draft modes).
|
||||
"""
|
||||
cfg = resolving_view(server_args)
|
||||
from sglang.srt.utils import is_hip
|
||||
|
||||
supported_draft_backends = (
|
||||
@@ -574,7 +577,7 @@ def _resolve_dflash_draft_attention_backend(server_args: ServerArgs) -> None:
|
||||
# Use triton on ROCm (no FlashInfer), flashinfer on CUDA.
|
||||
fallback_backend = "triton" if is_hip() else "flashinfer"
|
||||
|
||||
draft_backend = server_args.speculative_draft_attention_backend
|
||||
draft_backend = cfg.speculative_draft_attention_backend
|
||||
if draft_backend is None:
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
attention_backends_of,
|
||||
@@ -589,10 +592,10 @@ def _resolve_dflash_draft_attention_backend(server_args: ServerArgs) -> None:
|
||||
from sglang.srt.utils.hf_transformers_utils import get_config
|
||||
|
||||
draft_hf_config = get_config(
|
||||
server_args.speculative_draft_model_path,
|
||||
trust_remote_code=server_args.trust_remote_code,
|
||||
revision=server_args.speculative_draft_model_revision,
|
||||
model_override_args=json.loads(server_args.json_model_override_args),
|
||||
cfg.speculative_draft_model_path,
|
||||
trust_remote_code=cfg.trust_remote_code,
|
||||
revision=cfg.speculative_draft_model_revision,
|
||||
model_override_args=json.loads(cfg.json_model_override_args),
|
||||
)
|
||||
draft_text_config = (
|
||||
getattr(draft_hf_config, "text_config", None) or draft_hf_config
|
||||
@@ -635,7 +638,8 @@ def _resolve_dflash_draft_attention_backend(server_args: ServerArgs) -> None:
|
||||
|
||||
|
||||
def _handle_frozen_kv_mtp(server_args: ServerArgs) -> None:
|
||||
if server_args.max_running_requests is None:
|
||||
cfg = resolving_view(server_args)
|
||||
if cfg.max_running_requests is None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_frozen_kv_mtp",
|
||||
@@ -645,7 +649,7 @@ def _handle_frozen_kv_mtp(server_args: ServerArgs) -> None:
|
||||
"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:
|
||||
if cfg.enable_mixed_chunk:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_frozen_kv_mtp",
|
||||
@@ -658,13 +662,14 @@ def _handle_frozen_kv_mtp(server_args: ServerArgs) -> None:
|
||||
|
||||
|
||||
def _handle_eagle_family(server_args: ServerArgs) -> None:
|
||||
cfg = resolving_view(server_args)
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
attention_backends_of,
|
||||
resolved_view,
|
||||
)
|
||||
|
||||
if (
|
||||
server_args.speculative_algorithm == "STANDALONE"
|
||||
cfg.speculative_algorithm == "STANDALONE"
|
||||
and resolved_view(server_args).enable_dp_attention
|
||||
):
|
||||
# TODO: support dp attention for standalone speculative decoding
|
||||
@@ -672,7 +677,7 @@ def _handle_eagle_family(server_args: ServerArgs) -> None:
|
||||
"Currently standalone speculative decoding does not support dp attention."
|
||||
)
|
||||
|
||||
if server_args.max_running_requests is None:
|
||||
if cfg.max_running_requests is None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_eagle_family",
|
||||
@@ -690,7 +695,7 @@ def _handle_eagle_family(server_args: ServerArgs) -> None:
|
||||
"speculative decoding."
|
||||
)
|
||||
|
||||
if server_args.enable_mixed_chunk:
|
||||
if cfg.enable_mixed_chunk:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_eagle_family",
|
||||
@@ -716,16 +721,16 @@ def _handle_eagle_family(server_args: ServerArgs) -> None:
|
||||
"PixtralForConditionalGeneration",
|
||||
"HYV3ForCausalLM",
|
||||
]:
|
||||
if server_args.speculative_draft_model_path is None:
|
||||
if cfg.speculative_draft_model_path is None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_eagle_family",
|
||||
speculative_draft_model_path=server_args.model_path,
|
||||
speculative_draft_model_path=cfg.model_path,
|
||||
)
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_eagle_family",
|
||||
speculative_draft_model_revision=server_args.revision,
|
||||
speculative_draft_model_revision=cfg.revision,
|
||||
)
|
||||
else:
|
||||
if model_arch not in [
|
||||
@@ -736,13 +741,10 @@ def _handle_eagle_family(server_args: ServerArgs) -> None:
|
||||
"DeepSeek MTP does not require setting speculative_draft_model_path."
|
||||
)
|
||||
|
||||
if (
|
||||
not server_args.speculative_adaptive
|
||||
and server_args.speculative_num_steps is None
|
||||
):
|
||||
if not cfg.speculative_adaptive and cfg.speculative_num_steps is None:
|
||||
assert (
|
||||
server_args.speculative_eagle_topk is None
|
||||
and server_args.speculative_num_draft_tokens is None
|
||||
cfg.speculative_eagle_topk is None
|
||||
and cfg.speculative_num_draft_tokens is None
|
||||
)
|
||||
|
||||
steps, topk, draft_tokens = _auto_choose_speculative_params(
|
||||
@@ -757,29 +759,29 @@ def _handle_eagle_family(server_args: ServerArgs) -> None:
|
||||
)
|
||||
|
||||
if "trtllm_mha" in attention_backends_of(resolved_view(server_args)):
|
||||
if server_args.speculative_eagle_topk > 1:
|
||||
if cfg.speculative_eagle_topk > 1:
|
||||
raise ValueError(
|
||||
"trtllm_mha backend only supports topk = 1 for speculative decoding."
|
||||
)
|
||||
|
||||
if server_args.speculative_use_rejection_sampling:
|
||||
if cfg.speculative_use_rejection_sampling:
|
||||
# Resolved alias by now: NEXTN -> EAGLE, Gemma4 draft -> FROZEN_KV_MTP.
|
||||
# Only the EAGLE/EAGLE3 draft workers emit a target-vocab proposal that
|
||||
# the rejection-sampling kernel consumes; everything else (STANDALONE,
|
||||
# FROZEN_KV_MTP, NGRAM, DFLASH) is unsupported.
|
||||
if server_args.speculative_algorithm not in ("EAGLE", "EAGLE3"):
|
||||
if cfg.speculative_algorithm not in ("EAGLE", "EAGLE3"):
|
||||
raise NotImplementedError(
|
||||
"--speculative-use-rejection-sampling is only supported for "
|
||||
"EAGLE / EAGLE3 / NEXTN, not "
|
||||
f"speculative_algorithm={server_args.speculative_algorithm}."
|
||||
f"speculative_algorithm={cfg.speculative_algorithm}."
|
||||
)
|
||||
if server_args.speculative_eagle_topk != 1:
|
||||
if cfg.speculative_eagle_topk != 1:
|
||||
raise ValueError(
|
||||
"--speculative-use-rejection-sampling requires --speculative-eagle-topk=1."
|
||||
)
|
||||
if (
|
||||
server_args.speculative_accept_threshold_single != 1.0
|
||||
or server_args.speculative_accept_threshold_acc != 1.0
|
||||
cfg.speculative_accept_threshold_single != 1.0
|
||||
or cfg.speculative_accept_threshold_acc != 1.0
|
||||
):
|
||||
raise ValueError(
|
||||
"--speculative-use-rejection-sampling is incompatible with "
|
||||
@@ -787,7 +789,7 @@ def _handle_eagle_family(server_args: ServerArgs) -> None:
|
||||
"--speculative-accept-threshold-acc; rejection sampling ignores "
|
||||
"the accept thresholds."
|
||||
)
|
||||
if server_args.enable_deterministic_inference:
|
||||
if cfg.enable_deterministic_inference:
|
||||
raise ValueError(
|
||||
"--speculative-use-rejection-sampling is incompatible with "
|
||||
"--enable-deterministic-inference; the sampling kernel draws "
|
||||
@@ -798,7 +800,7 @@ def _handle_eagle_family(server_args: ServerArgs) -> None:
|
||||
|
||||
if (
|
||||
resolved_view(server_args).enable_multi_layer_eagle
|
||||
and server_args.speculative_eagle_topk != 1
|
||||
and cfg.speculative_eagle_topk != 1
|
||||
):
|
||||
raise ValueError(
|
||||
"--speculative-use-rejection-sampling with multi-layer EAGLE "
|
||||
@@ -811,9 +813,8 @@ def _handle_eagle_family(server_args: ServerArgs) -> None:
|
||||
)
|
||||
|
||||
if (
|
||||
server_args.speculative_eagle_topk == 1
|
||||
and server_args.speculative_num_draft_tokens
|
||||
!= server_args.speculative_num_steps + 1
|
||||
cfg.speculative_eagle_topk == 1
|
||||
and cfg.speculative_num_draft_tokens != cfg.speculative_num_steps + 1
|
||||
):
|
||||
logger.warning(
|
||||
"speculative_num_draft_tokens is adjusted to speculative_num_steps + 1 when speculative_eagle_topk == 1"
|
||||
@@ -821,7 +822,7 @@ def _handle_eagle_family(server_args: ServerArgs) -> None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_eagle_family",
|
||||
speculative_num_draft_tokens=server_args.speculative_num_steps + 1,
|
||||
speculative_num_draft_tokens=cfg.speculative_num_steps + 1,
|
||||
)
|
||||
|
||||
# topk > 1 + page_size > 1 needs the two-pass cascade draft-decode (shared prefix
|
||||
@@ -830,7 +831,7 @@ def _handle_eagle_family(server_args: ServerArgs) -> None:
|
||||
_PAGE_TREE_SPEC_BACKENDS = ("flashinfer", "fa3", "triton")
|
||||
view = resolved_view(server_args)
|
||||
if (
|
||||
server_args.speculative_eagle_topk > 1
|
||||
cfg.speculative_eagle_topk > 1
|
||||
and view.page_size > 1
|
||||
and view.attention_backend not in _PAGE_TREE_SPEC_BACKENDS
|
||||
):
|
||||
@@ -842,14 +843,15 @@ def _handle_eagle_family(server_args: ServerArgs) -> None:
|
||||
|
||||
|
||||
def _handle_ngram(server_args: ServerArgs) -> None:
|
||||
if server_args.device not in ("cuda", "cpu"):
|
||||
cfg = resolving_view(server_args)
|
||||
if cfg.device not in ("cuda", "cpu"):
|
||||
raise ValueError(
|
||||
"Ngram speculative decoding only supports CUDA or CPU devices."
|
||||
)
|
||||
|
||||
_disable_overlap_schedule_for_cpu(server_args)
|
||||
|
||||
if server_args.max_running_requests is None:
|
||||
if cfg.max_running_requests is None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_ngram",
|
||||
@@ -867,9 +869,9 @@ def _handle_ngram(server_args: ServerArgs) -> None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_ngram",
|
||||
speculative_eagle_topk=server_args.speculative_ngram_max_bfs_breadth,
|
||||
speculative_eagle_topk=cfg.speculative_ngram_max_bfs_breadth,
|
||||
)
|
||||
if server_args.speculative_num_draft_tokens is None:
|
||||
if cfg.speculative_num_draft_tokens is None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_ngram",
|
||||
@@ -879,31 +881,31 @@ def _handle_ngram(server_args: ServerArgs) -> None:
|
||||
"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:
|
||||
if cfg.speculative_num_steps is None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_ngram",
|
||||
speculative_num_steps=server_args.speculative_num_draft_tokens
|
||||
// server_args.speculative_eagle_topk,
|
||||
speculative_num_steps=cfg.speculative_num_draft_tokens
|
||||
// cfg.speculative_eagle_topk,
|
||||
)
|
||||
if server_args.speculative_ngram_external_corpus_path is not None:
|
||||
if server_args.speculative_ngram_external_sam_budget <= 0:
|
||||
if cfg.speculative_ngram_external_corpus_path is not None:
|
||||
if cfg.speculative_ngram_external_sam_budget <= 0:
|
||||
raise ValueError(
|
||||
"--speculative-ngram-external-sam-budget must be positive when "
|
||||
"--speculative-ngram-external-corpus-path is set."
|
||||
)
|
||||
if server_args.speculative_ngram_external_corpus_max_tokens <= 0:
|
||||
if cfg.speculative_ngram_external_corpus_max_tokens <= 0:
|
||||
raise ValueError(
|
||||
"--speculative-ngram-external-corpus-max-tokens must be positive when "
|
||||
"--speculative-ngram-external-corpus-path is set."
|
||||
)
|
||||
if (
|
||||
server_args.speculative_ngram_external_sam_budget
|
||||
> server_args.speculative_num_draft_tokens - 1
|
||||
cfg.speculative_ngram_external_sam_budget
|
||||
> cfg.speculative_num_draft_tokens - 1
|
||||
):
|
||||
raise ValueError(
|
||||
"speculative_ngram_external_sam_budget must be less than or equal to "
|
||||
f"speculative_num_draft_tokens - 1 ({server_args.speculative_num_draft_tokens - 1})."
|
||||
f"speculative_num_draft_tokens - 1 ({cfg.speculative_num_draft_tokens - 1})."
|
||||
)
|
||||
logger.warning(
|
||||
"The mixed chunked prefill are disabled because of "
|
||||
@@ -914,12 +916,12 @@ def _handle_ngram(server_args: ServerArgs) -> None:
|
||||
|
||||
view = resolved_view(server_args)
|
||||
if (
|
||||
server_args.speculative_eagle_topk > 1
|
||||
cfg.speculative_eagle_topk > 1
|
||||
and view.page_size > 1
|
||||
and view.attention_backend != "flashinfer"
|
||||
):
|
||||
raise ValueError(
|
||||
f"speculative_eagle_topk({server_args.speculative_eagle_topk}) > 1 "
|
||||
f"speculative_eagle_topk({cfg.speculative_eagle_topk}) > 1 "
|
||||
f"with page_size({view.page_size}) > 1 is unstable "
|
||||
"and produces incorrect results for paged attention backends. "
|
||||
"This combination is only supported for the 'flashinfer' backend."
|
||||
@@ -950,31 +952,32 @@ def _maybe_disable_adaptive(server_args: ServerArgs) -> None:
|
||||
|
||||
|
||||
def _init_adaptive_speculative_params(server_args: ServerArgs) -> None:
|
||||
cfg = resolving_view(server_args)
|
||||
from sglang.srt.speculative.adaptive_spec_params import (
|
||||
resolve_candidate_steps_from_config,
|
||||
)
|
||||
|
||||
candidate_steps = resolve_candidate_steps_from_config(
|
||||
cfg_path=server_args.speculative_adaptive_config,
|
||||
cfg_path=cfg.speculative_adaptive_config,
|
||||
)
|
||||
|
||||
if server_args.speculative_eagle_topk is None:
|
||||
if cfg.speculative_eagle_topk is None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_init_adaptive_speculative_params",
|
||||
speculative_eagle_topk=1,
|
||||
)
|
||||
|
||||
if server_args.speculative_num_steps is None:
|
||||
if cfg.speculative_num_steps is None:
|
||||
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:
|
||||
if cfg.speculative_num_steps not in candidate_steps:
|
||||
raise ValueError(
|
||||
f"--speculative-num-steps={server_args.speculative_num_steps} "
|
||||
f"--speculative-num-steps={cfg.speculative_num_steps} "
|
||||
f"is not in the adaptive config candidate_steps {candidate_steps}. "
|
||||
"Pass one of those values."
|
||||
)
|
||||
@@ -982,7 +985,7 @@ def _init_adaptive_speculative_params(server_args: ServerArgs) -> None:
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_init_adaptive_speculative_params",
|
||||
speculative_num_draft_tokens=server_args.speculative_num_steps + 1,
|
||||
speculative_num_draft_tokens=cfg.speculative_num_steps + 1,
|
||||
)
|
||||
|
||||
|
||||
@@ -992,7 +995,8 @@ def _auto_choose_speculative_params(server_args: ServerArgs, model_arch: str) ->
|
||||
|
||||
You can tune them on your own models and prompts with scripts/playground/bench_speculative.py
|
||||
"""
|
||||
if server_args.speculative_algorithm == "STANDALONE":
|
||||
cfg = resolving_view(server_args)
|
||||
if cfg.speculative_algorithm == "STANDALONE":
|
||||
return (3, 1, 4)
|
||||
if model_arch in ["LlamaForCausalLM"]:
|
||||
return (5, 4, 8)
|
||||
|
||||
@@ -589,46 +589,46 @@ class ModelConfig:
|
||||
context_length: Optional[int] = None,
|
||||
**kwargs,
|
||||
):
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
|
||||
cfg = resolving_view(server_args)
|
||||
quantization = (
|
||||
server_args.speculative_draft_model_quantization
|
||||
cfg.speculative_draft_model_quantization
|
||||
if is_draft_model
|
||||
else server_args.quantization
|
||||
else cfg.quantization
|
||||
)
|
||||
override_config_file = (
|
||||
server_args.decrypted_draft_config_file
|
||||
cfg.decrypted_draft_config_file
|
||||
if is_draft_model
|
||||
else server_args.decrypted_config_file
|
||||
else cfg.decrypted_config_file
|
||||
)
|
||||
return ModelConfig(
|
||||
model_path=model_path or server_args.model_path,
|
||||
trust_remote_code=server_args.trust_remote_code,
|
||||
revision=model_revision or server_args.revision,
|
||||
model_path=model_path or cfg.model_path,
|
||||
trust_remote_code=cfg.trust_remote_code,
|
||||
revision=model_revision or cfg.revision,
|
||||
context_length=(
|
||||
context_length
|
||||
if context_length is not None
|
||||
else server_args.context_length
|
||||
context_length if context_length is not None else cfg.context_length
|
||||
),
|
||||
model_override_args=server_args.json_model_override_args,
|
||||
is_embedding=server_args.is_embedding,
|
||||
enable_multimodal=server_args.enable_multimodal,
|
||||
dtype=server_args.dtype,
|
||||
model_override_args=cfg.json_model_override_args,
|
||||
is_embedding=cfg.is_embedding,
|
||||
enable_multimodal=cfg.enable_multimodal,
|
||||
dtype=cfg.dtype,
|
||||
quantization=quantization,
|
||||
model_impl=server_args.model_impl,
|
||||
sampling_defaults=server_args.sampling_defaults,
|
||||
quantize_and_serve=server_args.quantize_and_serve,
|
||||
model_impl=cfg.model_impl,
|
||||
sampling_defaults=cfg.sampling_defaults,
|
||||
quantize_and_serve=cfg.quantize_and_serve,
|
||||
override_config_file=override_config_file,
|
||||
is_multi_layer_eagle=server_args.enable_multi_layer_eagle,
|
||||
language_only=server_args.language_only,
|
||||
language_model_only=server_args.language_model_only,
|
||||
encoder_only=server_args.encoder_only,
|
||||
is_multi_layer_eagle=cfg.enable_multi_layer_eagle,
|
||||
language_only=cfg.language_only,
|
||||
language_model_only=cfg.language_model_only,
|
||||
encoder_only=cfg.encoder_only,
|
||||
is_draft_model=is_draft_model,
|
||||
is_draft_quantization_explicit=(
|
||||
is_draft_model
|
||||
and server_args._speculative_draft_quantization_explicitly_set
|
||||
is_draft_model and cfg._speculative_draft_quantization_explicitly_set
|
||||
),
|
||||
disable_hybrid_swa_memory=server_args.disable_hybrid_swa_memory,
|
||||
model_config_parser=server_args.model_config_parser,
|
||||
speculative_algorithm=server_args.speculative_algorithm,
|
||||
disable_hybrid_swa_memory=cfg.disable_hybrid_swa_memory,
|
||||
model_config_parser=cfg.model_config_parser,
|
||||
speculative_algorithm=cfg.speculative_algorithm,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
@@ -25,13 +25,16 @@ class DllmConfig:
|
||||
def from_server_args(
|
||||
server_args: ServerArgs,
|
||||
):
|
||||
if server_args.dllm_algorithm is None:
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
|
||||
cfg = resolving_view(server_args)
|
||||
if cfg.dllm_algorithm is None:
|
||||
return None
|
||||
|
||||
model_config = ModelConfig.from_server_args(
|
||||
server_args,
|
||||
model_path=server_args.model_path,
|
||||
model_revision=server_args.revision,
|
||||
model_path=cfg.model_path,
|
||||
model_revision=cfg.revision,
|
||||
)
|
||||
DLLM_PARAMS = {
|
||||
"LLaDA2MoeModelLM": {"block_size": 32, "mask_id": 156895},
|
||||
@@ -48,13 +51,11 @@ class DllmConfig:
|
||||
raise RuntimeError(f"Unknown diffusion LLM: {arch}")
|
||||
|
||||
max_running_requests = (
|
||||
1
|
||||
if server_args.max_running_requests is None
|
||||
else server_args.max_running_requests
|
||||
1 if cfg.max_running_requests is None else cfg.max_running_requests
|
||||
)
|
||||
|
||||
algorithm_config = {}
|
||||
if server_args.dllm_algorithm_config is not None:
|
||||
if cfg.dllm_algorithm_config is not None:
|
||||
try:
|
||||
import yaml
|
||||
except ImportError:
|
||||
@@ -62,17 +63,17 @@ class DllmConfig:
|
||||
"Please install PyYAML to use YAML config files. "
|
||||
"`pip install pyyaml`"
|
||||
)
|
||||
with open(server_args.dllm_algorithm_config, "r") as f:
|
||||
with open(cfg.dllm_algorithm_config, "r") as f:
|
||||
algorithm_config = yaml.safe_load(f)
|
||||
|
||||
# Parse common algorithm configurations
|
||||
block_size = algorithm_config.get("block_size", block_size)
|
||||
|
||||
return DllmConfig(
|
||||
algorithm=server_args.dllm_algorithm,
|
||||
algorithm=cfg.dllm_algorithm,
|
||||
algorithm_config=algorithm_config,
|
||||
block_size=block_size,
|
||||
mask_id=mask_id,
|
||||
max_running_requests=max_running_requests,
|
||||
first_done_first_out_mode=server_args.dllm_fdfo,
|
||||
first_done_first_out_mode=cfg.dllm_fdfo,
|
||||
)
|
||||
|
||||
@@ -43,6 +43,9 @@ def set_default_server_args(args: "ServerArgs"):
|
||||
"""
|
||||
Set default server arguments for NPU backend.
|
||||
"""
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
|
||||
cfg = resolving_view(args)
|
||||
|
||||
# NPU only works with "ascend" attention backend for now
|
||||
declare_resolution(
|
||||
@@ -60,7 +63,7 @@ def set_default_server_args(args: "ServerArgs"):
|
||||
"set_default_server_args",
|
||||
decode_attention_backend="ascend",
|
||||
)
|
||||
if args.page_size is None:
|
||||
if cfg.page_size is None:
|
||||
declare_resolution(
|
||||
args,
|
||||
"set_default_server_args",
|
||||
@@ -68,33 +71,33 @@ def set_default_server_args(args: "ServerArgs"):
|
||||
)
|
||||
|
||||
# NPU memory settings
|
||||
decode = args.cuda_graph_config.decode
|
||||
decode = cfg.cuda_graph_config.decode
|
||||
npu_mem = get_npu_memory_capacity()
|
||||
if npu_mem <= 32 * 1024:
|
||||
# Ascend 910B4,910B4_1
|
||||
# (chunked_prefill_size 4k, max_bs 16 if tp < 4 else 64)
|
||||
if args.chunked_prefill_size is None:
|
||||
if cfg.chunked_prefill_size is None:
|
||||
declare_resolution(
|
||||
args,
|
||||
"set_default_server_args",
|
||||
chunked_prefill_size=4 * 1024,
|
||||
)
|
||||
if decode.max_bs is None:
|
||||
if args.tp_size < 4:
|
||||
if cfg.tp_size < 4:
|
||||
decode.max_bs = 16
|
||||
else:
|
||||
decode.max_bs = 64
|
||||
elif npu_mem <= 64 * 1024:
|
||||
# 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:
|
||||
if cfg.chunked_prefill_size is None:
|
||||
declare_resolution(
|
||||
args,
|
||||
"set_default_server_args",
|
||||
chunked_prefill_size=8 * 1024,
|
||||
)
|
||||
if decode.max_bs is None:
|
||||
if args.tp_size < 4:
|
||||
if cfg.tp_size < 4:
|
||||
decode.max_bs = 64
|
||||
else:
|
||||
decode.max_bs = 256
|
||||
@@ -107,7 +110,7 @@ def set_default_server_args(args: "ServerArgs"):
|
||||
)
|
||||
|
||||
# handles hierarchical cache configs
|
||||
if args.enable_hierarchical_cache:
|
||||
if cfg.enable_hierarchical_cache:
|
||||
declare_resolution(
|
||||
args,
|
||||
"set_default_server_args",
|
||||
|
||||
@@ -239,18 +239,21 @@ _STRATEGY: Optional[ContextParallelStrategy] = None
|
||||
|
||||
def init_cp_strategy(server_args: ServerArgs) -> None:
|
||||
"""Bind the configured CP strategy for this process."""
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
|
||||
cfg = resolving_view(server_args)
|
||||
global _STRATEGY
|
||||
|
||||
if not getattr(server_args, "enable_prefill_cp", False):
|
||||
if not cfg.enable_prefill_cp:
|
||||
_STRATEGY = None
|
||||
return
|
||||
|
||||
cp_size = getattr(server_args, "attn_cp_size", 1)
|
||||
cp_size = cfg.attn_cp_size
|
||||
if cp_size <= 1:
|
||||
_STRATEGY = None
|
||||
return
|
||||
|
||||
kind = ContextParallelStrategyKind.from_string(server_args.cp_strategy)
|
||||
kind = ContextParallelStrategyKind.from_string(cfg.cp_strategy)
|
||||
if kind == ContextParallelStrategyKind.ZIGZAG:
|
||||
from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy
|
||||
|
||||
@@ -262,7 +265,7 @@ def init_cp_strategy(server_args: ServerArgs) -> None:
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported cp_strategy kind {kind} for "
|
||||
f"cp_strategy={server_args.cp_strategy!r}"
|
||||
f"cp_strategy={cfg.cp_strategy!r}"
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -42,12 +42,15 @@ if TYPE_CHECKING:
|
||||
|
||||
def supports_prefill_cp_bcg(server_args: ServerArgs) -> bool:
|
||||
"""Return whether the selected prefill-CP configuration supports BCG."""
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
|
||||
cfg = resolving_view(server_args)
|
||||
resolved = server_args._resolved()
|
||||
prefill_attention_backend, _ = server_args._resolved_attention_backends()
|
||||
return (
|
||||
server_args.enable_prefill_cp
|
||||
and resolved.attn_cp_size == server_args.tp_size
|
||||
and server_args.cp_strategy == "zigzag"
|
||||
cfg.enable_prefill_cp
|
||||
and resolved.attn_cp_size == cfg.tp_size
|
||||
and cfg.cp_strategy == "zigzag"
|
||||
and prefill_attention_backend == "trtllm_mha"
|
||||
)
|
||||
|
||||
|
||||
@@ -14,6 +14,9 @@ def validate_experimental_sgl_marlin_server_args(
|
||||
server_args: Any, resolved_args: Any
|
||||
) -> None:
|
||||
"""Validate startup options before the experimental runner is constructed."""
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
|
||||
cfg = resolving_view(server_args)
|
||||
|
||||
if resolved_args.ep_size > 1 and resolved_args.moe_a2a_backend != "none":
|
||||
raise ValueError("experimental_sgl_marlin EP requires --moe-a2a-backend none")
|
||||
@@ -21,16 +24,16 @@ def validate_experimental_sgl_marlin_server_args(
|
||||
# A provided adapter path implicitly enables LoRA later unless it was
|
||||
# explicitly disabled. No-LoRA delegates to the stock Marlin fused path.
|
||||
lora_enabled = bool(resolved_args.enable_lora) or (
|
||||
resolved_args.enable_lora is None and bool(server_args.lora_paths)
|
||||
resolved_args.enable_lora is None and bool(cfg.lora_paths)
|
||||
)
|
||||
if not lora_enabled:
|
||||
return
|
||||
|
||||
if not server_args.lora_use_virtual_experts:
|
||||
if not cfg.lora_use_virtual_experts:
|
||||
raise ValueError(
|
||||
"experimental_sgl_marlin LoRA requires --lora-use-virtual-experts"
|
||||
)
|
||||
if server_args.lora_backend != "triton":
|
||||
if cfg.lora_backend != "triton":
|
||||
# The temporary dense/sink kernels consume Triton SGEMM batch metadata
|
||||
# directly; other global backends are not adapted in this tree.
|
||||
raise ValueError("experimental_sgl_marlin LoRA requires --lora-backend triton")
|
||||
@@ -38,12 +41,12 @@ def validate_experimental_sgl_marlin_server_args(
|
||||
return
|
||||
|
||||
if (
|
||||
server_args.init_expert_location != "trivial"
|
||||
or server_args.ep_num_redundant_experts != 0
|
||||
or server_args.enable_eplb
|
||||
or server_args.elastic_ep_backend is not None
|
||||
or server_args.enable_elastic_expert_backup
|
||||
or server_args.elastic_ep_rejoin
|
||||
cfg.init_expert_location != "trivial"
|
||||
or cfg.ep_num_redundant_experts != 0
|
||||
or cfg.enable_eplb
|
||||
or cfg.elastic_ep_backend is not None
|
||||
or cfg.enable_elastic_expert_backup
|
||||
or cfg.elastic_ep_rejoin
|
||||
):
|
||||
raise ValueError(
|
||||
"experimental_sgl_marlin EP requires trivial expert placement "
|
||||
|
||||
@@ -196,10 +196,14 @@ def prepare_raw_kimi_server_args(
|
||||
server_args: Any, loader_config: dict[str, Any]
|
||||
) -> None:
|
||||
"""Resolve a raw GGUF model path into the normal loader inputs."""
|
||||
model_path = Path(server_args.model_path).expanduser()
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
|
||||
cfg = resolving_view(server_args)
|
||||
model_path = Path(cfg.model_path).expanduser()
|
||||
if not model_path.is_file() or model_path.suffix.lower() != ".gguf":
|
||||
return
|
||||
tokenizer_path = server_args.tokenizer_path
|
||||
tokenizer_path = cfg.tokenizer_path
|
||||
if tokenizer_path and Path(tokenizer_path).expanduser() == model_path:
|
||||
tokenizer_path = None
|
||||
assets = ensure_kimi_assets(
|
||||
@@ -503,7 +507,11 @@ def prepare_raw_deepseek_server_args(
|
||||
server_args: Any, loader_config: dict[str, Any]
|
||||
) -> None:
|
||||
"""Resolve a raw DeepSeek V4 GGUF into metadata and Expert Pack inputs."""
|
||||
source = Path(server_args.model_path).expanduser().resolve(strict=True)
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
|
||||
cfg = resolving_view(server_args)
|
||||
source = Path(cfg.model_path).expanduser().resolve(strict=True)
|
||||
if not source.is_file():
|
||||
return
|
||||
repo = _repo_root()
|
||||
@@ -546,7 +554,11 @@ def prepare_raw_expert_pack_server_args(
|
||||
server_args: Any, loader_config: dict[str, Any]
|
||||
) -> None:
|
||||
"""Dispatch a raw GGUF to the model-specific expert-pack preparation path."""
|
||||
source = Path(server_args.model_path).expanduser()
|
||||
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
|
||||
cfg = resolving_view(server_args)
|
||||
source = Path(cfg.model_path).expanduser()
|
||||
if not source.is_file():
|
||||
return
|
||||
name = source.name.upper()
|
||||
|
||||
@@ -29,7 +29,7 @@ import jinja2.ext
|
||||
import jinja2.nodes
|
||||
import jinja2.sandbox
|
||||
|
||||
from sglang.srt.arg_groups.overrides import declare_late_resolution
|
||||
from sglang.srt.arg_groups.overrides import declare_late_resolution, resolving_view
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -666,11 +666,12 @@ def _architecture_auto_parsers(server_args, needs: Tuple[str, ...]) -> Dict[str,
|
||||
"""The parsers the model architecture implies, for the fields still on auto."""
|
||||
from sglang.srt.utils.hf_transformers_utils import get_config
|
||||
|
||||
cfg = resolving_view(server_args)
|
||||
config = get_config(
|
||||
server_args.model_path,
|
||||
trust_remote_code=server_args.trust_remote_code,
|
||||
revision=getattr(server_args, "revision", None),
|
||||
model_config_parser=getattr(server_args, "model_config_parser", "auto"),
|
||||
cfg.model_path,
|
||||
trust_remote_code=cfg.trust_remote_code,
|
||||
revision=getattr(cfg, "revision", None),
|
||||
model_config_parser=getattr(cfg, "model_config_parser", "auto"),
|
||||
)
|
||||
architectures = getattr(config, "architectures", None) or []
|
||||
arch = architectures[0] if architectures else ""
|
||||
@@ -708,17 +709,18 @@ def resolve_auto_parsers(server_args) -> None:
|
||||
the schedulers it forks, the HTTP server, and the tokenizer workers it is
|
||||
serialized for.
|
||||
"""
|
||||
cfg = resolving_view(server_args)
|
||||
needs = tuple(
|
||||
attr
|
||||
for attr in ("reasoning_parser", "tool_call_parser")
|
||||
if getattr(server_args, attr) == "auto"
|
||||
if getattr(cfg, attr) == "auto"
|
||||
)
|
||||
if not needs:
|
||||
return
|
||||
|
||||
from sglang.srt.utils.hf_transformers_utils import get_tokenizer
|
||||
|
||||
chat_template_arg = getattr(server_args, "chat_template", None)
|
||||
chat_template_arg = getattr(cfg, "chat_template", None)
|
||||
try:
|
||||
explicit_jinja_template = _load_explicit_jinja_template(chat_template_arg)
|
||||
except Exception as e:
|
||||
@@ -731,8 +733,8 @@ def resolve_auto_parsers(server_args) -> None:
|
||||
tokenizer = None
|
||||
try:
|
||||
tokenizer = get_tokenizer(
|
||||
server_args.model_path,
|
||||
trust_remote_code=server_args.trust_remote_code,
|
||||
cfg.model_path,
|
||||
trust_remote_code=cfg.trust_remote_code,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load tokenizer for auto-detection: {e}")
|
||||
|
||||
+949
-851
File diff suppressed because it is too large
Load Diff
@@ -49,19 +49,19 @@ DEFAULT_ADAPTIVE_CONFIG: dict[str, dict] = {
|
||||
|
||||
def adaptive_unsupported_reason(server_args: ServerArgs) -> str | None:
|
||||
"""Return why adaptive spec cannot run under the given server args, or None if supported."""
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
|
||||
cfg = resolving_view(server_args)
|
||||
from sglang.srt.arg_groups.overrides import resolved_view
|
||||
|
||||
if server_args.speculative_algorithm not in ("EAGLE", "EAGLE3"):
|
||||
if cfg.speculative_algorithm not in ("EAGLE", "EAGLE3"):
|
||||
return (
|
||||
f"speculative_algorithm={server_args.speculative_algorithm} "
|
||||
f"speculative_algorithm={cfg.speculative_algorithm} "
|
||||
"(only EAGLE/EAGLE3 are supported)"
|
||||
)
|
||||
if (
|
||||
server_args.speculative_eagle_topk is not None
|
||||
and server_args.speculative_eagle_topk != 1
|
||||
):
|
||||
if cfg.speculative_eagle_topk is not None and cfg.speculative_eagle_topk != 1:
|
||||
return (
|
||||
f"speculative_eagle_topk={server_args.speculative_eagle_topk} "
|
||||
f"speculative_eagle_topk={cfg.speculative_eagle_topk} "
|
||||
"(only topk=1 is supported)"
|
||||
)
|
||||
if resolved_view(server_args).enable_dp_attention:
|
||||
@@ -74,12 +74,12 @@ def adaptive_unsupported_reason(server_args: ServerArgs) -> str | None:
|
||||
"enable_multi_layer_eagle=True is not supported "
|
||||
"(MultiLayerEagleWorkerV2 does not implement adaptive)"
|
||||
)
|
||||
if server_args.enable_two_batch_overlap:
|
||||
if cfg.enable_two_batch_overlap:
|
||||
return (
|
||||
"enable_two_batch_overlap=True is not supported "
|
||||
"(adaptive state swap would discard the TboAttnBackend wrapper)"
|
||||
)
|
||||
if server_args.enable_pdmux:
|
||||
if cfg.enable_pdmux:
|
||||
return (
|
||||
"enable_pdmux=True is not supported "
|
||||
"(adaptive state swap does not update decode_attn_backend_group)"
|
||||
|
||||
@@ -254,6 +254,9 @@ class SpeculativeAlgorithm(Enum):
|
||||
def create_worker(
|
||||
self, server_args: ServerArgs
|
||||
) -> Optional[Union[Type[BaseSpecWorker], Type[TpModelWorker], Type[NGRAMWorker]]]:
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
|
||||
cfg = resolving_view(server_args)
|
||||
assert (
|
||||
not self.is_none()
|
||||
), "Cannot create worker for NONE speculative algorithm."
|
||||
@@ -283,7 +286,7 @@ class SpeculativeAlgorithm(Enum):
|
||||
|
||||
# EAGLE / EAGLE3 / STANDALONE / MULTI_LAYER always use the V2 worker,
|
||||
# even with overlap disabled (scheduler drives it synchronously).
|
||||
if self.is_eagle() and server_args.enable_multi_layer_eagle:
|
||||
if self.is_eagle() and cfg.enable_multi_layer_eagle:
|
||||
from sglang.srt.speculative.multi_layer_eagle_worker_v2 import (
|
||||
MultiLayerEagleWorkerV2,
|
||||
)
|
||||
|
||||
@@ -108,7 +108,10 @@ class CustomSpecAlgo:
|
||||
pass
|
||||
|
||||
def create_worker(self, server_args: ServerArgs) -> Type:
|
||||
if not server_args.disable_overlap_schedule and not self.supports_overlap:
|
||||
from sglang.srt.arg_groups.overrides import resolving_view
|
||||
|
||||
cfg = resolving_view(server_args)
|
||||
if not cfg.disable_overlap_schedule and not self.supports_overlap:
|
||||
raise ValueError(
|
||||
f"Speculative algorithm {self.name} does not support overlap scheduling."
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user