[Feature] Add DFLASH speculative decoding support (#22077)
Co-authored-by: Jian Chen <141193260+jianc99@users.noreply.github.com> Co-authored-by: Zhijian Liu <5782437+zhijian-liu@users.noreply.github.com> Co-authored-by: Richard Gong <8001209+gongy@users.noreply.github.com> Co-authored-by: David Wang <21328423+dcw02@users.noreply.github.com> Co-authored-by: yilian49 <43861414+yilian49@users.noreply.github.com> Co-authored-by: xm:D <38322020+xiaomin-d@users.noreply.github.com>
This commit is contained in:
co-authored by
Jian Chen
Zhijian Liu
Richard Gong
yilian49
xm:D
parent
e14876742a
commit
f08726fd56
@@ -499,6 +499,8 @@ class ServerArgs:
|
||||
speculative_num_steps: Optional[int] = None
|
||||
speculative_eagle_topk: Optional[int] = None
|
||||
speculative_num_draft_tokens: Optional[int] = None
|
||||
speculative_dflash_block_size: Optional[int] = None
|
||||
speculative_dflash_draft_window_size: Optional[int] = None
|
||||
speculative_accept_threshold_single: float = 1.0
|
||||
speculative_accept_threshold_acc: float = 1.0
|
||||
speculative_token_map: Optional[str] = None
|
||||
@@ -3027,6 +3029,134 @@ class ServerArgs:
|
||||
if self.speculative_algorithm == "NEXTN":
|
||||
self.speculative_algorithm = "EAGLE"
|
||||
|
||||
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_dflash_draft_window_size is not None:
|
||||
window_size = int(self.speculative_dflash_draft_window_size)
|
||||
if window_size <= 0:
|
||||
raise ValueError(
|
||||
"DFLASH requires --speculative-dflash-draft-window-size "
|
||||
f"to be positive, got {window_size}."
|
||||
)
|
||||
self.speculative_dflash_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(
|
||||
"DFLASH --speculative-dflash-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 in ("EAGLE", "EAGLE3", "STANDALONE"):
|
||||
if self.speculative_algorithm == "STANDALONE" and self.enable_dp_attention:
|
||||
# TODO: support dp attention for standalone speculative decoding
|
||||
@@ -4832,7 +4962,7 @@ class ServerArgs:
|
||||
parser.add_argument(
|
||||
"--speculative-algorithm",
|
||||
type=str,
|
||||
choices=["EAGLE", "EAGLE3", "NEXTN", "STANDALONE", "NGRAM"],
|
||||
choices=["DFLASH", "EAGLE", "EAGLE3", "NEXTN", "STANDALONE", "NGRAM"],
|
||||
help="Speculative algorithm.",
|
||||
)
|
||||
parser.add_argument(
|
||||
@@ -4876,6 +5006,21 @@ class ServerArgs:
|
||||
help="The number of tokens sampled from the draft model in Speculative Decoding.",
|
||||
default=ServerArgs.speculative_num_draft_tokens,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--speculative-dflash-block-size",
|
||||
type=int,
|
||||
help="DFLASH only. Block size (verify window length). Alias of --speculative-num-draft-tokens for DFLASH.",
|
||||
default=ServerArgs.speculative_dflash_block_size,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--speculative-dflash-draft-window-size",
|
||||
type=int,
|
||||
help="DFLASH only. Sliding window size for the draft-model KV cache. "
|
||||
"When set, the draft worker keeps a recent target-token window in its "
|
||||
"local cache (paged backends may retain up to one extra page on the left "
|
||||
"for alignment). Default is full context.",
|
||||
default=ServerArgs.speculative_dflash_draft_window_size,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--speculative-accept-threshold-single",
|
||||
type=float,
|
||||
|
||||
Reference in New Issue
Block a user