[Spec] NGRAMWorker on BaseSpecWorker; algo-owned verify-tree shape params (#27799)
This commit is contained in:
@@ -169,11 +169,10 @@ def get_alloc_len_per_decode(server_args: Optional[ServerArgs] = None) -> int:
|
|||||||
spec_tokens = server_args.max_speculative_num_draft_tokens
|
spec_tokens = server_args.max_speculative_num_draft_tokens
|
||||||
page_size = server_args.page_size
|
page_size = server_args.page_size
|
||||||
|
|
||||||
# NGRAM drafts are a flat candidate list written to contiguous slots after
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
# seq_lens (no per-topk page duplication), so the flat formula applies at
|
|
||||||
# any page_size.
|
spec_algo = SpeculativeAlgorithm.from_string(server_args.speculative_algorithm)
|
||||||
is_ngram = (server_args.speculative_algorithm or "").upper() == "NGRAM"
|
if page_size == 1 or spec_topk == 1 or not spec_algo.has_draft_kv():
|
||||||
if page_size == 1 or spec_topk == 1 or is_ngram:
|
|
||||||
return max(spec_steps * spec_topk, spec_tokens)
|
return max(spec_steps * spec_topk, spec_tokens)
|
||||||
else:
|
else:
|
||||||
# page_size > 1 + topk > 1 (spec v2 tree): worst-case page-aligned tree
|
# page_size > 1 + topk > 1 (spec v2 tree): worst-case page-aligned tree
|
||||||
|
|||||||
@@ -368,6 +368,19 @@ class EagleDraftInputV2Mixin:
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class EagleVerifyInputV2Mixin:
|
class EagleVerifyInputV2Mixin:
|
||||||
|
@property
|
||||||
|
def max_tree_depth(self: EagleVerifyInput) -> int:
|
||||||
|
"""Longest root-to-leaf chain of the verify tree, incl. the root;
|
||||||
|
bounds the accept_index row width. EAGLE trees are depth-bounded by
|
||||||
|
the draft loop. Algorithms with other tree shapes override this."""
|
||||||
|
return self.spec_steps + 1
|
||||||
|
|
||||||
|
@property
|
||||||
|
def tree_topk(self: EagleVerifyInput) -> int:
|
||||||
|
"""Branching factor passed to the tree-verify kernels; -1 means an
|
||||||
|
irregular tree (no fixed per-level branching)."""
|
||||||
|
return self.topk
|
||||||
|
|
||||||
def prepare_for_v2_verify(
|
def prepare_for_v2_verify(
|
||||||
self: EagleVerifyInput,
|
self: EagleVerifyInput,
|
||||||
req_to_token_pool: ReqToTokenPool,
|
req_to_token_pool: ReqToTokenPool,
|
||||||
@@ -493,17 +506,8 @@ class EagleVerifyInputV2Mixin:
|
|||||||
candidates = self.draft_token.reshape(bs, self.draft_token_num)
|
candidates = self.draft_token.reshape(bs, self.draft_token_num)
|
||||||
predict_shape = list(next_token_logits.shape)[:-1]
|
predict_shape = list(next_token_logits.shape)[:-1]
|
||||||
predict = torch.zeros(predict_shape, dtype=torch.int32, device=device).flatten()
|
predict = torch.zeros(predict_shape, dtype=torch.int32, device=device).flatten()
|
||||||
# Longest root-to-leaf chain of the verify tree, incl. the root; bounds
|
|
||||||
# the accept_index row width. EAGLE trees are depth-bounded by the draft
|
|
||||||
# loop (spec_steps + 1); NGRAM trees are node-budgeted with no depth cap
|
|
||||||
# (a single corpus match can chain all draft_token_num nodes).
|
|
||||||
max_tree_depth = (
|
|
||||||
self.draft_token_num
|
|
||||||
if batch.spec_algorithm.is_ngram()
|
|
||||||
else self.spec_steps + 1
|
|
||||||
)
|
|
||||||
accept_index = torch.full(
|
accept_index = torch.full(
|
||||||
(bs, max_tree_depth), -1, dtype=torch.int32, device=device
|
(bs, self.max_tree_depth), -1, dtype=torch.int32, device=device
|
||||||
)
|
)
|
||||||
num_correct_drafts = torch.empty((bs,), dtype=torch.int32, device=device)
|
num_correct_drafts = torch.empty((bs,), dtype=torch.int32, device=device)
|
||||||
|
|
||||||
@@ -520,7 +524,7 @@ class EagleVerifyInputV2Mixin:
|
|||||||
retrieve_next_token=self.retrieve_next_token,
|
retrieve_next_token=self.retrieve_next_token,
|
||||||
retrieve_next_sibling=self.retrieve_next_sibling,
|
retrieve_next_sibling=self.retrieve_next_sibling,
|
||||||
target_predict=target_predict,
|
target_predict=target_predict,
|
||||||
topk=-1 if batch.spec_algorithm.is_ngram() else self.topk,
|
topk=self.tree_topk,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# Apply temperature and get target probs
|
# Apply temperature and get target probs
|
||||||
@@ -598,7 +602,7 @@ class EagleVerifyInputV2Mixin:
|
|||||||
num_correct_drafts=num_correct_drafts, # mutable
|
num_correct_drafts=num_correct_drafts, # mutable
|
||||||
simulate_acc_len=SIMULATE_ACC_LEN,
|
simulate_acc_len=SIMULATE_ACC_LEN,
|
||||||
bs=bs,
|
bs=bs,
|
||||||
spec_steps=max_tree_depth - 1,
|
spec_steps=self.max_tree_depth - 1,
|
||||||
)
|
)
|
||||||
|
|
||||||
# `num_correct_drafts` stays drafts-only inside this function; the returned
|
# `num_correct_drafts` stays drafts-only inside this function; the returned
|
||||||
|
|||||||
@@ -49,6 +49,18 @@ class NgramVerifyInput(SpecInput, EagleDraftInputV2Mixin, EagleVerifyInputV2Mixi
|
|||||||
custom_mask.device if custom_mask is not None else new_seq_lens.device
|
custom_mask.device if custom_mask is not None else new_seq_lens.device
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def max_tree_depth(self) -> int:
|
||||||
|
# NGRAM trees are node-budgeted with no depth cap: the corpus BFS only
|
||||||
|
# stops on the node budget, so a single long match can chain all
|
||||||
|
# draft_token_num nodes (spec_steps is meaningless for this tree).
|
||||||
|
return self.draft_token_num
|
||||||
|
|
||||||
|
@property
|
||||||
|
def tree_topk(self) -> int:
|
||||||
|
# Irregular tree: per-level branching follows the corpus matches.
|
||||||
|
return -1
|
||||||
|
|
||||||
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
|
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
|
||||||
return self.draft_token_num, self.draft_token_num
|
return self.draft_token_num, self.draft_token_num
|
||||||
|
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ from sglang.srt.managers.tp_worker import TpModelWorker
|
|||||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||||
from sglang.srt.observability.req_time_stats import set_time_batch
|
from sglang.srt.observability.req_time_stats import set_time_batch
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
from sglang.srt.speculative.base_spec_worker import BaseDraftWorker, BaseSpecWorker
|
||||||
from sglang.srt.speculative.cpp_ngram.ngram_corpus import NgramCorpus
|
from sglang.srt.speculative.cpp_ngram.ngram_corpus import NgramCorpus
|
||||||
from sglang.srt.speculative.ngram_info import NgramVerifyInput
|
from sglang.srt.speculative.ngram_info import NgramVerifyInput
|
||||||
from sglang.srt.speculative.spec_utils import (
|
from sglang.srt.speculative.spec_utils import (
|
||||||
@@ -30,7 +31,7 @@ logger = logging.getLogger(__name__)
|
|||||||
USE_FULL_MASK = True
|
USE_FULL_MASK = True
|
||||||
|
|
||||||
|
|
||||||
class NGRAMWorker:
|
class NGRAMWorker(BaseSpecWorker):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
@@ -45,7 +46,7 @@ class NGRAMWorker:
|
|||||||
):
|
):
|
||||||
self.server_args = server_args
|
self.server_args = server_args
|
||||||
self.enable_overlap = not server_args.disable_overlap_schedule
|
self.enable_overlap = not server_args.disable_overlap_schedule
|
||||||
self.target_worker = target_worker
|
self._target_worker = target_worker
|
||||||
self.model_runner = target_worker.model_runner
|
self.model_runner = target_worker.model_runner
|
||||||
self.tp_rank = tp_rank
|
self.tp_rank = tp_rank
|
||||||
self.page_size = server_args.page_size
|
self.page_size = server_args.page_size
|
||||||
@@ -99,6 +100,15 @@ class NGRAMWorker:
|
|||||||
loaded,
|
loaded,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def target_worker(self) -> TpModelWorker:
|
||||||
|
return self._target_worker
|
||||||
|
|
||||||
|
@property
|
||||||
|
def draft_worker(self) -> Optional[BaseDraftWorker]:
|
||||||
|
# NGRAM has no draft model; drafts come from the CPU-side corpus.
|
||||||
|
return None
|
||||||
|
|
||||||
def clear_cache_pool(self):
|
def clear_cache_pool(self):
|
||||||
self.ngram_corpus.reset()
|
self.ngram_corpus.reset()
|
||||||
self._prev_decode_rids = set()
|
self._prev_decode_rids = set()
|
||||||
|
|||||||
@@ -118,6 +118,12 @@ class SpeculativeAlgorithm(Enum):
|
|||||||
def supports_target_verify_for_draft(self) -> bool:
|
def supports_target_verify_for_draft(self) -> bool:
|
||||||
return self.is_dflash()
|
return self.is_dflash()
|
||||||
|
|
||||||
|
def has_draft_kv(self) -> bool:
|
||||||
|
"""Whether the draft phase writes KV chains. NGRAM does not (its tree
|
||||||
|
lives only in the verify mask), so per-decode KV sizing needs no
|
||||||
|
per-topk page rounding; see get_alloc_len_per_decode."""
|
||||||
|
return not self.is_ngram()
|
||||||
|
|
||||||
def create_future_map(
|
def create_future_map(
|
||||||
self,
|
self,
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
|
|||||||
@@ -78,6 +78,10 @@ class CustomSpecAlgo:
|
|||||||
def supports_target_verify_for_draft(self) -> bool:
|
def supports_target_verify_for_draft(self) -> bool:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
def has_draft_kv(self) -> bool:
|
||||||
|
# Conservative default: the larger KV reserve.
|
||||||
|
return True
|
||||||
|
|
||||||
def supports_spec_v2(self) -> bool:
|
def supports_spec_v2(self) -> bool:
|
||||||
return self.supports_overlap
|
return self.supports_overlap
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user