[Auto Sync] Update base_grammar_backend.py, llguidance_back... (20250911) (#10333)

Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
This commit is contained in:
Lianmin Zheng
2025-09-12 12:55:55 -07:00
committed by GitHub
co-authored by github-actions[bot] <github-actions[bot]@users.noreply.github.com>
parent 151e287d1a
commit 2269cf1e2f
5 changed files with 77 additions and 25 deletions
@@ -13,6 +13,7 @@
# ==============================================================================
"""Constrained decoding with xgrammar backend."""
import dataclasses
import json
import logging
from typing import List, Optional, Tuple, Union
@@ -31,6 +32,7 @@ from sglang.srt.constrained.base_grammar_backend import (
INVALID_GRAMMAR_OBJ,
BaseGrammarBackend,
BaseGrammarObject,
GrammarStats,
)
from sglang.srt.utils import is_hip
@@ -41,9 +43,9 @@ else:
from sglang.srt.constrained.triton_ops.bitmask_ops import (
apply_token_bitmask_inplace_triton,
)
logger = logging.getLogger(__name__)
MAX_ROLLBACK_TOKENS = 200
@@ -56,17 +58,20 @@ class XGrammarGrammar(BaseGrammarObject):
ctx: CompiledGrammar,
override_stop_tokens: Optional[Union[List[int], int]],
key_string: Optional[str] = None, # TODO (sk): for debugging, remove later
grammar_stats: Optional[GrammarStats] = GrammarStats(),
) -> None:
super().__init__()
self.matcher = matcher
self.vocab_size = vocab_size
self.ctx = ctx
self.override_stop_tokens = override_stop_tokens
self.finished = False
self.accepted_tokens = []
self.key_string = key_string
self.grammar_stats = grammar_stats
def accept_token(self, token: int):
if not self.is_terminated():
self.current_token = token
accepted = self.matcher.accept_token(token)
if not accepted:
# log for debugging
@@ -120,6 +125,9 @@ class XGrammarGrammar(BaseGrammarObject):
self.ctx,
self.override_stop_tokens,
self.key_string,
dataclasses.replace(
self.grammar_stats, is_cache_hit=True, tree_traversal_time=[]
),
)
def try_jump_forward(self, tokenizer) -> Optional[Tuple[List[int], str]]:
@@ -150,7 +158,7 @@ class XGrammarGrammar(BaseGrammarObject):
assert self.matcher.accept_token(new_output_ids[i])
def __repr__(self):
return f"XGrammarGrammar({self.key_string=}, {self.accepted_tokens=})"
return f"XGrammarGrammar({self.key_string=}, {self.accepted_tokens=}, {self.current_token=})"
class XGrammarGrammarBackend(BaseGrammarBackend):
@@ -177,14 +185,21 @@ class XGrammarGrammarBackend(BaseGrammarBackend):
self.vocab_size = vocab_size
self.override_stop_tokens = override_stop_tokens
def _from_context(self, ctx: CompiledGrammar, key_string: str) -> XGrammarGrammar:
def _from_context(
self, ctx: CompiledGrammar, key_string: str, grammar_stats: GrammarStats
) -> XGrammarGrammar:
matcher = GrammarMatcher(
ctx,
max_rollback_tokens=MAX_ROLLBACK_TOKENS,
override_stop_tokens=self.override_stop_tokens,
)
return XGrammarGrammar(
matcher, self.vocab_size, ctx, self.override_stop_tokens, key_string
matcher,
self.vocab_size,
ctx,
self.override_stop_tokens,
key_string,
grammar_stats,
)
def dispatch_json(self, key_string: str) -> Optional[XGrammarGrammar]:
@@ -198,7 +213,7 @@ class XGrammarGrammarBackend(BaseGrammarBackend):
except (RuntimeError, json.decoder.JSONDecodeError) as e:
logging.error(f"Hit invalid json_schema: {key_string=}, {e=}")
return INVALID_GRAMMAR_OBJ
return self._from_context(ctx, key_string)
return self._from_context(ctx, key_string, GrammarStats())
def dispatch_ebnf(self, key_string: str) -> Optional[XGrammarGrammar]:
try:
@@ -206,7 +221,7 @@ class XGrammarGrammarBackend(BaseGrammarBackend):
except RuntimeError as e:
logging.error(f"Hit invalid ebnf: {key_string=}, {e=}")
return INVALID_GRAMMAR_OBJ
return self._from_context(ctx, key_string)
return self._from_context(ctx, key_string, GrammarStats())
def dispatch_regex(self, key_string: str) -> Optional[XGrammarGrammar]:
try:
@@ -214,7 +229,7 @@ class XGrammarGrammarBackend(BaseGrammarBackend):
except RuntimeError as e:
logging.error(f"Hit invalid regex: {key_string=}, {e=}")
return INVALID_GRAMMAR_OBJ
return self._from_context(ctx, key_string)
return self._from_context(ctx, key_string, GrammarStats())
def dispatch_structural_tag(self, key_string: str) -> Optional[XGrammarGrammar]:
try:
@@ -233,7 +248,7 @@ class XGrammarGrammarBackend(BaseGrammarBackend):
except (RuntimeError, json.decoder.JSONDecodeError) as e:
logging.error(f"Hit invalid structural_tag: {key_string=}, {e=}")
return INVALID_GRAMMAR_OBJ
return self._from_context(ctx, key_string)
return self._from_context(ctx, key_string, GrammarStats())
def reset(self):
self.grammar_compiler.clear_cache()