[Spec] Clean up draft-window-size handling; extract spec arg setup to arg_groups (#25424)

This commit is contained in:
Liangsheng Yin
2026-05-16 03:23:48 -07:00
committed by GitHub
parent d1eb472a7a
commit aec4022e58
7 changed files with 574 additions and 502 deletions
@@ -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)
@@ -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
+12 -15
View File
@@ -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]
+21 -473
View File
@@ -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.
@@ -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
@@ -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."
)