Custom spec algorithm can handle server args (#28162)
This commit is contained in:
@@ -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."""
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user