diff --git a/python/sglang/srt/constrained/xgrammar_backend.py b/python/sglang/srt/constrained/xgrammar_backend.py index 19fdc4bbf..a0b44930c 100644 --- a/python/sglang/srt/constrained/xgrammar_backend.py +++ b/python/sglang/srt/constrained/xgrammar_backend.py @@ -26,7 +26,8 @@ from xgrammar import ( StructuralTag, StructuralTagItem, TokenizerInfo, - allocate_token_bitmask, + bitmask_dtype, + get_bitmask_shape, ) from sglang.srt.constrained.base_grammar_backend import ( @@ -56,6 +57,18 @@ logger = logging.getLogger(__name__) 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): def __init__( @@ -100,7 +113,7 @@ class XGrammarGrammar(BaseGrammarObject): def allocate_vocab_mask( self, 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) def fill_vocab_mask(self, vocab_mask: torch.Tensor, idx: int) -> None: self.matcher.fill_next_token_bitmask(vocab_mask, idx) @@ -228,7 +241,7 @@ class XGrammarGrammarBackend(BaseGrammarBackend): @staticmethod 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 def move_vocab_mask(vocab_mask: torch.Tensor, device) -> torch.Tensor: diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index a8cea65ec..14dbcfefe 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -3008,6 +3008,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): forward_mode=self.forward_mode, out_cache_loc=self.out_cache_loc, return_logprob=self.return_logprob, + has_grammar=self.has_grammar, decoding_reqs=self.decoding_reqs, spec_algorithm=self.spec_algorithm, spec_info=self.spec_info, diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index c46a8c596..baa191862 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -1646,19 +1646,34 @@ class Scheduler( and last_batch_is_extend ) - # We do not support overlap + spec + grammar yet, - # so we need to turn off overlap for this batch. - # TODO(lsyin): support overlap + spec + grammar + # Spec algorithms that don't advance the grammar FSM inside verify() (see + # supports_grammar_overlap) still need overlap forced off for grammar decode + # batches, so the FSM is advanced before the next batch's bitmask. need_grammar_sync = ( batch and not batch.spec_algorithm.is_none() + and not batch.spec_algorithm.supports_grammar_overlap() and batch.has_grammar and batch.forward_mode.is_decode() 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 + 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") def process_input_requests(self, recv_reqs: List): now = time.monotonic() @@ -3337,15 +3352,19 @@ class Scheduler( # Spec_v2 fires on_publish mid-worker (between verify and # draft_extend) so schedule prep can overlap with draft_extend. # Non-spec has no later work — scheduler publishes after return. - fwd_kwargs = ( - { - "on_publish": partial( - self.future_map.publish, future_indices + fwd_kwargs = {} + if not batch.spec_algorithm.is_none(): + fwd_kwargs["on_publish"] = partial( + self.future_map.publish, future_indices + ) + # Grammar-overlap-capable workers advance the grammar FSM + # inside verify() before building the bitmask; hand them the + # 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 ) - } - if not batch.spec_algorithm.is_none() - else {} - ) # FIXME: pp is not compatible with overlap batch_result = self.model_worker.forward_batch_generation( diff --git a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py index aca94e616..b68d59161 100644 --- a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py @@ -273,7 +273,9 @@ class SchedulerBatchResultProcessor: if req.grammar is not None: 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: @@ -487,17 +489,24 @@ class SchedulerBatchResultProcessor: ) return hidden_state_offset - def _apply_prefill_grammar(self, *, req: Req, next_token_id: int) -> None: - # FIXME: this try-except block is for handling unexpected xgrammar issue. - try: - req.grammar.accept_token(next_token_id) - except ValueError as e: - # Grammar accept_token can raise ValueError if the token is not in the grammar. - # This can happen if the grammar is not set correctly or the token is invalid. - logger.error( - f"Grammar accept_token failed for req {req.rid} with token {next_token_id}: {e}" - ) - req.to_finish = FINISH_ABORT() + 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. + try: + req.grammar.accept_token(next_token_id) + except ValueError as e: + # Grammar accept_token can raise ValueError if the token is not in the + # grammar. This can happen if the grammar is not set correctly or the + # token is invalid. + logger.error( + f"Grammar accept_token failed for req {req.rid} with token {next_token_id}: {e}" + ) + req.to_finish = FINISH_ABORT() req.grammar.finished = req.finished() def _apply_chunked_prefill_logprobs( @@ -565,6 +574,13 @@ class SchedulerBatchResultProcessor: 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 = [] # 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. @@ -580,11 +596,9 @@ class SchedulerBatchResultProcessor: pass else: if req.grammar is not None: - # Stop accepting once the grammar terminates, so the - # over-drafted suffix is never committed to KV nor emitted. - # This advances the grammar FSM; the result loop only syncs - # grammar.finished. - accept_tokens = self._accept_grammar_tokens(req, accept_tokens) + # FSM already advanced + truncated by advance_grammar_fsm; reuse + # the retained (grammar-legal) run instead of advancing again. + accept_tokens = result.grammar_retained_tokens[i] # Commit the full accepted run (drafts + bonus). num_accept_tokens = len(accept_tokens) @@ -636,6 +650,64 @@ class SchedulerBatchResultProcessor: req.to_finish = FINISH_ABORT() 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( self, batch: ScheduleBatch, diff --git a/python/sglang/srt/managers/utils.py b/python/sglang/srt/managers/utils.py index 74d2a4699..3ef02c784 100644 --- a/python/sglang/srt/managers/utils.py +++ b/python/sglang/srt/managers/utils.py @@ -66,6 +66,13 @@ class GenerationBatchResult: future_indices: Optional[torch.Tensor] = 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? # sync path: forward stream -> output processor accept_lens: Optional[torch.Tensor] = None diff --git a/python/sglang/srt/speculative/eagle_worker_common.py b/python/sglang/srt/speculative/eagle_worker_common.py index 9f1303464..34b627043 100644 --- a/python/sglang/srt/speculative/eagle_worker_common.py +++ b/python/sglang/srt/speculative/eagle_worker_common.py @@ -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.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 ( CaptureHiddenMode, ForwardBatch, @@ -467,6 +467,7 @@ def run_eagle_verify( device: str, metadata_ready_pre_pad: bool, finalize_tree_path: bool, + grammar_barrier=None, ) -> GenerationBatchResult: """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: - retrieve_next_token_cpu = verify_input.retrieve_next_token.cpu() - retrieve_next_sibling_cpu = verify_input.retrieve_next_sibling.cpu() - draft_tokens_cpu = verify_input.draft_token.view( - verify_input.retrieve_next_token.shape - ).cpu() + retrieve_next_token_cpu = _async_d2h(verify_input.retrieve_next_token) + retrieve_next_sibling_cpu = _async_d2h(verify_input.retrieve_next_sibling) + draft_tokens_cpu = _async_d2h( + verify_input.draft_token.view(verify_input.retrieve_next_token.shape) + ) + grammar_copy_done = torch.get_device_module(device).Event() + grammar_copy_done.record() if metadata_ready_pre_pad: # Multi-layer eagle preserved-verbatim behavior: metadata init is @@ -561,6 +570,17 @@ def run_eagle_verify( # Generate vocab mask for constrained decoding vocab_mask = None 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. vocab_mask = generate_token_bitmask( batch.reqs, @@ -573,7 +593,12 @@ def run_eagle_verify( if vocab_mask 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 # and will be applied to produce wrong results batch.sampling_info.vocab_mask = None diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index e8493012c..3c743f7fb 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -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: # Target prefill target_capture_mode = ( @@ -1192,7 +1194,7 @@ class EAGLEWorkerV2(BaseSpecWorker): verify_input: EagleVerifyInput = self.draft_worker.draft(batch) assert verify_input.is_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. if on_publish is not None: on_publish(batch_output.new_seq_lens) @@ -1495,7 +1497,7 @@ class EAGLEWorkerV2(BaseSpecWorker): ) dw._rebuild_topk1_chain_buffers() - def verify(self, batch: ScheduleBatch): + def verify(self, batch: ScheduleBatch, grammar_barrier=None): return run_eagle_verify( batch, target_worker=self.target_worker, @@ -1509,6 +1511,7 @@ class EAGLEWorkerV2(BaseSpecWorker): device=self.device, metadata_ready_pre_pad=False, finalize_tree_path=True, + grammar_barrier=grammar_barrier, ) def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput): diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py b/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py index e85853aac..98c210b4e 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py @@ -708,7 +708,9 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2): 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 # change is the idle draft-input (FrozenKVMTPDraftInput + recurrent hidden # 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) assert verify_input.is_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. if on_publish is not None: on_publish(batch_output.new_seq_lens) diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index 309efdc30..fb2e1de93 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -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: # Target prefill target_capture_mode = ( @@ -974,14 +976,14 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker): verify_input: EagleVerifyInput = self.draft_worker.draft(batch) assert verify_input.is_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. if on_publish is not None: on_publish(batch_output.new_seq_lens) self.draft_worker._draft_extend_for_decode(batch, batch_output) return batch_output - def verify(self, batch: ScheduleBatch): + def verify(self, batch: ScheduleBatch, grammar_barrier=None): return run_eagle_verify( batch, target_worker=self.target_worker, @@ -995,4 +997,5 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker): device=self.device, metadata_ready_pre_pad=False, finalize_tree_path=False, + grammar_barrier=grammar_barrier, ) diff --git a/python/sglang/srt/speculative/spec_info.py b/python/sglang/srt/speculative/spec_info.py index c9d27b760..558e88f61 100644 --- a/python/sglang/srt/speculative/spec_info.py +++ b/python/sglang/srt/speculative/spec_info.py @@ -133,6 +133,11 @@ class SpeculativeAlgorithm(Enum): graphs in the decode cuda graph runner.""" 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: """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 diff --git a/python/sglang/srt/speculative/spec_registry.py b/python/sglang/srt/speculative/spec_registry.py index f1b3c2447..4ab16a2a3 100644 --- a/python/sglang/srt/speculative/spec_registry.py +++ b/python/sglang/srt/speculative/spec_registry.py @@ -95,6 +95,11 @@ class CustomSpecAlgo: def supports_ragged_verify(self) -> bool: 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: # Conservative default: the larger KV reserve. return True diff --git a/test/registered/unit/managers/test_batch_result_processor_spec_grammar.py b/test/registered/unit/managers/test_batch_result_processor_spec_grammar.py index 0507b350e..ef5e2056a 100644 --- a/test/registered/unit/managers/test_batch_result_processor_spec_grammar.py +++ b/test/registered/unit/managers/test_batch_result_processor_spec_grammar.py @@ -43,9 +43,19 @@ class _FakeSpecAlgorithm: return False +class _FakeForwardMode: + def is_decode(self) -> bool: + return True + + def is_extend(self) -> bool: + return False + + class _FakeBatch: def __init__(self, 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() @@ -95,6 +105,8 @@ def _make_result(num_draft_tokens, accept_lens, flat_tokens): num_correct_drafts_per_req_cpu=None, block_accept_lens=None, cap_lens=None, + copy_done=None, + grammar_advanced=False, )