From aec4022e58c6f74c72f9b7be449ad3e81238c14e Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Sat, 16 May 2026 03:23:48 -0700 Subject: [PATCH] [Spec] Clean up draft-window-size handling; extract spec arg setup to arg_groups (#25424) --- .../sglang/srt/arg_groups/argparse_actions.py | 82 +++ .../sglang/srt/arg_groups/speculative_hook.py | 443 ++++++++++++++++ python/sglang/srt/models/llama_eagle3.py | 27 +- python/sglang/srt/server_args.py | 494 +----------------- .../sglang/srt/speculative/dflash_worker.py | 5 +- .../srt/speculative/frozen_kv_mtp_worker.py | 2 +- .../unit/server_args/test_server_args.py | 23 +- 7 files changed, 574 insertions(+), 502 deletions(-) create mode 100644 python/sglang/srt/arg_groups/argparse_actions.py create mode 100644 python/sglang/srt/arg_groups/speculative_hook.py diff --git a/python/sglang/srt/arg_groups/argparse_actions.py b/python/sglang/srt/arg_groups/argparse_actions.py new file mode 100644 index 000000000..5540a10a7 --- /dev/null +++ b/python/sglang/srt/arg_groups/argparse_actions.py @@ -0,0 +1,82 @@ +import argparse +import json +import logging + +logger = logging.getLogger(__name__) + + +class LoRAPathAction(argparse.Action): + def __call__(self, parser, namespace, values, option_string=None): + lora_paths = [] + if values: + assert isinstance(values, list), "Expected a list of LoRA paths." + for lora_path in values: + lora_path = lora_path.strip() + if lora_path.startswith("{") and lora_path.endswith("}"): + obj = json.loads(lora_path) + assert "lora_path" in obj and "lora_name" in obj, ( + f"{repr(lora_path)} looks like a JSON str, " + "but it does not contain 'lora_name' and 'lora_path' keys." + ) + lora_paths.append(obj) + else: + lora_paths.append(lora_path) + + setattr(namespace, self.dest, lora_paths) + + +def print_deprecated_warning(message: str): + logger.warning(f"\033[1;33m{message}\033[0m") + + +class DeprecatedAction(argparse.Action): + def __init__(self, option_strings, dest, nargs=0, **kwargs): + super(DeprecatedAction, self).__init__( + option_strings, dest, nargs=nargs, **kwargs + ) + + def __call__(self, parser, namespace, values, option_string=None): + print_deprecated_warning( + f"The command line argument '{option_string}' is deprecated and will be removed in future versions." + ) + + +class DeprecatedStoreTrueAction(argparse.Action): + """Deprecated flag that still stores True and prints a warning.""" + + def __init__( + self, + option_strings, + dest, + new_flag=None, + nargs=0, + const=True, + default=False, + **kwargs, + ): + self.new_flag = new_flag + super().__init__( + option_strings, dest, nargs=nargs, const=const, default=default, **kwargs + ) + + def __call__(self, parser, namespace, values, option_string=None): + replacement = f" Use '{self.new_flag}' instead." if self.new_flag else "" + print_deprecated_warning( + f"'{option_string}' is deprecated and will be removed in a future release.{replacement}" + ) + setattr(namespace, self.dest, True) + + +class DeprecatedAliasStoreAction(argparse.Action): + """Deprecated alias that stores its value and prints a warning.""" + + def __init__(self, option_strings, dest, new_flag=None, **kwargs): + self.new_flag = new_flag + super().__init__(option_strings, dest, **kwargs) + + def __call__(self, parser, namespace, values, option_string=None): + replacement = f" Use '{self.new_flag}' instead." if self.new_flag else "" + print_deprecated_warning( + f"'{option_string}' is deprecated and will be removed in a future release.{replacement}" + ) + setattr(namespace, self.dest, values) diff --git a/python/sglang/srt/arg_groups/speculative_hook.py b/python/sglang/srt/arg_groups/speculative_hook.py new file mode 100644 index 000000000..c1f720062 --- /dev/null +++ b/python/sglang/srt/arg_groups/speculative_hook.py @@ -0,0 +1,443 @@ +import json +import logging +from typing import TYPE_CHECKING, Optional + +from sglang.srt.environ import envs + +if TYPE_CHECKING: + from sglang.srt.server_args import ServerArgs + +logger = logging.getLogger(__name__) + + +def _resolve_speculative_algorithm_alias( + speculative_algorithm: Optional[str], + speculative_draft_model_path: Optional[str], + trust_remote_code: bool = False, +) -> Optional[str]: + """Resolve CLI speculative algorithm; NEXTN/EAGLE may become FROZEN_KV_MTP for Gemma4 assistant drafts.""" + + is_gemma4_draft = False + if speculative_draft_model_path: + from sglang.srt.utils.hf_transformers_utils import get_config + + cfg = get_config( + speculative_draft_model_path, trust_remote_code=trust_remote_code + ) + is_gemma4_draft = "Gemma4AssistantForCausalLM" in ( + getattr(cfg, "architectures", None) or [] + ) + + if speculative_algorithm == "EAGLE3" and is_gemma4_draft: + raise ValueError( + "Gemma4AssistantForCausalLM draft requires " + "--speculative-algorithm NEXTN or EAGLE; EAGLE3 is " + "not supported for this draft architecture." + ) + + if speculative_algorithm == "NEXTN" or speculative_algorithm == "EAGLE": + if is_gemma4_draft: + logger.info( + "Detected Gemma4AssistantForCausalLM draft; " + f"promoting --speculative-algorithm {speculative_algorithm} to FROZEN_KV_MTP." + ) + return "FROZEN_KV_MTP" + return "EAGLE" + + return speculative_algorithm + + +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 + ): + server_args.speculative_draft_model_revision = "main" + + if server_args.speculative_moe_runner_backend is None: + server_args.speculative_moe_runner_backend = server_args.moe_runner_backend + + if server_args.speculative_algorithm is not None: + server_args.speculative_algorithm = server_args.speculative_algorithm.upper() + + server_args.speculative_algorithm = _resolve_speculative_algorithm_alias( + server_args.speculative_algorithm, + server_args.speculative_draft_model_path, + trust_remote_code=server_args.trust_remote_code, + ) + + # Validate --speculative-draft-window-size once, regardless of algorithm. + # Consumed by DFLASH (compact draft KV cache) and Llama EAGLE-3 (drafter attention SWA). + if server_args.speculative_draft_window_size is not None: + window_size = int(server_args.speculative_draft_window_size) + if window_size <= 0: + raise ValueError( + f"--speculative-draft-window-size must be positive, got {window_size}." + ) + server_args.speculative_draft_window_size = window_size + if server_args.speculative_algorithm not in ("EAGLE3", "DFLASH"): + logger.warning( + "--speculative-draft-window-size has no effect with " + "speculative_algorithm=%s (honored by Llama EAGLE-3 and DFLASH only).", + server_args.speculative_algorithm, + ) + + if server_args.speculative_algorithm is not None: + from sglang.srt.speculative.spec_info import SpeculativeAlgorithm + from sglang.srt.speculative.spec_registry import CustomSpecAlgo + + algo = SpeculativeAlgorithm.from_string(server_args.speculative_algorithm) + + # TODO: move the per-algorithm validation below into spec module hooks. + if isinstance(algo, CustomSpecAlgo) and algo.validate_server_args is not None: + algo.validate_server_args(server_args) + + if server_args.speculative_skip_dp_mlp_sync: + assert server_args.speculative_algorithm == "EAGLE", ( + "--speculative-skip-dp-mlp-sync is only supported with " + f"speculative_algorithm == EAGLE, got {server_args.speculative_algorithm}." + ) + + 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 server_args.speculative_adaptive: + _maybe_disable_adaptive(server_args) + + +def _handle_dflash(server_args: "ServerArgs") -> None: + if server_args.enable_dp_attention: + raise ValueError( + "Currently DFLASH speculative decoding does not support dp attention." + ) + + if server_args.pp_size != 1: + raise ValueError( + "Currently DFLASH speculative decoding only supports pp_size == 1." + ) + + if server_args.speculative_draft_model_path is None: + raise ValueError( + "DFLASH speculative decoding requires setting --speculative-draft-model-path." + ) + + # DFLASH does not use EAGLE-style `num_steps`/`topk`, but those fields still + # affect generic scheduler/KV-cache accounting (buffer sizing, KV freeing, + # RoPE reservation). Force them to 1 to avoid surprising memory behavior. + # + # For DFlash, the natural unit is `block_size` (verify window length). + if server_args.speculative_num_steps is None: + server_args.speculative_num_steps = 1 + elif int(server_args.speculative_num_steps) != 1: + logger.warning( + "DFLASH only supports speculative_num_steps == 1; overriding speculative_num_steps=%s to 1.", + server_args.speculative_num_steps, + ) + server_args.speculative_num_steps = 1 + + if server_args.speculative_eagle_topk is None: + server_args.speculative_eagle_topk = 1 + elif int(server_args.speculative_eagle_topk) != 1: + logger.warning( + "DFLASH only supports speculative_eagle_topk == 1; overriding speculative_eagle_topk=%s to 1.", + server_args.speculative_eagle_topk, + ) + server_args.speculative_eagle_topk = 1 + + if server_args.speculative_dflash_block_size is not None: + if int(server_args.speculative_dflash_block_size) <= 0: + raise ValueError( + "DFLASH requires --speculative-dflash-block-size to be positive, " + f"got {server_args.speculative_dflash_block_size}." + ) + if server_args.speculative_num_draft_tokens is not None and int( + server_args.speculative_num_draft_tokens + ) != int(server_args.speculative_dflash_block_size): + raise ValueError( + "Both --speculative-num-draft-tokens and --speculative-dflash-block-size are set " + "but they differ. For DFLASH they must match. " + f"speculative_num_draft_tokens={server_args.speculative_num_draft_tokens}, " + f"speculative_dflash_block_size={server_args.speculative_dflash_block_size}." + ) + server_args.speculative_num_draft_tokens = int( + server_args.speculative_dflash_block_size + ) + + if server_args.speculative_num_draft_tokens is None: + from sglang.srt.speculative.dflash_utils import ( + parse_dflash_draft_config, + ) + + model_override_args = json.loads(server_args.json_model_override_args) + inferred_block_size = None + try: + from sglang.srt.utils.hf_transformers_utils import get_config + + draft_hf_config = get_config( + server_args.speculative_draft_model_path, + trust_remote_code=server_args.trust_remote_code, + revision=server_args.speculative_draft_model_revision, + model_override_args=model_override_args, + ) + inferred_block_size = parse_dflash_draft_config( + draft_hf_config=draft_hf_config + ).resolve_block_size(default=None) + except Exception as e: + logger.warning( + "Failed to infer DFLASH block_size from draft model config; " + "defaulting speculative_num_draft_tokens to 16. Error: %s", + e, + ) + + if inferred_block_size is None: + inferred_block_size = 16 + logger.warning( + "speculative_num_draft_tokens is not set; defaulting to %d for DFLASH.", + inferred_block_size, + ) + server_args.speculative_num_draft_tokens = inferred_block_size + + if server_args.speculative_draft_window_size is not None: + draft_tokens = int(server_args.speculative_num_draft_tokens) + if server_args.speculative_draft_window_size < draft_tokens: + raise ValueError( + "--speculative-draft-window-size must be >= " + "--speculative-num-draft-tokens (block_size). " + f"window_size={server_args.speculative_draft_window_size}, block_size={draft_tokens}." + ) + + if server_args.max_running_requests is None: + server_args.max_running_requests = 48 + logger.warning( + "Max running requests is reset to 48 for speculative decoding. You can override this by explicitly setting --max-running-requests." + ) + + server_args.disable_overlap_schedule = True + logger.warning( + "Overlap scheduler is disabled when using DFLASH speculative decoding (spec v2 is not supported yet)." + ) + + if server_args.enable_mixed_chunk: + server_args.enable_mixed_chunk = False + logger.warning( + "Mixed chunked prefill is disabled because of using dflash speculative decoding." + ) + + +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( + "Max running requests is reset to 48 for speculative decoding. You can override this by explicitly setting --max-running-requests." + ) + + server_args.disable_overlap_schedule = True + logger.warning( + "Overlap scheduler is disabled when using Frozen-KV MTP speculative decoding (spec v2 is not supported yet)." + ) + + if server_args.enable_mixed_chunk: + server_args.enable_mixed_chunk = False + logger.warning( + "Mixed chunked prefill is disabled because of using " + "Frozen-KV MTP speculative decoding." + ) + + +def _handle_eagle_family(server_args: "ServerArgs") -> None: + if ( + server_args.speculative_algorithm == "STANDALONE" + and server_args.enable_dp_attention + ): + # TODO: support dp attention for standalone speculative decoding + raise ValueError( + "Currently standalone speculative decoding does not support dp attention." + ) + + if server_args.max_running_requests is None: + server_args.max_running_requests = 48 + logger.warning( + "Max running requests is reset to 48 for speculative decoding. You can override this by explicitly setting --max-running-requests." + ) + + spec_v1_reason = None + if ( + server_args.speculative_eagle_topk is not None + and server_args.speculative_eagle_topk > 1 + and not server_args.disable_overlap_schedule + ): + server_args.disable_overlap_schedule = True + spec_v1_reason = "spec v2 currently only supports topk = 1" + elif ( + not envs.SGLANG_ENABLE_SPEC_V2.get() + and not server_args.disable_overlap_schedule + ): + server_args.disable_overlap_schedule = True + spec_v1_reason = "SGLANG_ENABLE_SPEC_V2=False" + + if server_args.disable_overlap_schedule: + logger.warning( + "Spec v1 is used for eagle/eagle3/standalone speculative decoding because %s.", + spec_v1_reason or "overlap schedule is disabled", + ) + else: + logger.warning( + "Spec v2 is enabled by default for eagle/eagle3/standalone speculative decoding." + ) + + if server_args.enable_mixed_chunk: + server_args.enable_mixed_chunk = False + logger.warning( + "Mixed chunked prefill is disabled because of using " + "eagle speculative decoding." + ) + + model_arch = server_args.get_model_config().hf_config.architectures[0] + if model_arch in [ + "DeepseekV32ForCausalLM", + "DeepseekV3ForCausalLM", + "DeepseekV4ForCausalLM", + "Glm4MoeForCausalLM", + "Glm4MoeLiteForCausalLM", + "GlmMoeDsaForCausalLM", + "BailingMoeForCausalLM", + "BailingMoeV2ForCausalLM", + "BailingMoeV2_5ForCausalLM", + "MistralLarge3ForCausalLM", + "PixtralForConditionalGeneration", + "HYV3ForCausalLM", + ]: + if server_args.speculative_draft_model_path is None: + server_args.speculative_draft_model_path = server_args.model_path + server_args.speculative_draft_model_revision = server_args.revision + else: + if model_arch not in [ + "MistralLarge3ForCausalLM", + "PixtralForConditionalGeneration", + ]: + logger.warning( + "DeepSeek MTP does not require setting speculative_draft_model_path." + ) + + if server_args.speculative_num_steps is None: + assert ( + server_args.speculative_eagle_topk is None + and server_args.speculative_num_draft_tokens is None + ) + from sglang.srt.server_args import auto_choose_speculative_params + + ( + server_args.speculative_num_steps, + server_args.speculative_eagle_topk, + server_args.speculative_num_draft_tokens, + ) = auto_choose_speculative_params(server_args) + + if ( + server_args.attention_backend == "trtllm_mha" + or server_args.decode_attention_backend == "trtllm_mha" + or server_args.prefill_attention_backend == "trtllm_mha" + ): + if server_args.speculative_eagle_topk > 1: + raise ValueError( + "trtllm_mha backend only supports topk = 1 for speculative decoding." + ) + + if ( + server_args.speculative_eagle_topk == 1 + and server_args.speculative_num_draft_tokens + != server_args.speculative_num_steps + 1 + ): + logger.warning( + "speculative_num_draft_tokens is adjusted to speculative_num_steps + 1 when speculative_eagle_topk == 1" + ) + server_args.speculative_num_draft_tokens = server_args.speculative_num_steps + 1 + + if ( + server_args.speculative_eagle_topk > 1 + and server_args.page_size > 1 + and server_args.attention_backend not in ["flashinfer", "fa3"] + ): + raise ValueError( + "speculative_eagle_topk > 1 with page_size > 1 is unstable and produces incorrect results for paged attention backends. This combination is only supported for the 'flashinfer' backend." + ) + + +def _handle_ngram(server_args: "ServerArgs") -> None: + if not server_args.device.startswith("cuda"): + raise ValueError("Ngram speculative decoding only supports CUDA device.") + + if server_args.max_running_requests is None: + server_args.max_running_requests = 48 + logger.warning( + "Max running requests is reset to 48 for speculative decoding. You can override this by explicitly setting --max-running-requests." + ) + + server_args.disable_overlap_schedule = True + server_args.enable_mixed_chunk = False + server_args.speculative_eagle_topk = server_args.speculative_ngram_max_bfs_breadth + if server_args.speculative_num_draft_tokens is None: + server_args.speculative_num_draft_tokens = 12 + logger.warning( + "speculative_num_draft_tokens is set to 12 by default for ngram speculative decoding. " + "You can override this by explicitly setting --speculative-num-draft-tokens." + ) + if server_args.speculative_ngram_external_corpus_path is not None: + if server_args.speculative_ngram_external_sam_budget <= 0: + raise ValueError( + "--speculative-ngram-external-sam-budget must be positive when " + "--speculative-ngram-external-corpus-path is set." + ) + if server_args.speculative_ngram_external_corpus_max_tokens <= 0: + raise ValueError( + "--speculative-ngram-external-corpus-max-tokens must be positive when " + "--speculative-ngram-external-corpus-path is set." + ) + if ( + server_args.speculative_ngram_external_sam_budget + > server_args.speculative_num_draft_tokens - 1 + ): + raise ValueError( + "speculative_ngram_external_sam_budget must be less than or equal to " + f"speculative_num_draft_tokens - 1 ({server_args.speculative_num_draft_tokens - 1})." + ) + logger.warning( + "The overlap scheduler and mixed chunked prefill are disabled because of " + "using ngram speculative decoding." + ) + + if ( + server_args.speculative_eagle_topk > 1 + and server_args.page_size > 1 + and server_args.attention_backend != "flashinfer" + ): + raise ValueError( + f"speculative_eagle_topk({server_args.speculative_eagle_topk}) > 1 " + f"with page_size({server_args.page_size}) > 1 is unstable " + "and produces incorrect results for paged attention backends. " + "This combination is only supported for the 'flashinfer' backend." + ) + if server_args.enable_dp_attention: + # TODO: support dp attention for ngram speculative decoding + raise ValueError( + "Currently ngram speculative decoding does not support dp attention." + ) + + +def _maybe_disable_adaptive(server_args: "ServerArgs") -> None: + from sglang.srt.speculative.adaptive_spec_params import ( + adaptive_unsupported_reason, + ) + + reason = adaptive_unsupported_reason(server_args) + if reason is not None: + logger.warning( + f"speculative_adaptive disabled: {reason}. " + "Falling back to static speculative params." + ) + server_args.speculative_adaptive = False diff --git a/python/sglang/srt/models/llama_eagle3.py b/python/sglang/srt/models/llama_eagle3.py index cfaf7f525..1b2531c76 100644 --- a/python/sglang/srt/models/llama_eagle3.py +++ b/python/sglang/srt/models/llama_eagle3.py @@ -47,7 +47,6 @@ class LlamaDecoderLayer(LlamaDecoderLayer): config: LlamaConfig, layer_id: int = 0, quant_config: Optional[QuantizationConfig] = None, - draft_window_size: Optional[int] = None, prefix: str = "", ) -> None: super().__init__(config, layer_id, quant_config, prefix) @@ -67,9 +66,6 @@ class LlamaDecoderLayer(LlamaDecoderLayer): prefix=add_prefix("qkv_proj", prefix), ) - if draft_window_size is not None: - self.self_attn.attn.sliding_window_size = draft_window_size - if config.model_type == "llama4_text": inter_size = config.intermediate_size_mlp else: @@ -120,7 +116,6 @@ class LlamaModel(nn.Module): self, config: LlamaConfig, quant_config: Optional[QuantizationConfig] = None, - draft_window_size: Optional[int] = None, prefix: str = "", ) -> None: super().__init__() @@ -179,7 +174,7 @@ class LlamaModel(nn.Module): self.layers = nn.ModuleList( [ - LlamaDecoderLayer(config, i, quant_config, draft_window_size, prefix) + LlamaDecoderLayer(config, i, quant_config, prefix) for i in range(config.num_hidden_layers) ] ) @@ -259,12 +254,20 @@ class LlamaForCausalLMEagle3(LlamaForCausalLM): self.quant_config = quant_config self.pp_group = get_pp_group() + # Cache draft SWA size from server args once; consumed both by the post-init + # attention patch below and by `get_attention_sliding_window_size` later. + self._draft_window_size: Optional[int] = ( + get_global_server_args().speculative_draft_window_size + ) + self.model = LlamaModel( config, quant_config=quant_config, - draft_window_size=self.get_attention_sliding_window_size(), prefix=add_prefix("model", prefix), ) + if self._draft_window_size is not None: + for layer in self.model.layers: + layer.self_attn.attn.sliding_window_size = self._draft_window_size # Llama 3.2 1B Instruct set tie_word_embeddings to True # Llama 3.1 8B Instruct set tie_word_embeddings to False self.load_lm_head_from_target = False @@ -348,14 +351,8 @@ class LlamaForCausalLMEagle3(LlamaForCausalLM): def get_hot_token_id(self): return self.hot_token_id - def get_attention_sliding_window_size(self): - server_args = get_global_server_args() - draft_window_size: Optional[int] = ( - int(server_args.speculative_draft_window_size) - if server_args.speculative_draft_window_size is not None - else None - ) - return draft_window_size + def get_attention_sliding_window_size(self) -> Optional[int]: + return self._draft_window_size EntryClass = [LlamaForCausalLMEagle3] diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index f7d7cf67b..29d446eb2 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -27,6 +27,12 @@ import random import tempfile from typing import Any, Callable, Dict, List, Literal, Optional, Union +from sglang.srt.arg_groups.argparse_actions import ( + DeprecatedAction, + DeprecatedAliasStoreAction, + DeprecatedStoreTrueAction, + LoRAPathAction, +) from sglang.srt.configs.linear_attn_model_registry import get_linear_attn_spec_by_arch from sglang.srt.connector import ConnectorType from sglang.srt.environ import envs @@ -316,43 +322,6 @@ def add_linear_attn_kernel_backend_choices(choices): LINEAR_ATTN_KERNEL_BACKEND_CHOICES.extend(choices) -def _resolve_speculative_algorithm_alias( - speculative_algorithm: Optional[str], - speculative_draft_model_path: Optional[str], - trust_remote_code: bool = False, -) -> Optional[str]: - """Resolve CLI speculative algorithm; NEXTN/EAGLE may become FROZEN_KV_MTP for Gemma4 assistant drafts.""" - - is_gemma4_draft = False - if speculative_draft_model_path: - from sglang.srt.utils.hf_transformers_utils import get_config - - cfg = get_config( - speculative_draft_model_path, trust_remote_code=trust_remote_code - ) - is_gemma4_draft = "Gemma4AssistantForCausalLM" in ( - getattr(cfg, "architectures", None) or [] - ) - - if speculative_algorithm == "EAGLE3" and is_gemma4_draft: - raise ValueError( - "Gemma4AssistantForCausalLM draft requires " - "--speculative-algorithm NEXTN or EAGLE; EAGLE3 is " - "not supported for this draft architecture." - ) - - if speculative_algorithm == "NEXTN" or speculative_algorithm == "EAGLE": - if is_gemma4_draft: - logger.info( - "Detected Gemma4AssistantForCausalLM draft; " - f"promoting --speculative-algorithm {speculative_algorithm} to FROZEN_KV_MTP." - ) - return "FROZEN_KV_MTP" - return "EAGLE" - - return speculative_algorithm - - @dataclasses.dataclass class ServerArgs: """ @@ -972,7 +941,9 @@ class ServerArgs: self._handle_pipeline_parallelism() # Handle speculative decoding logic. - self._handle_speculative_decoding() + from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding + + handle_speculative_decoding(self) # Handle model loading format. self._handle_load_format() @@ -3547,376 +3518,6 @@ class ServerArgs: ) return False - def _handle_speculative_decoding(self): - if ( - self.speculative_draft_model_path is not None - and self.speculative_draft_model_revision is None - ): - self.speculative_draft_model_revision = "main" - - if self.speculative_moe_runner_backend is None: - self.speculative_moe_runner_backend = self.moe_runner_backend - - if self.speculative_algorithm is not None: - self.speculative_algorithm = self.speculative_algorithm.upper() - - self.speculative_algorithm = _resolve_speculative_algorithm_alias( - self.speculative_algorithm, - self.speculative_draft_model_path, - trust_remote_code=self.trust_remote_code, - ) - - if self.speculative_algorithm is not None: - from sglang.srt.speculative.spec_info import SpeculativeAlgorithm - from sglang.srt.speculative.spec_registry import CustomSpecAlgo - - algo = SpeculativeAlgorithm.from_string(self.speculative_algorithm) - - # TODO: move the per-algorithm validation below into spec module hooks. - if ( - isinstance(algo, CustomSpecAlgo) - and algo.validate_server_args is not None - ): - algo.validate_server_args(self) - - if self.speculative_skip_dp_mlp_sync: - assert self.speculative_algorithm == "EAGLE", ( - "--speculative-skip-dp-mlp-sync is only supported with " - f"speculative_algorithm == EAGLE, got {self.speculative_algorithm}." - ) - - if self.speculative_algorithm == "DFLASH": - if self.enable_dp_attention: - raise ValueError( - "Currently DFLASH speculative decoding does not support dp attention." - ) - - if self.pp_size != 1: - raise ValueError( - "Currently DFLASH speculative decoding only supports pp_size == 1." - ) - - if self.speculative_draft_model_path is None: - raise ValueError( - "DFLASH speculative decoding requires setting --speculative-draft-model-path." - ) - - # DFLASH does not use EAGLE-style `num_steps`/`topk`, but those fields still - # affect generic scheduler/KV-cache accounting (buffer sizing, KV freeing, - # RoPE reservation). Force them to 1 to avoid surprising memory behavior. - # - # For DFlash, the natural unit is `block_size` (verify window length). - if self.speculative_num_steps is None: - self.speculative_num_steps = 1 - elif int(self.speculative_num_steps) != 1: - logger.warning( - "DFLASH only supports speculative_num_steps == 1; overriding speculative_num_steps=%s to 1.", - self.speculative_num_steps, - ) - self.speculative_num_steps = 1 - - if self.speculative_eagle_topk is None: - self.speculative_eagle_topk = 1 - elif int(self.speculative_eagle_topk) != 1: - logger.warning( - "DFLASH only supports speculative_eagle_topk == 1; overriding speculative_eagle_topk=%s to 1.", - self.speculative_eagle_topk, - ) - self.speculative_eagle_topk = 1 - - if self.speculative_dflash_block_size is not None: - if int(self.speculative_dflash_block_size) <= 0: - raise ValueError( - "DFLASH requires --speculative-dflash-block-size to be positive, " - f"got {self.speculative_dflash_block_size}." - ) - if self.speculative_num_draft_tokens is not None and int( - self.speculative_num_draft_tokens - ) != int(self.speculative_dflash_block_size): - raise ValueError( - "Both --speculative-num-draft-tokens and --speculative-dflash-block-size are set " - "but they differ. For DFLASH they must match. " - f"speculative_num_draft_tokens={self.speculative_num_draft_tokens}, " - f"speculative_dflash_block_size={self.speculative_dflash_block_size}." - ) - self.speculative_num_draft_tokens = int( - self.speculative_dflash_block_size - ) - - window_size = None - if self.speculative_draft_window_size is not None: - window_size = int(self.speculative_draft_window_size) - if window_size <= 0: - raise ValueError( - f"--speculative-draft-window-size must be positive, got {window_size}." - ) - self.speculative_draft_window_size = window_size - - if self.speculative_num_draft_tokens is None: - from sglang.srt.speculative.dflash_utils import ( - parse_dflash_draft_config, - ) - - model_override_args = json.loads(self.json_model_override_args) - inferred_block_size = None - try: - from sglang.srt.utils.hf_transformers_utils import get_config - - draft_hf_config = get_config( - self.speculative_draft_model_path, - trust_remote_code=self.trust_remote_code, - revision=self.speculative_draft_model_revision, - model_override_args=model_override_args, - ) - inferred_block_size = parse_dflash_draft_config( - draft_hf_config=draft_hf_config - ).resolve_block_size(default=None) - except Exception as e: - logger.warning( - "Failed to infer DFLASH block_size from draft model config; " - "defaulting speculative_num_draft_tokens to 16. Error: %s", - e, - ) - - if inferred_block_size is None: - inferred_block_size = 16 - logger.warning( - "speculative_num_draft_tokens is not set; defaulting to %d for DFLASH.", - inferred_block_size, - ) - self.speculative_num_draft_tokens = inferred_block_size - - if window_size is not None: - draft_tokens = int(self.speculative_num_draft_tokens) - if window_size < draft_tokens: - raise ValueError( - "--speculative-draft-window-size must be >= " - "--speculative-num-draft-tokens (block_size). " - f"window_size={window_size}, block_size={draft_tokens}." - ) - - if self.max_running_requests is None: - self.max_running_requests = 48 - logger.warning( - "Max running requests is reset to 48 for speculative decoding. You can override this by explicitly setting --max-running-requests." - ) - - self.disable_overlap_schedule = True - logger.warning( - "Overlap scheduler is disabled when using DFLASH speculative decoding (spec v2 is not supported yet)." - ) - - if self.enable_mixed_chunk: - self.enable_mixed_chunk = False - logger.warning( - "Mixed chunked prefill is disabled because of using dflash speculative decoding." - ) - - if self.speculative_algorithm == "FROZEN_KV_MTP": - if self.max_running_requests is None: - self.max_running_requests = 48 - logger.warning( - "Max running requests is reset to 48 for speculative decoding. You can override this by explicitly setting --max-running-requests." - ) - - self.disable_overlap_schedule = True - logger.warning( - "Overlap scheduler is disabled when using Frozen-KV MTP speculative decoding (spec v2 is not supported yet)." - ) - - if self.enable_mixed_chunk: - self.enable_mixed_chunk = False - logger.warning( - "Mixed chunked prefill is disabled because of using " - "Frozen-KV MTP speculative decoding." - ) - - if self.speculative_algorithm in ("EAGLE", "EAGLE3", "STANDALONE"): - if self.speculative_algorithm == "STANDALONE" and self.enable_dp_attention: - # TODO: support dp attention for standalone speculative decoding - raise ValueError( - "Currently standalone speculative decoding does not support dp attention." - ) - - if self.max_running_requests is None: - self.max_running_requests = 48 - logger.warning( - "Max running requests is reset to 48 for speculative decoding. You can override this by explicitly setting --max-running-requests." - ) - - spec_v1_reason = None - if ( - self.speculative_eagle_topk is not None - and self.speculative_eagle_topk > 1 - and not self.disable_overlap_schedule - ): - self.disable_overlap_schedule = True - spec_v1_reason = "spec v2 currently only supports topk = 1" - elif ( - not envs.SGLANG_ENABLE_SPEC_V2.get() - and not self.disable_overlap_schedule - ): - self.disable_overlap_schedule = True - spec_v1_reason = "SGLANG_ENABLE_SPEC_V2=False" - - if self.disable_overlap_schedule: - logger.warning( - "Spec v1 is used for eagle/eagle3/standalone speculative decoding because %s.", - spec_v1_reason or "overlap schedule is disabled", - ) - else: - logger.warning( - "Spec v2 is enabled by default for eagle/eagle3/standalone speculative decoding." - ) - - if self.enable_mixed_chunk: - self.enable_mixed_chunk = False - logger.warning( - "Mixed chunked prefill is disabled because of using " - "eagle speculative decoding." - ) - - model_arch = self.get_model_config().hf_config.architectures[0] - if model_arch in [ - "DeepseekV32ForCausalLM", - "DeepseekV3ForCausalLM", - "DeepseekV4ForCausalLM", - "Glm4MoeForCausalLM", - "Glm4MoeLiteForCausalLM", - "GlmMoeDsaForCausalLM", - "BailingMoeForCausalLM", - "BailingMoeV2ForCausalLM", - "BailingMoeV2_5ForCausalLM", - "MistralLarge3ForCausalLM", - "PixtralForConditionalGeneration", - "HYV3ForCausalLM", - ]: - if self.speculative_draft_model_path is None: - self.speculative_draft_model_path = self.model_path - self.speculative_draft_model_revision = self.revision - else: - if model_arch not in [ - "MistralLarge3ForCausalLM", - "PixtralForConditionalGeneration", - ]: - logger.warning( - "DeepSeek MTP does not require setting speculative_draft_model_path." - ) - - if self.speculative_num_steps is None: - assert ( - self.speculative_eagle_topk is None - and self.speculative_num_draft_tokens is None - ) - ( - self.speculative_num_steps, - self.speculative_eagle_topk, - self.speculative_num_draft_tokens, - ) = auto_choose_speculative_params(self) - - if ( - self.attention_backend == "trtllm_mha" - or self.decode_attention_backend == "trtllm_mha" - or self.prefill_attention_backend == "trtllm_mha" - ): - if self.speculative_eagle_topk > 1: - raise ValueError( - "trtllm_mha backend only supports topk = 1 for speculative decoding." - ) - - if ( - self.speculative_eagle_topk == 1 - and self.speculative_num_draft_tokens != self.speculative_num_steps + 1 - ): - logger.warning( - "speculative_num_draft_tokens is adjusted to speculative_num_steps + 1 when speculative_eagle_topk == 1" - ) - self.speculative_num_draft_tokens = self.speculative_num_steps + 1 - - if ( - self.speculative_eagle_topk > 1 - and self.page_size > 1 - and self.attention_backend not in ["flashinfer", "fa3"] - ): - raise ValueError( - "speculative_eagle_topk > 1 with page_size > 1 is unstable and produces incorrect results for paged attention backends. This combination is only supported for the 'flashinfer' backend." - ) - - if self.speculative_algorithm == "NGRAM": - if not self.device.startswith("cuda"): - raise ValueError( - "Ngram speculative decoding only supports CUDA device." - ) - - if self.max_running_requests is None: - self.max_running_requests = 48 - logger.warning( - "Max running requests is reset to 48 for speculative decoding. You can override this by explicitly setting --max-running-requests." - ) - - self.disable_overlap_schedule = True - self.enable_mixed_chunk = False - self.speculative_eagle_topk = self.speculative_ngram_max_bfs_breadth - if self.speculative_num_draft_tokens is None: - self.speculative_num_draft_tokens = 12 - logger.warning( - "speculative_num_draft_tokens is set to 12 by default for ngram speculative decoding. " - "You can override this by explicitly setting --speculative-num-draft-tokens." - ) - if self.speculative_ngram_external_corpus_path is not None: - if self.speculative_ngram_external_sam_budget <= 0: - raise ValueError( - "--speculative-ngram-external-sam-budget must be positive when " - "--speculative-ngram-external-corpus-path is set." - ) - if self.speculative_ngram_external_corpus_max_tokens <= 0: - raise ValueError( - "--speculative-ngram-external-corpus-max-tokens must be positive when " - "--speculative-ngram-external-corpus-path is set." - ) - if ( - self.speculative_ngram_external_sam_budget - > self.speculative_num_draft_tokens - 1 - ): - raise ValueError( - "speculative_ngram_external_sam_budget must be less than or equal to " - f"speculative_num_draft_tokens - 1 ({self.speculative_num_draft_tokens - 1})." - ) - logger.warning( - "The overlap scheduler and mixed chunked prefill are disabled because of " - "using ngram speculative decoding." - ) - - if ( - self.speculative_eagle_topk > 1 - and self.page_size > 1 - and self.attention_backend != "flashinfer" - ): - raise ValueError( - f"speculative_eagle_topk({self.speculative_eagle_topk}) > 1 " - f"with page_size({self.page_size}) > 1 is unstable " - "and produces incorrect results for paged attention backends. " - "This combination is only supported for the 'flashinfer' backend." - ) - if self.enable_dp_attention: - # TODO: support dp attention for ngram speculative decoding - raise ValueError( - "Currently ngram speculative decoding does not support dp attention." - ) - - if self.speculative_adaptive: - from sglang.srt.speculative.adaptive_spec_params import ( - adaptive_unsupported_reason, - ) - - reason = adaptive_unsupported_reason(self) - if reason is not None: - logger.warning( - f"speculative_adaptive disabled: {reason}. " - "Falling back to static speculative params." - ) - self.speculative_adaptive = False - def _handle_load_format(self): if ( self.load_format == "auto" or self.load_format == "gguf" @@ -5870,17 +5471,26 @@ class ServerArgs: ) parser.add_argument( "--speculative-draft-window-size", - "--speculative-dflash-draft-window-size", type=int, dest="speculative_draft_window_size", - help="Sliding window size for the draft model (honored by EAGLE-3 and DFLASH). " - "For EAGLE-3, the drafter only attends to the most recent N keys " + help="Sliding window size for the draft model. Honored by Llama EAGLE-3 " + "(`LlamaForCausalLMEagle3`) and DFLASH only; other EAGLE-3 backends (e.g. " + "MLA-based drafters) silently ignore it. " + "For Llama EAGLE-3, the drafter only attends to the most recent N keys " "(verifier hidden states + its own outputs); the verifier is unaffected. " "For DFLASH, the draft worker keeps a recent target-token window in its " "local KV cache (paged backends may retain up to one extra page on the " "left for alignment). Default is full attention/context.", default=ServerArgs.speculative_draft_window_size, ) + parser.add_argument( + "--speculative-dflash-draft-window-size", + type=int, + dest="speculative_draft_window_size", + action=DeprecatedAliasStoreAction, + new_flag="--speculative-draft-window-size", + help=argparse.SUPPRESS, + ) parser.add_argument( "--speculative-moe-runner-backend", type=str, @@ -7855,68 +7465,6 @@ class PortArgs: ) -class LoRAPathAction(argparse.Action): - def __call__(self, parser, namespace, values, option_string=None): - lora_paths = [] - if values: - assert isinstance(values, list), "Expected a list of LoRA paths." - for lora_path in values: - lora_path = lora_path.strip() - if lora_path.startswith("{") and lora_path.endswith("}"): - obj = json.loads(lora_path) - assert "lora_path" in obj and "lora_name" in obj, ( - f"{repr(lora_path)} looks like a JSON str, " - "but it does not contain 'lora_name' and 'lora_path' keys." - ) - lora_paths.append(obj) - else: - lora_paths.append(lora_path) - - setattr(namespace, self.dest, lora_paths) - - -def print_deprecated_warning(message: str): - logger.warning(f"\033[1;33m{message}\033[0m") - - -class DeprecatedAction(argparse.Action): - def __init__(self, option_strings, dest, nargs=0, **kwargs): - super(DeprecatedAction, self).__init__( - option_strings, dest, nargs=nargs, **kwargs - ) - - def __call__(self, parser, namespace, values, option_string=None): - print_deprecated_warning( - f"The command line argument '{option_string}' is deprecated and will be removed in future versions." - ) - - -class DeprecatedStoreTrueAction(argparse.Action): - """Deprecated flag that still stores True and prints a warning.""" - - def __init__( - self, - option_strings, - dest, - new_flag=None, - nargs=0, - const=True, - default=False, - **kwargs, - ): - self.new_flag = new_flag - super().__init__( - option_strings, dest, nargs=nargs, const=const, default=default, **kwargs - ) - - def __call__(self, parser, namespace, values, option_string=None): - replacement = f" Use '{self.new_flag}' instead." if self.new_flag else "" - print_deprecated_warning( - f"'{option_string}' is deprecated and will be removed in a future release.{replacement}" - ) - setattr(namespace, self.dest, True) - - def auto_choose_speculative_params(self: ServerArgs): """ Automatically choose the parameters for speculative decoding. diff --git a/python/sglang/srt/speculative/dflash_worker.py b/python/sglang/srt/speculative/dflash_worker.py index 4b9e4fe9b..b3d688c84 100644 --- a/python/sglang/srt/speculative/dflash_worker.py +++ b/python/sglang/srt/speculative/dflash_worker.py @@ -73,10 +73,9 @@ class DFlashWorker: self.target_worker = target_worker self.model_runner = target_worker.model_runner self.page_size = server_args.page_size + # Normalized in arg_groups.speculative_hook.handle_speculative_decoding. self.draft_window_size: Optional[int] = ( - int(server_args.speculative_draft_window_size) - if server_args.speculative_draft_window_size is not None - else None + server_args.speculative_draft_window_size ) self.use_compact_draft_cache = self.draft_window_size is not None self.device = target_worker.device diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_worker.py b/python/sglang/srt/speculative/frozen_kv_mtp_worker.py index 02d1454b3..480dd7144 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_worker.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_worker.py @@ -110,7 +110,7 @@ class FrozenKVMTPWorker(TpModelWorker): "FrozenKVMTPWorker should only be instantiated for " "SpeculativeAlgorithm.FROZEN_KV_MTP, got " f"{self.speculative_algorithm.name}. The dispatch happens in " - "server_args._handle_speculative_decoding -> " + "arg_groups.speculative_hook.handle_speculative_decoding -> " "_resolve_speculative_algorithm_alias." ) diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 69fc94dd8..fdcea5039 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -3,6 +3,7 @@ import tempfile import unittest from unittest.mock import MagicMock, patch +from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding from sglang.srt.server_args import PortArgs, ServerArgs, prepare_server_args from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import ( @@ -497,21 +498,23 @@ class TestNgramExternalSamArgs(CustomTestCase): return args def test_external_sam_budget_must_fit_draft_budget(self): + args = self._make_dummy_ngram_args( + speculative_num_draft_tokens=4, + speculative_ngram_external_corpus_path="/tmp/ngram-corpus.jsonl", + speculative_ngram_external_sam_budget=4, + ) with self.assertRaises(ValueError) as context: - self._make_dummy_ngram_args( - speculative_num_draft_tokens=4, - speculative_ngram_external_corpus_path="/tmp/ngram-corpus.jsonl", - speculative_ngram_external_sam_budget=4, - )._handle_speculative_decoding() + handle_speculative_decoding(args) self.assertIn("speculative_num_draft_tokens - 1", str(context.exception)) def test_external_corpus_max_tokens_must_be_positive(self): + args = self._make_dummy_ngram_args( + speculative_ngram_external_corpus_path="/tmp/ngram-corpus.jsonl", + speculative_ngram_external_sam_budget=2, + speculative_ngram_external_corpus_max_tokens=0, + ) with self.assertRaises(ValueError) as context: - self._make_dummy_ngram_args( - speculative_ngram_external_corpus_path="/tmp/ngram-corpus.jsonl", - speculative_ngram_external_sam_budget=2, - speculative_ngram_external_corpus_max_tokens=0, - )._handle_speculative_decoding() + handle_speculative_decoding(args) self.assertIn("external-corpus-max-tokens", str(context.exception))