[Spec] Share the grammar mask build and verify-tree staging across spec workers (#32393)

This commit is contained in:
Liangsheng Yin
2026-07-25 03:40:21 -07:00
committed by GitHub
parent 3c5bf1f6d2
commit 17afd8421f
3 changed files with 108 additions and 63 deletions
@@ -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, _async_d2h from sglang.srt.managers.utils import GenerationBatchResult
from sglang.srt.model_executor.forward_batch_info import ( from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode, CaptureHiddenMode,
ForwardBatch, ForwardBatch,
@@ -23,8 +23,9 @@ from sglang.srt.speculative.eagle_utils import (
eagle_sample, eagle_sample,
) )
from sglang.srt.speculative.spec_utils import ( from sglang.srt.speculative.spec_utils import (
GrammarTree,
build_grammar_vocab_mask,
commit_mamba_states_after_verify, commit_mamba_states_after_verify,
generate_token_bitmask,
move_accept_tokens_to_target_kvcache, move_accept_tokens_to_target_kvcache,
record_stream_each, record_stream_each,
record_stream_for_v2_verify, record_stream_for_v2_verify,
@@ -528,21 +529,16 @@ def run_eagle_verify(
), ),
) )
# Prepare grammar data on CPU if needed. Use async pinned D2H copies (not # Must stay ahead of the target verify launch below.
# blocking .cpu()) and record an event. The copies are issued before the grammar_tree = (
# target verify launch below so they run right after the draft, but the GrammarTree.from_device(
# host does not block here. We wait on grammar_copy_done only just before verify_input.retrieve_next_token,
# the CPU bitmask traversal reads the buffers, so the traversal (and these verify_input.retrieve_next_sibling,
# copies) overlap the target verify forward instead of stalling the GPU. verify_input.draft_token.view(verify_input.retrieve_next_token.shape),
grammar_copy_done = None )
if batch.has_grammar: if batch.has_grammar
retrieve_next_token_cpu = _async_d2h(verify_input.retrieve_next_token) else None
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: if metadata_ready_pre_pad:
# Multi-layer eagle preserved-verbatim behavior: metadata init is # Multi-layer eagle preserved-verbatim behavior: metadata init is
@@ -576,33 +572,14 @@ def run_eagle_verify(
# overlap the target verify forward. No-op if there is nothing pending. # overlap the target verify forward. No-op if there is nothing pending.
if grammar_barrier is not None: if grammar_barrier is not None:
grammar_barrier() grammar_barrier()
# Wait for the async draft/verify-input D2H copies above to land before vocab_mask = build_grammar_vocab_mask(
# the CPU traversal reads them. The event was recorded right after the reqs=batch.reqs,
# copies (before the target verify launch), so this wait — and the verify_input=verify_input,
# traversal below — overlap the target verify forward. tree=grammar_tree,
grammar_copy_done.synchronize() sampling_info=batch.sampling_info,
# Generate the logit mask for structured output. device=verify_input.retrieve_next_token.device,
vocab_mask = generate_token_bitmask(
batch.reqs,
verify_input,
retrieve_next_token_cpu,
retrieve_next_sibling_cpu,
draft_tokens_cpu,
batch.sampling_info.vocab_size,
) )
if vocab_mask is not None:
assert verify_input.grammar is not None
# 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
# Sample # Sample
maybe_detect_nan(logits_output.next_token_logits, "verify: target model logits") maybe_detect_nan(logits_output.next_token_logits, "verify: target model logits")
maybe_detect_inf(logits_output.next_token_logits, "verify: target model logits") maybe_detect_inf(logits_output.next_token_logits, "verify: target model logits")
+10 -20
View File
@@ -21,8 +21,9 @@ from sglang.srt.speculative.cpp_ngram.ngram_corpus import NgramCorpus
from sglang.srt.speculative.eagle_utils import eagle_sample from sglang.srt.speculative.eagle_utils import eagle_sample
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 (
GrammarTree,
build_grammar_vocab_mask,
commit_mamba_states_after_verify, commit_mamba_states_after_verify,
generate_token_bitmask,
move_accept_tokens_to_target_kvcache, move_accept_tokens_to_target_kvcache,
prepare_mamba_track_for_verify, prepare_mamba_track_for_verify,
record_stream_for_v2_verify, record_stream_for_v2_verify,
@@ -430,29 +431,18 @@ class NGRAMWorker(BaseSpecWorker):
retrieve_next_token_cpu, retrieve_next_sibling_cpu = _derive_tree_links( retrieve_next_token_cpu, retrieve_next_sibling_cpu = _derive_tree_links(
mask, bs, self.draft_token_num mask, bs, self.draft_token_num
) )
draft_tokens_cpu = ( vocab_mask = build_grammar_vocab_mask(
torch.from_numpy(req_drafts).to(torch.int64).view(bs, -1) reqs=batch.reqs,
) verify_input=verify_input,
vocab_mask = generate_token_bitmask( tree=GrammarTree.from_host(
batch.reqs,
verify_input,
retrieve_next_token_cpu, retrieve_next_token_cpu,
retrieve_next_sibling_cpu, retrieve_next_sibling_cpu,
draft_tokens_cpu, torch.from_numpy(req_drafts).to(torch.int64).view(bs, -1),
batch.sampling_info.vocab_size, ),
sampling_info=batch.sampling_info,
device=verify_input.retrieve_next_token.device,
) )
if vocab_mask is not None:
assert verify_input.grammar is not None
# non_blocking is safe: the bitmask source is pinned, and stream
# order keeps the copy ahead of apply_vocab_mask.
vocab_mask = vocab_mask.to(
verify_input.retrieve_next_token.device, non_blocking=True
)
# NOTE (sk): 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
# Sample # Sample
maybe_detect_nan( maybe_detect_nan(
logits_output.next_token_logits, "verify: target model logits" logits_output.next_token_logits, "verify: target model logits"
@@ -38,6 +38,7 @@ from sglang.srt.distributed.parallel_state import (
) )
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.schedule_batch import set_mamba_track_indices_from_reqs from sglang.srt.managers.schedule_batch import set_mamba_track_indices_from_reqs
from sglang.srt.managers.utils import _async_d2h
from sglang.srt.mem_cache.allocation import ( from sglang.srt.mem_cache.allocation import (
assign_req_to_token_pool as assign_req_to_token_pool, assign_req_to_token_pool as assign_req_to_token_pool,
) )
@@ -69,8 +70,10 @@ if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
from sglang.srt.managers.tp_worker import TpModelWorker from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.eagle_info import EagleVerifyInput from sglang.srt.speculative.eagle_info import EagleVerifyInput
from sglang.srt.speculative.spec_info import SpecInput
if _is_cuda: if _is_cuda:
@@ -549,6 +552,81 @@ def generate_token_bitmask(
return allocate_token_bitmask return allocate_token_bitmask
class GrammarTree:
"""The verify tree the grammar bitmask is built over, on the host.
``from_device`` starts an async copy, so build it before the target verify
launch; ``from_host`` is for algorithms that build the tree there (NGRAM).
"""
def __init__(self, host: Tuple[torch.Tensor, ...], done_event):
self._host = host
self._done = done_event
@classmethod
def from_device(
cls,
retrieve_next_token: torch.Tensor,
retrieve_next_sibling: torch.Tensor,
draft_token: torch.Tensor,
) -> GrammarTree:
tensors = (retrieve_next_token, retrieve_next_sibling, draft_token)
host = tuple(_async_d2h(t) for t in tensors)
# Sources may be mixed -- an algorithm can synthesize part of the tree on
# the host -- so the event has to key off whichever one is on device.
device = next((t.device for t in tensors if t.device.type != "cpu"), None)
if device is None:
return cls(host, None)
done = torch.get_device_module(device).Event()
done.record()
return cls(host, done)
@classmethod
def from_host(
cls,
retrieve_next_token: torch.Tensor,
retrieve_next_sibling: torch.Tensor,
draft_token: torch.Tensor,
) -> GrammarTree:
return cls((retrieve_next_token, retrieve_next_sibling, draft_token), None)
def resolve(self) -> Tuple[torch.Tensor, ...]:
if self._done is not None:
self._done.synchronize()
return self._host
def build_grammar_vocab_mask(
*,
reqs: List[Req],
verify_input: SpecInput,
tree: GrammarTree,
sampling_info: SamplingBatchInfo,
device,
) -> Optional[torch.Tensor]:
"""Build the constrained-decoding bitmask over a verify tree and stage it on device.
Call it after the target verify launch: resolving the tree and traversing it are
both host work, so both overlap that forward.
"""
vocab_mask = generate_token_bitmask(
reqs,
verify_input,
*tree.resolve(),
sampling_info.vocab_size,
)
if vocab_mask is None:
return None
assert verify_input.grammar is not None
# non_blocking is safe: the bitmask is pinned (see xgrammar_backend), and stream
# order keeps the copy ahead of the sampler's apply_vocab_mask.
vocab_mask = vocab_mask.to(device, non_blocking=True)
# Otherwise the extend stage's leftover mask is applied instead.
sampling_info.vocab_mask = None
return vocab_mask
def load_token_map(token_map_path: str) -> List[int]: def load_token_map(token_map_path: str) -> List[int]:
if not os.path.exists(token_map_path): if not os.path.exists(token_map_path):
repo_id = os.path.dirname(token_map_path) repo_id = os.path.dirname(token_map_path)