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 import logging
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
@@ -7,7 +9,7 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__) 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.""" """Apply DeepSeek V4 model-specific server arg defaults and constraints."""
from sglang.srt.server_args import ServerArgs 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.""" """Validate DeepSeek V4 context-parallel configuration."""
if not server_args.enable_prefill_cp: if not server_args.enable_prefill_cp:
return return
@@ -1,3 +1,5 @@
from __future__ import annotations
import logging import logging
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
@@ -21,7 +23,7 @@ def _hisparse_default_backend(kv_cache_dtype: str) -> str:
def apply_hisparse_dsa_backend_defaults( def apply_hisparse_dsa_backend_defaults(
server_args: "ServerArgs", server_args: ServerArgs,
user_set_prefill: bool, user_set_prefill: bool,
user_set_decode: bool, user_set_decode: bool,
kv_cache_dtype: str, kv_cache_dtype: str,
@@ -46,7 +48,7 @@ def apply_hisparse_dsa_backend_defaults(
return True 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).""" """Validate --enable-hisparse constraints (model class, radix cache, DSA backend)."""
if not server_args.enable_hisparse: if not server_args.enable_hisparse:
return return
@@ -1,3 +1,5 @@
from __future__ import annotations
import logging import logging
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
@@ -9,7 +11,7 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__) 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.""" """Apply NemotronH model-specific server arg defaults and constraints."""
model_config = server_args.get_model_config() model_config = server_args.get_model_config()
is_modelopt = model_config.quantization in [ is_modelopt = model_config.quantization in [
@@ -1,3 +1,5 @@
from __future__ import annotations
import logging import logging
import os import os
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
@@ -10,7 +12,7 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__) 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.""" """Validate and normalize PD-disaggregation server args."""
# "mooncake_tcp" is mooncake with the TCP transport forced: set MC_FORCE_TCP # "mooncake_tcp" is mooncake with the TCP transport forced: set MC_FORCE_TCP
# so mooncake installs TcpTransport instead of RDMA, rewrite the backend to # so mooncake installs TcpTransport instead of RDMA, rewrite the backend to
@@ -1,3 +1,5 @@
from __future__ import annotations
import json import json
import logging import logging
import os import os
@@ -49,7 +51,7 @@ def _resolve_speculative_algorithm_alias(
return speculative_algorithm return speculative_algorithm
def handle_speculative_decoding(server_args: "ServerArgs") -> None: def handle_speculative_decoding(server_args: ServerArgs) -> None:
if ( if (
server_args.speculative_draft_model_path is not None server_args.speculative_draft_model_path is not None
and server_args.speculative_draft_model_revision is 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, server_args.speculative_algorithm,
) )
algo = None
if server_args.speculative_algorithm is not None: if server_args.speculative_algorithm is not None:
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_registry import CustomSpecAlgo 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: if server_args.speculative_adaptive:
_init_adaptive_speculative_params(server_args) _init_adaptive_speculative_params(server_args)
if server_args.speculative_algorithm == "DFLASH": if algo is not None:
_handle_dflash(server_args) algo.handle_server_args(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)
def _handle_dflash(server_args: "ServerArgs") -> None: def _handle_dflash(server_args: ServerArgs) -> None:
if server_args.enable_dp_attention: if server_args.enable_dp_attention:
raise ValueError( raise ValueError(
"Currently DFLASH speculative decoding does not support dp attention." "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: if server_args.max_running_requests is None:
server_args.max_running_requests = 48 server_args.max_running_requests = 48
logger.warning( 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 ( if (
server_args.speculative_algorithm == "STANDALONE" server_args.speculative_algorithm == "STANDALONE"
and server_args.enable_dp_attention 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"): if not server_args.device.startswith("cuda"):
raise ValueError("Ngram speculative decoding only supports CUDA device.") 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 ( from sglang.srt.speculative.adaptive_spec_params import (
adaptive_unsupported_reason, adaptive_unsupported_reason,
) )
@@ -450,7 +447,7 @@ def _maybe_disable_adaptive(server_args: "ServerArgs") -> None:
server_args.speculative_adaptive = False 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 ( from sglang.srt.speculative.adaptive_spec_params import (
resolve_candidate_steps_from_config, 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 server_args.speculative_num_draft_tokens = server_args.speculative_num_steps + 1
def _auto_choose_speculative_params( def _auto_choose_speculative_params(server_args: ServerArgs, model_arch: str) -> tuple:
server_args: "ServerArgs", model_arch: str
) -> tuple:
""" """
Automatically choose the parameters for speculative decoding. Automatically choose the parameters for speculative decoding.
@@ -159,6 +159,27 @@ class SpeculativeAlgorithm(Enum):
def need_topk(self) -> bool: def need_topk(self) -> bool:
return self.is_eagle() or self.is_standalone() 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( def get_num_tokens_per_bs_for_target_verify(
self, num_draft_tokens: int, is_draft_worker: bool self, num_draft_tokens: int, is_draft_worker: bool
) -> int: ) -> int:
@@ -89,6 +89,9 @@ class CustomSpecAlgo:
# Conservative default: the larger KV reserve. # Conservative default: the larger KV reserve.
return True return True
def handle_server_args(self, server_args: ServerArgs) -> None:
pass
def create_worker(self, server_args: ServerArgs) -> Type: def create_worker(self, server_args: ServerArgs) -> Type:
if not server_args.disable_overlap_schedule and not self.supports_overlap: if not server_args.disable_overlap_schedule and not self.supports_overlap:
raise ValueError( raise ValueError(
@@ -1,8 +1,10 @@
"""Unit tests for the speculative algorithm plugin registry.""" """Unit tests for the speculative algorithm plugin registry."""
import unittest import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock 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_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_registry import ( from sglang.srt.speculative.spec_registry import (
_REGISTRY, _REGISTRY,
@@ -207,6 +209,39 @@ class TestValidatorHook(_RegistryIsolated):
validator.assert_not_called() 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): class TestSubclassOverride(_RegistryIsolated):
"""Plugins can subclass CustomSpecAlgo to override is_*() / create_worker.""" """Plugins can subclass CustomSpecAlgo to override is_*() / create_worker."""