Support grammar + spec + reasoning (#14163)
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user