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