[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:
co-authored by
github-actions[bot] <github-actions[bot]@users.noreply.github.com>
parent
151e287d1a
commit
2269cf1e2f
@@ -14,8 +14,9 @@
|
||||
"""The baseclass of a backend for grammar-guided constrained decoding."""
|
||||
|
||||
import logging
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, field
|
||||
from threading import Event
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
|
||||
@@ -26,10 +27,22 @@ from sglang.srt.server_args import ServerArgs
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class GrammarStats:
|
||||
compilation_time: Optional[float] = None
|
||||
schema_count: Optional[int] = None
|
||||
ebnf_size: Optional[int] = None
|
||||
is_cache_hit: bool = False
|
||||
is_grammar_aborted: bool = False
|
||||
tree_traversal_time: List[float] = field(default_factory=list)
|
||||
|
||||
|
||||
class BaseGrammarObject:
|
||||
|
||||
def __init__(self):
|
||||
self._finished = False
|
||||
self.grammar_stats = None
|
||||
self.current_token = None
|
||||
|
||||
def accept_token(self, token: int) -> None:
|
||||
"""
|
||||
@@ -137,19 +150,26 @@ class BaseGrammarBackend:
|
||||
return self._not_supported("structural_tag", key_string)
|
||||
|
||||
def _init_value_dispatch(self, key: Tuple[str, str]) -> Optional[BaseGrammarObject]:
|
||||
s = time.perf_counter()
|
||||
key_type, key_string = key
|
||||
if key_type == "json":
|
||||
return self.dispatch_json(key_string)
|
||||
grammar = self.dispatch_json(key_string)
|
||||
elif key_type == "regex":
|
||||
return self.dispatch_regex(key_string)
|
||||
grammar = self.dispatch_regex(key_string)
|
||||
elif key_type == "ebnf":
|
||||
return self.dispatch_ebnf(key_string)
|
||||
grammar = self.dispatch_ebnf(key_string)
|
||||
elif key_type == "structural_tag":
|
||||
return self.dispatch_structural_tag(key_string)
|
||||
grammar = self.dispatch_structural_tag(key_string)
|
||||
elif key_type == "structural_pattern":
|
||||
return self.dispatch_structural_pattern(key_string)
|
||||
grammar = self.dispatch_structural_pattern(key_string)
|
||||
elif key_type == "structural_pattern_v2":
|
||||
grammar = self.dispatch_structural_pattern_v2(key_string)
|
||||
else:
|
||||
return self.dispatch_fallback(key_type, key_string)
|
||||
grammar = self.dispatch_fallback(key_type, key_string)
|
||||
|
||||
if grammar is not None and grammar.grammar_stats is not None:
|
||||
grammar.grammar_stats.compilation_time = time.perf_counter() - s
|
||||
return grammar
|
||||
|
||||
def get_cached_or_future_value(
|
||||
self, key: Tuple[str, str]
|
||||
@@ -167,20 +187,36 @@ class BaseGrammarBackend:
|
||||
self.cache.clear()
|
||||
|
||||
|
||||
GRAMMAR_BACKEND_REGISTRY = {}
|
||||
|
||||
|
||||
def register_grammar_backend(name, init_func):
|
||||
GRAMMAR_BACKEND_REGISTRY[name] = init_func
|
||||
|
||||
|
||||
def create_grammar_backend(
|
||||
server_args: ServerArgs,
|
||||
tokenizer,
|
||||
vocab_size: int,
|
||||
eos_token_ids: Optional[set] = None,
|
||||
) -> Optional[BaseGrammarBackend]:
|
||||
if server_args.grammar_backend == "outlines":
|
||||
name = server_args.grammar_backend
|
||||
|
||||
# Custom grammar backend has the highest priority
|
||||
if name in GRAMMAR_BACKEND_REGISTRY:
|
||||
return GRAMMAR_BACKEND_REGISTRY[name](
|
||||
server_args, tokenizer, vocab_size, eos_token_ids
|
||||
)
|
||||
|
||||
# Default grammar backends
|
||||
if name == "outlines":
|
||||
from sglang.srt.constrained.outlines_backend import OutlinesGrammarBackend
|
||||
|
||||
grammar_backend = OutlinesGrammarBackend(
|
||||
tokenizer,
|
||||
whitespace_pattern=server_args.constrained_json_whitespace_pattern,
|
||||
)
|
||||
elif server_args.grammar_backend == "xgrammar":
|
||||
elif name == "xgrammar":
|
||||
from sglang.srt.constrained.xgrammar_backend import XGrammarGrammarBackend
|
||||
|
||||
# Convert Set[int] to List[int] if needed
|
||||
@@ -189,17 +225,17 @@ def create_grammar_backend(
|
||||
grammar_backend = XGrammarGrammarBackend(
|
||||
tokenizer, vocab_size=vocab_size, model_eos_token_ids=eos_list
|
||||
)
|
||||
elif server_args.grammar_backend == "llguidance":
|
||||
elif name == "llguidance":
|
||||
from sglang.srt.constrained.llguidance_backend import GuidanceBackend
|
||||
|
||||
grammar_backend = GuidanceBackend(
|
||||
tokenizer=tokenizer,
|
||||
whitespace_pattern=server_args.constrained_json_whitespace_pattern,
|
||||
)
|
||||
elif server_args.grammar_backend == "none":
|
||||
elif name == "none":
|
||||
return None
|
||||
else:
|
||||
raise ValueError(f"Invalid grammar backend: {server_args.grammar_backend}")
|
||||
raise ValueError(f"Invalid grammar backend: {name}")
|
||||
|
||||
if server_args.reasoning_parser and hasattr(tokenizer, "think_end_id"):
|
||||
from sglang.srt.constrained.reasoner_grammar_backend import (
|
||||
|
||||
Reference in New Issue
Block a user