Custom spec algorithm can handle server args (#28162)

This commit is contained in:
jasonjk-park
2026-06-16 17:13:52 -07:00
committed by GitHub
parent ca84d52b78
commit d86a7e7018
8 changed files with 86 additions and 24 deletions
@@ -1,3 +1,5 @@
from __future__ import annotations
import logging
from typing import TYPE_CHECKING
@@ -7,7 +9,7 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__)
def apply_deepseek_v4_defaults(server_args: "ServerArgs", model_arch: str) -> None:
def apply_deepseek_v4_defaults(server_args: ServerArgs, model_arch: str) -> None:
"""Apply DeepSeek V4 model-specific server arg defaults and constraints."""
from sglang.srt.server_args import ServerArgs
@@ -47,7 +49,7 @@ def apply_deepseek_v4_defaults(server_args: "ServerArgs", model_arch: str) -> No
)
def validate_deepseek_v4_cp(server_args: "ServerArgs") -> None:
def validate_deepseek_v4_cp(server_args: ServerArgs) -> None:
"""Validate DeepSeek V4 context-parallel configuration."""
if not server_args.enable_prefill_cp:
return
@@ -1,3 +1,5 @@
from __future__ import annotations
import logging
from typing import TYPE_CHECKING
@@ -21,7 +23,7 @@ def _hisparse_default_backend(kv_cache_dtype: str) -> str:
def apply_hisparse_dsa_backend_defaults(
server_args: "ServerArgs",
server_args: ServerArgs,
user_set_prefill: bool,
user_set_decode: bool,
kv_cache_dtype: str,
@@ -46,7 +48,7 @@ def apply_hisparse_dsa_backend_defaults(
return True
def validate_hisparse(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:
return
@@ -1,3 +1,5 @@
from __future__ import annotations
import logging
from typing import TYPE_CHECKING
@@ -9,7 +11,7 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__)
def apply_nemotron_h_defaults(server_args: "ServerArgs", model_arch: str) -> None:
def apply_nemotron_h_defaults(server_args: ServerArgs, model_arch: str) -> None:
"""Apply NemotronH model-specific server arg defaults and constraints."""
model_config = server_args.get_model_config()
is_modelopt = model_config.quantization in [
@@ -1,3 +1,5 @@
from __future__ import annotations
import logging
import os
from typing import TYPE_CHECKING
@@ -10,7 +12,7 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__)
def handle_pd_disaggregation(server_args: "ServerArgs") -> None:
def handle_pd_disaggregation(server_args: ServerArgs) -> None:
"""Validate and normalize PD-disaggregation 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
@@ -1,3 +1,5 @@
from __future__ import annotations
import json
import logging
import os
@@ -49,7 +51,7 @@ def _resolve_speculative_algorithm_alias(
return speculative_algorithm
def handle_speculative_decoding(server_args: "ServerArgs") -> None:
def handle_speculative_decoding(server_args: ServerArgs) -> None:
if (
server_args.speculative_draft_model_path is not None
and server_args.speculative_draft_model_revision is None
@@ -100,6 +102,7 @@ def handle_speculative_decoding(server_args: "ServerArgs") -> None:
server_args.speculative_algorithm,
)
algo = None
if server_args.speculative_algorithm is not None:
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_registry import CustomSpecAlgo
@@ -121,17 +124,11 @@ def handle_speculative_decoding(server_args: "ServerArgs") -> None:
if server_args.speculative_adaptive:
_init_adaptive_speculative_params(server_args)
if server_args.speculative_algorithm == "DFLASH":
_handle_dflash(server_args)
elif server_args.speculative_algorithm == "FROZEN_KV_MTP":
_handle_frozen_kv_mtp(server_args)
elif server_args.speculative_algorithm in ("EAGLE", "EAGLE3", "STANDALONE"):
_handle_eagle_family(server_args)
elif server_args.speculative_algorithm == "NGRAM":
_handle_ngram(server_args)
if algo is not None:
algo.handle_server_args(server_args)
def _handle_dflash(server_args: "ServerArgs") -> None:
def _handle_dflash(server_args: ServerArgs) -> None:
if server_args.enable_dp_attention:
raise ValueError(
"Currently DFLASH speculative decoding does not support dp attention."
@@ -245,7 +242,7 @@ def _handle_dflash(server_args: "ServerArgs") -> None:
)
def _handle_frozen_kv_mtp(server_args: "ServerArgs") -> None:
def _handle_frozen_kv_mtp(server_args: ServerArgs) -> None:
if server_args.max_running_requests is None:
server_args.max_running_requests = 48
logger.warning(
@@ -260,7 +257,7 @@ def _handle_frozen_kv_mtp(server_args: "ServerArgs") -> None:
)
def _handle_eagle_family(server_args: "ServerArgs") -> None:
def _handle_eagle_family(server_args: ServerArgs) -> None:
if (
server_args.speculative_algorithm == "STANDALONE"
and server_args.enable_dp_attention
@@ -371,7 +368,7 @@ def _handle_eagle_family(server_args: "ServerArgs") -> None:
)
def _handle_ngram(server_args: "ServerArgs") -> None:
def _handle_ngram(server_args: ServerArgs) -> None:
if not server_args.device.startswith("cuda"):
raise ValueError("Ngram speculative decoding only supports CUDA device.")
@@ -436,7 +433,7 @@ def _handle_ngram(server_args: "ServerArgs") -> None:
)
def _maybe_disable_adaptive(server_args: "ServerArgs") -> None:
def _maybe_disable_adaptive(server_args: ServerArgs) -> None:
from sglang.srt.speculative.adaptive_spec_params import (
adaptive_unsupported_reason,
)
@@ -450,7 +447,7 @@ def _maybe_disable_adaptive(server_args: "ServerArgs") -> None:
server_args.speculative_adaptive = False
def _init_adaptive_speculative_params(server_args: "ServerArgs") -> None:
def _init_adaptive_speculative_params(server_args: ServerArgs) -> None:
from sglang.srt.speculative.adaptive_spec_params import (
resolve_candidate_steps_from_config,
)
@@ -475,9 +472,7 @@ def _init_adaptive_speculative_params(server_args: "ServerArgs") -> None:
server_args.speculative_num_draft_tokens = server_args.speculative_num_steps + 1
def _auto_choose_speculative_params(
server_args: "ServerArgs", model_arch: str
) -> tuple:
def _auto_choose_speculative_params(server_args: ServerArgs, model_arch: str) -> tuple:
"""
Automatically choose the parameters for speculative decoding.
@@ -159,6 +159,27 @@ class SpeculativeAlgorithm(Enum):
def need_topk(self) -> bool:
return self.is_eagle() or self.is_standalone()
def handle_server_args(self, server_args: ServerArgs) -> None:
"""Hook for per-algorithm server args mutation.
In-place updated.
"""
from sglang.srt.arg_groups.speculative_hook import (
_handle_dflash,
_handle_eagle_family,
_handle_frozen_kv_mtp,
_handle_ngram,
)
if self.is_dflash():
_handle_dflash(server_args)
elif self.is_frozen_kv_mtp():
_handle_frozen_kv_mtp(server_args)
elif self.is_eagle() or self.is_standalone():
_handle_eagle_family(server_args)
elif self.is_ngram():
_handle_ngram(server_args)
def get_num_tokens_per_bs_for_target_verify(
self, num_draft_tokens: int, is_draft_worker: bool
) -> int:
@@ -89,6 +89,9 @@ class CustomSpecAlgo:
# Conservative default: the larger KV reserve.
return True
def handle_server_args(self, server_args: ServerArgs) -> None:
pass
def create_worker(self, server_args: ServerArgs) -> Type:
if not server_args.disable_overlap_schedule and not self.supports_overlap:
raise ValueError(