[Spec] Clean up draft-window-size handling; extract spec arg setup to arg_groups (#25424)
This commit is contained in:
@@ -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
|
||||
@@ -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]
|
||||
|
||||
@@ -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."
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user