Overlap grammar (constrained decoding) with speculative decode verify (#31488)
Co-authored-by: Jason Park <jasonjk@fb.com>
This commit is contained in:
co-authored by
Jason Park
parent
8905cbd42f
commit
e7e8aaa73c
@@ -26,7 +26,8 @@ from xgrammar import (
|
|||||||
StructuralTag,
|
StructuralTag,
|
||||||
StructuralTagItem,
|
StructuralTagItem,
|
||||||
TokenizerInfo,
|
TokenizerInfo,
|
||||||
allocate_token_bitmask,
|
bitmask_dtype,
|
||||||
|
get_bitmask_shape,
|
||||||
)
|
)
|
||||||
|
|
||||||
from sglang.srt.constrained.base_grammar_backend import (
|
from sglang.srt.constrained.base_grammar_backend import (
|
||||||
@@ -56,6 +57,18 @@ logger = logging.getLogger(__name__)
|
|||||||
MAX_ROLLBACK_TOKENS = 200
|
MAX_ROLLBACK_TOKENS = 200
|
||||||
|
|
||||||
|
|
||||||
|
def _allocate_token_bitmask(vocab_size: int, batch_size: int) -> torch.Tensor:
|
||||||
|
# Always allocate a pinned bitmask so the later H2D to the device can be a
|
||||||
|
# genuine non_blocking copy (a pageable source silently downgrades it to a
|
||||||
|
# blocking copy).
|
||||||
|
return torch.full(
|
||||||
|
get_bitmask_shape(batch_size, vocab_size),
|
||||||
|
-1,
|
||||||
|
dtype=bitmask_dtype,
|
||||||
|
pin_memory=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class XGrammarGrammar(BaseGrammarObject):
|
class XGrammarGrammar(BaseGrammarObject):
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -100,7 +113,7 @@ class XGrammarGrammar(BaseGrammarObject):
|
|||||||
def allocate_vocab_mask(
|
def allocate_vocab_mask(
|
||||||
self, vocab_size: int, batch_size: int, device
|
self, vocab_size: int, batch_size: int, device
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
return allocate_token_bitmask(batch_size, vocab_size)
|
return _allocate_token_bitmask(vocab_size, batch_size)
|
||||||
|
|
||||||
def fill_vocab_mask(self, vocab_mask: torch.Tensor, idx: int) -> None:
|
def fill_vocab_mask(self, vocab_mask: torch.Tensor, idx: int) -> None:
|
||||||
self.matcher.fill_next_token_bitmask(vocab_mask, idx)
|
self.matcher.fill_next_token_bitmask(vocab_mask, idx)
|
||||||
@@ -228,7 +241,7 @@ class XGrammarGrammarBackend(BaseGrammarBackend):
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def allocate_vocab_mask(vocab_size: int, batch_size: int, device) -> torch.Tensor:
|
def allocate_vocab_mask(vocab_size: int, batch_size: int, device) -> torch.Tensor:
|
||||||
return allocate_token_bitmask(batch_size, vocab_size)
|
return _allocate_token_bitmask(vocab_size, batch_size)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def move_vocab_mask(vocab_mask: torch.Tensor, device) -> torch.Tensor:
|
def move_vocab_mask(vocab_mask: torch.Tensor, device) -> torch.Tensor:
|
||||||
|
|||||||
@@ -3008,6 +3008,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
forward_mode=self.forward_mode,
|
forward_mode=self.forward_mode,
|
||||||
out_cache_loc=self.out_cache_loc,
|
out_cache_loc=self.out_cache_loc,
|
||||||
return_logprob=self.return_logprob,
|
return_logprob=self.return_logprob,
|
||||||
|
has_grammar=self.has_grammar,
|
||||||
decoding_reqs=self.decoding_reqs,
|
decoding_reqs=self.decoding_reqs,
|
||||||
spec_algorithm=self.spec_algorithm,
|
spec_algorithm=self.spec_algorithm,
|
||||||
spec_info=self.spec_info,
|
spec_info=self.spec_info,
|
||||||
|
|||||||
@@ -1646,19 +1646,34 @@ class Scheduler(
|
|||||||
and last_batch_is_extend
|
and last_batch_is_extend
|
||||||
)
|
)
|
||||||
|
|
||||||
# We do not support overlap + spec + grammar yet,
|
# Spec algorithms that don't advance the grammar FSM inside verify() (see
|
||||||
# so we need to turn off overlap for this batch.
|
# supports_grammar_overlap) still need overlap forced off for grammar decode
|
||||||
# TODO(lsyin): support overlap + spec + grammar
|
# batches, so the FSM is advanced before the next batch's bitmask.
|
||||||
need_grammar_sync = (
|
need_grammar_sync = (
|
||||||
batch
|
batch
|
||||||
and not batch.spec_algorithm.is_none()
|
and not batch.spec_algorithm.is_none()
|
||||||
|
and not batch.spec_algorithm.supports_grammar_overlap()
|
||||||
and batch.has_grammar
|
and batch.has_grammar
|
||||||
and batch.forward_mode.is_decode()
|
and batch.forward_mode.is_decode()
|
||||||
and len(self.result_queue) > 0
|
and len(self.result_queue) > 0
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Algorithms that support grammar overlap advance the FSM inside verify()
|
||||||
|
# via the grammar barrier (overlapping the target forward), which resolves
|
||||||
|
# whatever result is still pending in the queue — including the
|
||||||
|
# extend->decode boundary — so no grammar-specific overlap disable is needed.
|
||||||
return disable_overlap_for_batch or need_grammar_sync
|
return disable_overlap_for_batch or need_grammar_sync
|
||||||
|
|
||||||
|
def _advance_pending_grammar(self):
|
||||||
|
"""Grammar barrier (spec-v2 overlap): advance the FSM over any not-yet
|
||||||
|
-processed decode result still in the queue, so a following verify()'s
|
||||||
|
bitmask sees the previous batch's committed tokens. Invoked mid-worker
|
||||||
|
(before generate_token_bitmask) so the CPU advance overlaps the target
|
||||||
|
verify forward. Idempotent; no-op when the queue is empty or has no grammar.
|
||||||
|
"""
|
||||||
|
for prev_batch, prev_result in self.result_queue:
|
||||||
|
self.batch_result_processor.advance_grammar_fsm(prev_result, prev_batch)
|
||||||
|
|
||||||
@scheduler_nvtx_method("scheduler.process_input_requests")
|
@scheduler_nvtx_method("scheduler.process_input_requests")
|
||||||
def process_input_requests(self, recv_reqs: List):
|
def process_input_requests(self, recv_reqs: List):
|
||||||
now = time.monotonic()
|
now = time.monotonic()
|
||||||
@@ -3337,14 +3352,18 @@ class Scheduler(
|
|||||||
# Spec_v2 fires on_publish mid-worker (between verify and
|
# Spec_v2 fires on_publish mid-worker (between verify and
|
||||||
# draft_extend) so schedule prep can overlap with draft_extend.
|
# draft_extend) so schedule prep can overlap with draft_extend.
|
||||||
# Non-spec has no later work — scheduler publishes after return.
|
# Non-spec has no later work — scheduler publishes after return.
|
||||||
fwd_kwargs = (
|
fwd_kwargs = {}
|
||||||
{
|
if not batch.spec_algorithm.is_none():
|
||||||
"on_publish": partial(
|
fwd_kwargs["on_publish"] = partial(
|
||||||
self.future_map.publish, future_indices
|
self.future_map.publish, future_indices
|
||||||
)
|
)
|
||||||
}
|
# Grammar-overlap-capable workers advance the grammar FSM
|
||||||
if not batch.spec_algorithm.is_none()
|
# inside verify() before building the bitmask; hand them the
|
||||||
else {}
|
# barrier that resolves the previous batch's committed
|
||||||
|
# tokens (overlapping the target forward).
|
||||||
|
if batch.spec_algorithm.supports_grammar_overlap():
|
||||||
|
fwd_kwargs["grammar_barrier"] = (
|
||||||
|
self._advance_pending_grammar
|
||||||
)
|
)
|
||||||
|
|
||||||
# FIXME: pp is not compatible with overlap
|
# FIXME: pp is not compatible with overlap
|
||||||
|
|||||||
@@ -273,7 +273,9 @@ class SchedulerBatchResultProcessor:
|
|||||||
|
|
||||||
if req.grammar is not None:
|
if req.grammar is not None:
|
||||||
self._apply_prefill_grammar(
|
self._apply_prefill_grammar(
|
||||||
req=req, next_token_id=next_token_id
|
req=req,
|
||||||
|
next_token_id=next_token_id,
|
||||||
|
already_advanced=result.grammar_advanced,
|
||||||
)
|
)
|
||||||
|
|
||||||
else:
|
else:
|
||||||
@@ -487,13 +489,20 @@ class SchedulerBatchResultProcessor:
|
|||||||
)
|
)
|
||||||
return hidden_state_offset
|
return hidden_state_offset
|
||||||
|
|
||||||
def _apply_prefill_grammar(self, *, req: Req, next_token_id: int) -> None:
|
def _apply_prefill_grammar(
|
||||||
|
self, *, req: Req, next_token_id: int, already_advanced: bool = False
|
||||||
|
) -> None:
|
||||||
|
# The grammar barrier may have already advanced the FSM over this prefilled
|
||||||
|
# token (spec overlap path); only advance if not, but always sync
|
||||||
|
# grammar.finished.
|
||||||
|
if not already_advanced:
|
||||||
# FIXME: this try-except block is for handling unexpected xgrammar issue.
|
# FIXME: this try-except block is for handling unexpected xgrammar issue.
|
||||||
try:
|
try:
|
||||||
req.grammar.accept_token(next_token_id)
|
req.grammar.accept_token(next_token_id)
|
||||||
except ValueError as e:
|
except ValueError as e:
|
||||||
# Grammar accept_token can raise ValueError if the token is not in the grammar.
|
# Grammar accept_token can raise ValueError if the token is not in the
|
||||||
# This can happen if the grammar is not set correctly or the token is invalid.
|
# grammar. This can happen if the grammar is not set correctly or the
|
||||||
|
# token is invalid.
|
||||||
logger.error(
|
logger.error(
|
||||||
f"Grammar accept_token failed for req {req.rid} with token {next_token_id}: {e}"
|
f"Grammar accept_token failed for req {req.rid} with token {next_token_id}: {e}"
|
||||||
)
|
)
|
||||||
@@ -565,6 +574,13 @@ class SchedulerBatchResultProcessor:
|
|||||||
result.num_correct_drafts_per_req_cpu, batch_size=len(batch.reqs)
|
result.num_correct_drafts_per_req_cpu, batch_size=len(batch.reqs)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Advance the grammar FSM over this batch's committed tokens (idempotent):
|
||||||
|
# the EAGLE overlap path already did this inside verify() via the grammar
|
||||||
|
# barrier; otherwise advance now. advance_grammar_fsm self-gates on per-req
|
||||||
|
# grammar (the queued batch.copy() does not carry has_grammar) and consumes
|
||||||
|
# result.grammar_retained_tokens below instead of re-advancing.
|
||||||
|
self.advance_grammar_fsm(result, batch)
|
||||||
|
|
||||||
predict_tokens = []
|
predict_tokens = []
|
||||||
# In adaptive spec-v2, the worker state may already have switched when this
|
# In adaptive spec-v2, the worker state may already have switched when this
|
||||||
# delayed result is processed. Use the draft token count recorded on result.
|
# delayed result is processed. Use the draft token count recorded on result.
|
||||||
@@ -580,11 +596,9 @@ class SchedulerBatchResultProcessor:
|
|||||||
pass
|
pass
|
||||||
else:
|
else:
|
||||||
if req.grammar is not None:
|
if req.grammar is not None:
|
||||||
# Stop accepting once the grammar terminates, so the
|
# FSM already advanced + truncated by advance_grammar_fsm; reuse
|
||||||
# over-drafted suffix is never committed to KV nor emitted.
|
# the retained (grammar-legal) run instead of advancing again.
|
||||||
# This advances the grammar FSM; the result loop only syncs
|
accept_tokens = result.grammar_retained_tokens[i]
|
||||||
# grammar.finished.
|
|
||||||
accept_tokens = self._accept_grammar_tokens(req, accept_tokens)
|
|
||||||
|
|
||||||
# Commit the full accepted run (drafts + bonus).
|
# Commit the full accepted run (drafts + bonus).
|
||||||
num_accept_tokens = len(accept_tokens)
|
num_accept_tokens = len(accept_tokens)
|
||||||
@@ -636,6 +650,64 @@ class SchedulerBatchResultProcessor:
|
|||||||
req.to_finish = FINISH_ABORT()
|
req.to_finish = FINISH_ABORT()
|
||||||
return retained
|
return retained
|
||||||
|
|
||||||
|
def advance_grammar_fsm(
|
||||||
|
self, result: GenerationBatchResult, batch: ScheduleBatch
|
||||||
|
) -> None:
|
||||||
|
"""Advance each req's grammar FSM over the tokens THIS batch committed, and
|
||||||
|
(for decode) memoize the grammar-truncated run on ``result``.
|
||||||
|
|
||||||
|
This is the single place the spec-v2 FSM advances. It is idempotent
|
||||||
|
(``result.grammar_advanced``) so it runs either eagerly — inside ``verify(N)``
|
||||||
|
via the scheduler's grammar barrier, so the advance overlaps the target-verify
|
||||||
|
forward — or lazily from the result processors on the non-overlap / non-EAGLE
|
||||||
|
paths. It handles both decode (the accepted spec run) and extend (the single
|
||||||
|
prefilled token) results, so the barrier can resolve whatever the previous
|
||||||
|
batch was — e.g. the extend->decode boundary.
|
||||||
|
"""
|
||||||
|
if result.grammar_advanced or not batch.has_grammar:
|
||||||
|
return
|
||||||
|
is_decode = batch.forward_mode.is_decode()
|
||||||
|
if not (is_decode or batch.forward_mode.is_extend()):
|
||||||
|
return
|
||||||
|
if result.copy_done is not None:
|
||||||
|
result.copy_done.synchronize()
|
||||||
|
next_token_ids = result.next_token_ids.tolist()
|
||||||
|
|
||||||
|
if not is_decode:
|
||||||
|
# Extend: advance over the single token each completed-prefill req emitted
|
||||||
|
# (mirrors process_batch_result_prefill's per-req token indexing).
|
||||||
|
for i, req in enumerate(batch.reqs):
|
||||||
|
if (
|
||||||
|
req.grammar is None
|
||||||
|
or req.is_retracted
|
||||||
|
or req.finished()
|
||||||
|
or req.inflight_middle_chunks > 0
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
self._accept_grammar_tokens(req, next_token_ids[i])
|
||||||
|
result.grammar_advanced = True
|
||||||
|
return
|
||||||
|
|
||||||
|
# Decode: only the spec-v2 path reaches here (the grammar barrier for
|
||||||
|
# spec-overlap workers and _resolve_spec_v2_tokens). Non-spec grammar decode
|
||||||
|
# advances its FSM in process_batch_result_decode and has no accept_lens, so
|
||||||
|
# bail out defensively.
|
||||||
|
if result.accept_lens is None:
|
||||||
|
return
|
||||||
|
accept_lens = result.accept_lens.tolist()
|
||||||
|
stride = result.speculative_num_draft_tokens
|
||||||
|
assert stride is not None, "spec-v2 result missing speculative_num_draft_tokens"
|
||||||
|
retained = [None] * len(batch.reqs)
|
||||||
|
for i, req in enumerate(batch.reqs):
|
||||||
|
if req.grammar is None or req.is_retracted or req.finished():
|
||||||
|
continue
|
||||||
|
accept_tokens = next_token_ids[i * stride : i * stride + accept_lens[i]]
|
||||||
|
# Stop accepting once the grammar terminates so the over-drafted suffix
|
||||||
|
# is never committed to KV nor emitted; this advances the FSM.
|
||||||
|
retained[i] = self._accept_grammar_tokens(req, accept_tokens)
|
||||||
|
result.grammar_retained_tokens = retained
|
||||||
|
result.grammar_advanced = True
|
||||||
|
|
||||||
def process_batch_result_idle(
|
def process_batch_result_idle(
|
||||||
self,
|
self,
|
||||||
batch: ScheduleBatch,
|
batch: ScheduleBatch,
|
||||||
|
|||||||
@@ -66,6 +66,13 @@ class GenerationBatchResult:
|
|||||||
future_indices: Optional[torch.Tensor] = None
|
future_indices: Optional[torch.Tensor] = None
|
||||||
speculative_num_draft_tokens: Optional[int] = None
|
speculative_num_draft_tokens: Optional[int] = None
|
||||||
|
|
||||||
|
# Grammar FSM advance memoization (spec-v2 overlap). advance_grammar_fsm sets
|
||||||
|
# these once — eagerly via the scheduler's grammar barrier inside verify(), or
|
||||||
|
# lazily in _resolve_spec_v2_tokens — and the latter consumes
|
||||||
|
# grammar_retained_tokens instead of re-advancing the FSM.
|
||||||
|
grammar_advanced: bool = False
|
||||||
|
grammar_retained_tokens: Optional[list] = None
|
||||||
|
|
||||||
# FIXME(lsyin): maybe move to a better place?
|
# FIXME(lsyin): maybe move to a better place?
|
||||||
# sync path: forward stream -> output processor
|
# sync path: forward stream -> output processor
|
||||||
accept_lens: Optional[torch.Tensor] = None
|
accept_lens: Optional[torch.Tensor] = None
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ from sglang.kernels.ops.speculative.cache_locs import (
|
|||||||
)
|
)
|
||||||
from sglang.kernels.ops.speculative.eagle import fill_bonus_tokens_func
|
from sglang.kernels.ops.speculative.eagle import fill_bonus_tokens_func
|
||||||
from sglang.srt.layers.logprob_processor import compute_spec_v2_logprobs
|
from sglang.srt.layers.logprob_processor import compute_spec_v2_logprobs
|
||||||
from sglang.srt.managers.utils import GenerationBatchResult
|
from sglang.srt.managers.utils import GenerationBatchResult, _async_d2h
|
||||||
from sglang.srt.model_executor.forward_batch_info import (
|
from sglang.srt.model_executor.forward_batch_info import (
|
||||||
CaptureHiddenMode,
|
CaptureHiddenMode,
|
||||||
ForwardBatch,
|
ForwardBatch,
|
||||||
@@ -467,6 +467,7 @@ def run_eagle_verify(
|
|||||||
device: str,
|
device: str,
|
||||||
metadata_ready_pre_pad: bool,
|
metadata_ready_pre_pad: bool,
|
||||||
finalize_tree_path: bool,
|
finalize_tree_path: bool,
|
||||||
|
grammar_barrier=None,
|
||||||
) -> GenerationBatchResult:
|
) -> GenerationBatchResult:
|
||||||
"""Shared verify step: target-verify forward, sampling, acceptance bookkeeping.
|
"""Shared verify step: target-verify forward, sampling, acceptance bookkeeping.
|
||||||
|
|
||||||
@@ -527,13 +528,21 @@ def run_eagle_verify(
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
# Prepare grammar data on CPU if needed
|
# Prepare grammar data on CPU if needed. Use async pinned D2H copies (not
|
||||||
|
# blocking .cpu()) and record an event. The copies are issued before the
|
||||||
|
# target verify launch below so they run right after the draft, but the
|
||||||
|
# host does not block here. We wait on grammar_copy_done only just before
|
||||||
|
# the CPU bitmask traversal reads the buffers, so the traversal (and these
|
||||||
|
# copies) overlap the target verify forward instead of stalling the GPU.
|
||||||
|
grammar_copy_done = None
|
||||||
if batch.has_grammar:
|
if batch.has_grammar:
|
||||||
retrieve_next_token_cpu = verify_input.retrieve_next_token.cpu()
|
retrieve_next_token_cpu = _async_d2h(verify_input.retrieve_next_token)
|
||||||
retrieve_next_sibling_cpu = verify_input.retrieve_next_sibling.cpu()
|
retrieve_next_sibling_cpu = _async_d2h(verify_input.retrieve_next_sibling)
|
||||||
draft_tokens_cpu = verify_input.draft_token.view(
|
draft_tokens_cpu = _async_d2h(
|
||||||
verify_input.retrieve_next_token.shape
|
verify_input.draft_token.view(verify_input.retrieve_next_token.shape)
|
||||||
).cpu()
|
)
|
||||||
|
grammar_copy_done = torch.get_device_module(device).Event()
|
||||||
|
grammar_copy_done.record()
|
||||||
|
|
||||||
if metadata_ready_pre_pad:
|
if metadata_ready_pre_pad:
|
||||||
# Multi-layer eagle preserved-verbatim behavior: metadata init is
|
# Multi-layer eagle preserved-verbatim behavior: metadata init is
|
||||||
@@ -561,6 +570,17 @@ def run_eagle_verify(
|
|||||||
# Generate vocab mask for constrained decoding
|
# Generate vocab mask for constrained decoding
|
||||||
vocab_mask = None
|
vocab_mask = None
|
||||||
if batch.has_grammar:
|
if batch.has_grammar:
|
||||||
|
# Grammar barrier: advance the previous batch's grammar FSM over its
|
||||||
|
# committed tokens before building this batch's bitmask. Runs after the
|
||||||
|
# target forward launch, so the FSM advance and the traversal below both
|
||||||
|
# overlap the target verify forward. No-op if there is nothing pending.
|
||||||
|
if grammar_barrier is not None:
|
||||||
|
grammar_barrier()
|
||||||
|
# Wait for the async draft/verify-input D2H copies above to land before
|
||||||
|
# the CPU traversal reads them. The event was recorded right after the
|
||||||
|
# copies (before the target verify launch), so this wait — and the
|
||||||
|
# traversal below — overlap the target verify forward.
|
||||||
|
grammar_copy_done.synchronize()
|
||||||
# Generate the logit mask for structured output.
|
# Generate the logit mask for structured output.
|
||||||
vocab_mask = generate_token_bitmask(
|
vocab_mask = generate_token_bitmask(
|
||||||
batch.reqs,
|
batch.reqs,
|
||||||
@@ -573,7 +593,12 @@ def run_eagle_verify(
|
|||||||
|
|
||||||
if vocab_mask is not None:
|
if vocab_mask is not None:
|
||||||
assert verify_input.grammar is not None
|
assert verify_input.grammar is not None
|
||||||
vocab_mask = vocab_mask.to(verify_input.retrieve_next_token.device)
|
# non_blocking H2D so the mask copy overlaps the tail of the target
|
||||||
|
# verify forward instead of syncing the host; stream ordering keeps
|
||||||
|
# it before eagle_sample's apply_vocab_mask below.
|
||||||
|
vocab_mask = vocab_mask.to(
|
||||||
|
verify_input.retrieve_next_token.device, non_blocking=True
|
||||||
|
)
|
||||||
# NOTE: otherwise, this vocab mask will be the one from the previous extend stage
|
# NOTE: otherwise, this vocab mask will be the one from the previous extend stage
|
||||||
# and will be applied to produce wrong results
|
# and will be applied to produce wrong results
|
||||||
batch.sampling_info.vocab_mask = None
|
batch.sampling_info.vocab_mask = None
|
||||||
|
|||||||
@@ -1119,7 +1119,9 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward_batch_generation(self, batch: ScheduleBatch, on_publish=None):
|
def forward_batch_generation(
|
||||||
|
self, batch: ScheduleBatch, on_publish=None, grammar_barrier=None
|
||||||
|
):
|
||||||
if batch.forward_mode.is_extend() or batch.is_extend_in_batch:
|
if batch.forward_mode.is_extend() or batch.is_extend_in_batch:
|
||||||
# Target prefill
|
# Target prefill
|
||||||
target_capture_mode = (
|
target_capture_mode = (
|
||||||
@@ -1192,7 +1194,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
verify_input: EagleVerifyInput = self.draft_worker.draft(batch)
|
verify_input: EagleVerifyInput = self.draft_worker.draft(batch)
|
||||||
assert verify_input.is_verify_input()
|
assert verify_input.is_verify_input()
|
||||||
batch.spec_info = verify_input
|
batch.spec_info = verify_input
|
||||||
batch_output = self.verify(batch)
|
batch_output = self.verify(batch, grammar_barrier=grammar_barrier)
|
||||||
# Publish before draft_extend so the fence is at verify-end.
|
# Publish before draft_extend so the fence is at verify-end.
|
||||||
if on_publish is not None:
|
if on_publish is not None:
|
||||||
on_publish(batch_output.new_seq_lens)
|
on_publish(batch_output.new_seq_lens)
|
||||||
@@ -1495,7 +1497,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
)
|
)
|
||||||
dw._rebuild_topk1_chain_buffers()
|
dw._rebuild_topk1_chain_buffers()
|
||||||
|
|
||||||
def verify(self, batch: ScheduleBatch):
|
def verify(self, batch: ScheduleBatch, grammar_barrier=None):
|
||||||
return run_eagle_verify(
|
return run_eagle_verify(
|
||||||
batch,
|
batch,
|
||||||
target_worker=self.target_worker,
|
target_worker=self.target_worker,
|
||||||
@@ -1509,6 +1511,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
metadata_ready_pre_pad=False,
|
metadata_ready_pre_pad=False,
|
||||||
finalize_tree_path=True,
|
finalize_tree_path=True,
|
||||||
|
grammar_barrier=grammar_barrier,
|
||||||
)
|
)
|
||||||
|
|
||||||
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
|
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
|
||||||
|
|||||||
@@ -708,7 +708,9 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
|
|||||||
self._draft_worker.draft_attn_backend,
|
self._draft_worker.draft_attn_backend,
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward_batch_generation(self, batch: ScheduleBatch, on_publish=None):
|
def forward_batch_generation(
|
||||||
|
self, batch: ScheduleBatch, on_publish=None, grammar_barrier=None
|
||||||
|
):
|
||||||
# Mirrors EAGLEWorkerV2.forward_batch_generation; the only frozen-specific
|
# Mirrors EAGLEWorkerV2.forward_batch_generation; the only frozen-specific
|
||||||
# change is the idle draft-input (FrozenKVMTPDraftInput + recurrent hidden
|
# change is the idle draft-input (FrozenKVMTPDraftInput + recurrent hidden
|
||||||
# size). The draft / seed-based draft-extend hooks are FrozenKVMTPDraftWorker's.
|
# size). The draft / seed-based draft-extend hooks are FrozenKVMTPDraftWorker's.
|
||||||
@@ -758,7 +760,7 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
|
|||||||
verify_input = self.draft_worker.draft(batch)
|
verify_input = self.draft_worker.draft(batch)
|
||||||
assert verify_input.is_verify_input()
|
assert verify_input.is_verify_input()
|
||||||
batch.spec_info = verify_input
|
batch.spec_info = verify_input
|
||||||
batch_output = self.verify(batch)
|
batch_output = self.verify(batch, grammar_barrier=grammar_barrier)
|
||||||
# Publish before draft-extend so the fence is at verify-end.
|
# Publish before draft-extend so the fence is at verify-end.
|
||||||
if on_publish is not None:
|
if on_publish is not None:
|
||||||
on_publish(batch_output.new_seq_lens)
|
on_publish(batch_output.new_seq_lens)
|
||||||
|
|||||||
@@ -927,7 +927,9 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward_batch_generation(self, batch: ScheduleBatch, on_publish=None):
|
def forward_batch_generation(
|
||||||
|
self, batch: ScheduleBatch, on_publish=None, grammar_barrier=None
|
||||||
|
):
|
||||||
if batch.forward_mode.is_extend() or batch.is_extend_in_batch:
|
if batch.forward_mode.is_extend() or batch.is_extend_in_batch:
|
||||||
# Target prefill
|
# Target prefill
|
||||||
target_capture_mode = (
|
target_capture_mode = (
|
||||||
@@ -974,14 +976,14 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
verify_input: EagleVerifyInput = self.draft_worker.draft(batch)
|
verify_input: EagleVerifyInput = self.draft_worker.draft(batch)
|
||||||
assert verify_input.is_verify_input()
|
assert verify_input.is_verify_input()
|
||||||
batch.spec_info = verify_input
|
batch.spec_info = verify_input
|
||||||
batch_output = self.verify(batch)
|
batch_output = self.verify(batch, grammar_barrier=grammar_barrier)
|
||||||
# Publish before draft_extend so the fence is at verify-end.
|
# Publish before draft_extend so the fence is at verify-end.
|
||||||
if on_publish is not None:
|
if on_publish is not None:
|
||||||
on_publish(batch_output.new_seq_lens)
|
on_publish(batch_output.new_seq_lens)
|
||||||
self.draft_worker._draft_extend_for_decode(batch, batch_output)
|
self.draft_worker._draft_extend_for_decode(batch, batch_output)
|
||||||
return batch_output
|
return batch_output
|
||||||
|
|
||||||
def verify(self, batch: ScheduleBatch):
|
def verify(self, batch: ScheduleBatch, grammar_barrier=None):
|
||||||
return run_eagle_verify(
|
return run_eagle_verify(
|
||||||
batch,
|
batch,
|
||||||
target_worker=self.target_worker,
|
target_worker=self.target_worker,
|
||||||
@@ -995,4 +997,5 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
metadata_ready_pre_pad=False,
|
metadata_ready_pre_pad=False,
|
||||||
finalize_tree_path=False,
|
finalize_tree_path=False,
|
||||||
|
grammar_barrier=grammar_barrier,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -133,6 +133,11 @@ class SpeculativeAlgorithm(Enum):
|
|||||||
graphs in the decode cuda graph runner."""
|
graphs in the decode cuda graph runner."""
|
||||||
return self.is_dspark()
|
return self.is_dspark()
|
||||||
|
|
||||||
|
def supports_grammar_overlap(self) -> bool:
|
||||||
|
# Whether the worker advances the grammar FSM inside verify() (via the
|
||||||
|
# scheduler's grammar barrier), letting spec + grammar decode overlap.
|
||||||
|
return self.is_eagle()
|
||||||
|
|
||||||
def has_draft_kv(self) -> bool:
|
def has_draft_kv(self) -> bool:
|
||||||
"""Whether the draft phase writes KV chains. NGRAM does not (its tree
|
"""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
|
lives only in the verify mask), so per-decode KV sizing needs no
|
||||||
|
|||||||
@@ -95,6 +95,11 @@ class CustomSpecAlgo:
|
|||||||
def supports_ragged_verify(self) -> bool:
|
def supports_ragged_verify(self) -> bool:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
def supports_grammar_overlap(self) -> bool:
|
||||||
|
# Whether the worker advances the grammar FSM inside verify() (via the
|
||||||
|
# scheduler's grammar barrier), letting spec + grammar decode overlap.
|
||||||
|
return False
|
||||||
|
|
||||||
def has_draft_kv(self) -> bool:
|
def has_draft_kv(self) -> bool:
|
||||||
# Conservative default: the larger KV reserve.
|
# Conservative default: the larger KV reserve.
|
||||||
return True
|
return True
|
||||||
|
|||||||
@@ -43,9 +43,19 @@ class _FakeSpecAlgorithm:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeForwardMode:
|
||||||
|
def is_decode(self) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
def is_extend(self) -> bool:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
class _FakeBatch:
|
class _FakeBatch:
|
||||||
def __init__(self, reqs):
|
def __init__(self, reqs):
|
||||||
self.reqs = reqs
|
self.reqs = reqs
|
||||||
|
self.has_grammar = any(req.grammar is not None for req in reqs)
|
||||||
|
self.forward_mode = _FakeForwardMode()
|
||||||
self.spec_algorithm = _FakeSpecAlgorithm()
|
self.spec_algorithm = _FakeSpecAlgorithm()
|
||||||
|
|
||||||
|
|
||||||
@@ -95,6 +105,8 @@ def _make_result(num_draft_tokens, accept_lens, flat_tokens):
|
|||||||
num_correct_drafts_per_req_cpu=None,
|
num_correct_drafts_per_req_cpu=None,
|
||||||
block_accept_lens=None,
|
block_accept_lens=None,
|
||||||
cap_lens=None,
|
cap_lens=None,
|
||||||
|
copy_done=None,
|
||||||
|
grammar_advanced=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user