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(
@@ -1,8 +1,10 @@
"""Unit tests for the speculative algorithm plugin registry."""
import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock
from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_registry import (
_REGISTRY,
@@ -207,6 +209,39 @@ class TestValidatorHook(_RegistryIsolated):
validator.assert_not_called()
class TestServerArgsHook(_RegistryIsolated):
def test_handle_speculative_decoding_invokes_custom_handle_server_args(self):
class CustomHandleServerArgs(CustomSpecAlgo):
def handle_server_args(self, server_args):
server_args.custom_spec_handle_seen = self.name
server_args.speculative_num_draft_tokens = 7
@SpeculativeAlgorithm.register(
"MY_HANDLE_ARGS", supports_overlap=True, spec_class=CustomHandleServerArgs
)
def _factory(server_args):
return MagicMock
server_args = SimpleNamespace(
speculative_draft_model_path=None,
speculative_draft_model_revision=None,
speculative_moe_runner_backend=None,
moe_runner_backend="auto",
speculative_algorithm="my_handle_args",
decrypted_draft_config_file=None,
trust_remote_code=False,
speculative_draft_window_size=None,
speculative_skip_dp_mlp_sync=False,
speculative_adaptive=False,
)
handle_speculative_decoding(server_args)
self.assertEqual(server_args.speculative_algorithm, "MY_HANDLE_ARGS")
self.assertEqual(server_args.custom_spec_handle_seen, "MY_HANDLE_ARGS")
self.assertEqual(server_args.speculative_num_draft_tokens, 7)
class TestSubclassOverride(_RegistryIsolated):
"""Plugins can subclass CustomSpecAlgo to override is_*() / create_worker."""