Support grammar + spec + reasoning (#14163)

This commit is contained in:
Liangsheng Yin
2025-11-30 21:19:57 +08:00
committed by GitHub
parent 340c613ab5
commit 0a9d64530d
5 changed files with 164 additions and 8 deletions

View File

@@ -29,14 +29,35 @@ class ReasonerGrammarObject(BaseGrammarObject):
super().__init__()
self.grammar = grammar
self.think_end_id = think_end_id
self.is_in_reasoning = True
# -1 means thinking has not ended yet
# 0 means just ended thinking in the last token
# + means number of tokens after thinking ended
self.tokens_after_think_end = -1
def transfer_state(self, token: int) -> int:
if self.tokens_after_think_end == -1 and token == self.think_end_id:
self.tokens_after_think_end = 0
elif self.tokens_after_think_end >= 0:
self.tokens_after_think_end += 1
def rollback_state(self):
if self.tokens_after_think_end == 0:
self.tokens_after_think_end = -1
elif self.tokens_after_think_end > 0:
self.tokens_after_think_end -= 1
def accept_token(self, token: int):
if token == self.think_end_id:
self.is_in_reasoning = False
if not self.is_in_reasoning and token != self.think_end_id:
if self.tokens_after_think_end >= 0:
self.grammar.accept_token(token)
self.transfer_state(token)
def rollback(self, k):
steps_after_think = min(k, self.tokens_after_think_end)
if steps_after_think > 0:
self.grammar.rollback(steps_after_think)
for _ in range(k):
self.rollback_state()
def allocate_vocab_mask(
self, vocab_size: int, batch_size: int, device
@@ -44,7 +65,7 @@ class ReasonerGrammarObject(BaseGrammarObject):
return self.grammar.allocate_vocab_mask(vocab_size, batch_size, device)
def fill_vocab_mask(self, vocab_mask: torch.Tensor, idx: int) -> None:
if not self.is_in_reasoning:
if self.tokens_after_think_end >= 0:
self.grammar.fill_vocab_mask(vocab_mask, idx)
def move_vocab_mask(self, vocab_mask: torch.Tensor, device) -> torch.Tensor:

View File

@@ -305,3 +305,39 @@ class XGrammarGrammarBackend(BaseGrammarBackend):
def reset(self):
self.grammar_compiler.clear_cache()
def demo_test():
from transformers import AutoConfig, AutoTokenizer
from sglang.test.test_utils import DEFAULT_MODEL_NAME_FOR_TEST
tokenizer = AutoTokenizer.from_pretrained(DEFAULT_MODEL_NAME_FOR_TEST)
hf_config = AutoConfig.from_pretrained(DEFAULT_MODEL_NAME_FOR_TEST)
# Should use vocab size from model config
vocab_size = hf_config.vocab_size
eos_token_id = tokenizer.eos_token_id
backend = XGrammarGrammarBackend(
tokenizer, vocab_size=vocab_size, model_eos_token_ids=[eos_token_id]
)
regex = r"hello (world|there)"
grammar = backend.dispatch_regex(regex)
tokens = [
tokenizer.encode(t, add_special_tokens=False)[0] for t in ["hello", " world"]
]
# Test termination
grammar.accept_token(tokens[0]) # accept "hello"
grammar.accept_token(tokens[1]) # accept " world"
grammar.accept_token(eos_token_id) # accept EOS
assert grammar.is_terminated()
# Test rollback the terminated state
grammar.rollback(1)
assert not grammar.is_terminated()
if __name__ == "__main__":
demo_test()

View File

@@ -581,8 +581,6 @@ def traverse_tree(
retrieve_next_token.shape == retrieve_next_sibling.shape == draft_tokens.shape
)
allocate_token_bitmask.fill_(0)
def dfs(
curr: int,
retrieve_next_token: torch.Tensor,