diff --git a/python/sglang/srt/arg_groups/deepseek_v4_hook.py b/python/sglang/srt/arg_groups/deepseek_v4_hook.py index f57159a09..cc40788df 100644 --- a/python/sglang/srt/arg_groups/deepseek_v4_hook.py +++ b/python/sglang/srt/arg_groups/deepseek_v4_hook.py @@ -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 diff --git a/python/sglang/srt/arg_groups/hisparse_hook.py b/python/sglang/srt/arg_groups/hisparse_hook.py index f9f1197ae..135d172cf 100644 --- a/python/sglang/srt/arg_groups/hisparse_hook.py +++ b/python/sglang/srt/arg_groups/hisparse_hook.py @@ -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 diff --git a/python/sglang/srt/arg_groups/nemotron_h_hook.py b/python/sglang/srt/arg_groups/nemotron_h_hook.py index a4e17670f..191c8bca2 100644 --- a/python/sglang/srt/arg_groups/nemotron_h_hook.py +++ b/python/sglang/srt/arg_groups/nemotron_h_hook.py @@ -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 [ diff --git a/python/sglang/srt/arg_groups/pd_disaggregation_hook.py b/python/sglang/srt/arg_groups/pd_disaggregation_hook.py index 43bd301b9..9e90322c4 100644 --- a/python/sglang/srt/arg_groups/pd_disaggregation_hook.py +++ b/python/sglang/srt/arg_groups/pd_disaggregation_hook.py @@ -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 diff --git a/python/sglang/srt/arg_groups/speculative_hook.py b/python/sglang/srt/arg_groups/speculative_hook.py index 690b684f9..f04ba2e09 100644 --- a/python/sglang/srt/arg_groups/speculative_hook.py +++ b/python/sglang/srt/arg_groups/speculative_hook.py @@ -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. diff --git a/python/sglang/srt/speculative/spec_info.py b/python/sglang/srt/speculative/spec_info.py index f3172bc0d..f891c15fa 100644 --- a/python/sglang/srt/speculative/spec_info.py +++ b/python/sglang/srt/speculative/spec_info.py @@ -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: diff --git a/python/sglang/srt/speculative/spec_registry.py b/python/sglang/srt/speculative/spec_registry.py index fa0becafb..ce3f3178f 100644 --- a/python/sglang/srt/speculative/spec_registry.py +++ b/python/sglang/srt/speculative/spec_registry.py @@ -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( diff --git a/test/registered/unit/spec/test_spec_registry.py b/test/registered/unit/spec/test_spec_registry.py index fe08b45db..141b854c4 100644 --- a/test/registered/unit/spec/test_spec_registry.py +++ b/test/registered/unit/spec/test_spec_registry.py @@ -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."""