fix: propagate grammar errors and improve llguidance backend (#20467)

This commit is contained in:
jellysnack
2026-03-16 02:11:18 +03:00
committed by GitHub
parent 116aef8504
commit 53f831691a
6 changed files with 64 additions and 28 deletions

View File

@@ -116,7 +116,15 @@ class BaseGrammarObject:
raise NotImplementedError()
INVALID_GRAMMAR_OBJ = BaseGrammarObject()
class InvalidGrammarObject(BaseGrammarObject):
"""Represents a grammar that failed to compile, carrying the original error message."""
def __init__(self, error_message: str = "Unknown grammar error"):
super().__init__()
self.error_message = error_message
def __repr__(self):
return f"InvalidGrammarObject(error_message={self.error_message!r})"
class BaseGrammarBackend:
@@ -126,7 +134,7 @@ class BaseGrammarBackend:
def _not_supported(self, key_type: str, key_string: str) -> BaseGrammarObject:
logger.warning(f"Skip unsupported {key_type=}, {key_string=}")
return INVALID_GRAMMAR_OBJ
return InvalidGrammarObject()
def dispatch_fallback(self, key_type: str, key_string: str) -> BaseGrammarObject:
"""

View File

@@ -8,7 +8,7 @@ from typing import TYPE_CHECKING, List
import torch
from sglang.srt.constrained.base_grammar_backend import (
INVALID_GRAMMAR_OBJ,
InvalidGrammarObject,
create_grammar_backend,
)
from sglang.srt.environ import envs
@@ -95,8 +95,12 @@ class GrammarManager:
req.grammar_key = key
add_to_grammar_queue = True
else:
if value is INVALID_GRAMMAR_OBJ: # We hit a cached invalid grammar.
error_msg = f"Invalid grammar request with cache hit: {key=}"
if isinstance(
value, InvalidGrammarObject
): # We hit a cached invalid grammar.
error_msg = (
f"Failed to compile {key[0]} grammar: {value.error_message}"
)
req.set_finish_with_abort(error_msg)
if add_to_grammar_queue:
@@ -174,8 +178,8 @@ class GrammarManager:
assert isinstance(req.grammar, futures.Future) and req.grammar_key
req.grammar = req.grammar.result()
self.grammar_backend.set_cache(req.grammar_key, req.grammar.copy())
if req.grammar is INVALID_GRAMMAR_OBJ:
error_msg = f"Invalid grammar request: {req.grammar_key=}"
if isinstance(req.grammar, InvalidGrammarObject):
error_msg = f"Failed to compile {req.grammar_key[0]} grammar: {req.grammar.error_message}"
req.set_finish_with_abort(error_msg)
# Return failed requests
@@ -185,7 +189,9 @@ class GrammarManager:
assert isinstance(req.grammar, futures.Future) and req.grammar_key
req.grammar.cancel()
self.grammar_backend.set_cache(req.grammar_key, INVALID_GRAMMAR_OBJ)
self.grammar_backend.set_cache(
req.grammar_key, InvalidGrammarObject("Grammar preprocessing timed out")
)
error_msg = f"Grammar preprocessing timed out: {req.grammar_key=}"
req.set_finish_with_abort(error_msg)

View File

@@ -28,9 +28,9 @@ from llguidance.torch import (
)
from sglang.srt.constrained.base_grammar_backend import (
INVALID_GRAMMAR_OBJ,
BaseGrammarBackend,
BaseGrammarObject,
InvalidGrammarObject,
)
from sglang.srt.constrained.utils import is_legacy_structural_tag
@@ -49,18 +49,36 @@ class GuidanceGrammar(BaseGrammarObject):
self.serialized_grammar,
log_level=int(os.environ.get("LLGUIDANCE_LOG_LEVEL", "1")),
)
self._check_err()
self.bitmask = None
self.eos_token = self.llguidance_tokenizer.eos_token
def accept_token(self, token: int):
if not self.ll_matcher.consume_token(token):
logger.warning(f"matcher error: {self.ll_matcher.get_error()}")
if self.finished:
return
if self.ll_matcher.is_stopped() and token == self.eos_token:
self.finished = True
return
self.ll_matcher.consume_token(token)
self._check_err()
def rollback(self, num_tokens: int) -> None:
if num_tokens <= 0:
return
if self.finished:
self.finished = False
# EOS token after stop isn't tracked in ll_matcher
num_tokens -= 1
self.ll_matcher.rollback(num_tokens)
self._check_err()
def is_terminated(self):
return self.finished
def fill_vocab_mask(self, vocab_mask: torch.Tensor, idx: int) -> None:
if self.ll_matcher.is_stopped():
self.finished = True
fill_next_token_bitmask(self.ll_matcher, vocab_mask, idx)
self._check_err()
def allocate_vocab_mask(
self, vocab_size: int, batch_size: int, device
@@ -105,6 +123,10 @@ class GuidanceGrammar(BaseGrammarObject):
):
pass
def _check_err(self) -> None:
if self.ll_matcher.is_error():
raise ValueError(self.ll_matcher.get_error())
class GuidanceBackend(BaseGrammarBackend):
@@ -130,7 +152,7 @@ class GuidanceBackend(BaseGrammarBackend):
)
except Exception as e:
logger.error(f"Hit invalid grammar: {serialized_grammar=}, {e=}")
return INVALID_GRAMMAR_OBJ
return InvalidGrammarObject(str(e))
def dispatch_json(self, key_string: str) -> BaseGrammarObject:
try:
@@ -143,7 +165,7 @@ class GuidanceBackend(BaseGrammarBackend):
)
except Exception as e:
logger.error(f"Hit invalid json_schema: {key_string=}, {e=}")
return INVALID_GRAMMAR_OBJ
return InvalidGrammarObject(str(e))
return self._from_serialized(serialized_grammar)
def dispatch_regex(self, key_string: str) -> BaseGrammarObject:
@@ -156,7 +178,7 @@ class GuidanceBackend(BaseGrammarBackend):
return self._from_serialized(serialized_grammar)
except ValueError as e:
logger.error(f"Hit invalid ebnf: {key_string=}, {e=}")
return INVALID_GRAMMAR_OBJ
return InvalidGrammarObject(str(e))
def dispatch_structural_tag(self, key_string: str) -> BaseGrammarObject:
try:
@@ -175,4 +197,4 @@ class GuidanceBackend(BaseGrammarBackend):
return self._from_serialized(g)
except Exception as e:
logger.error(f"Hit invalid structural_tag: {key_string=}, {e=}")
return INVALID_GRAMMAR_OBJ
return InvalidGrammarObject(str(e))

View File

@@ -24,9 +24,9 @@ from outlines.models.transformers import TransformerTokenizer
from pydantic import BaseModel
from sglang.srt.constrained.base_grammar_backend import (
INVALID_GRAMMAR_OBJ,
BaseGrammarBackend,
BaseGrammarObject,
InvalidGrammarObject,
)
from sglang.srt.constrained.outlines_jump_forward import OutlinesJumpForwardMap
@@ -152,7 +152,7 @@ class OutlinesGrammarBackend(BaseGrammarBackend):
guide = RegexGuide(regex, self.outlines_tokenizer)
except interegular.patterns.InvalidSyntax as e:
logger.error(f"Hit invalid regex schema: {regex=}, {e=}")
return INVALID_GRAMMAR_OBJ
return InvalidGrammarObject(str(e))
jump_forward_map = None
return OutlinesGrammar(guide, jump_forward_map)
@@ -171,7 +171,7 @@ class OutlinesGrammarBackend(BaseGrammarBackend):
)
except (NotImplementedError, json.decoder.JSONDecodeError, ValueError) as e:
logger.error(f"Hit invalid json_schema: {key_string=}, {e=}")
return INVALID_GRAMMAR_OBJ
return InvalidGrammarObject(str(e))
return self._compile_regex(regex)
def dispatch_regex(self, key_string: str):

View File

@@ -18,9 +18,9 @@ from typing import List, Optional, Tuple
import torch
from .base_grammar_backend import (
INVALID_GRAMMAR_OBJ,
BaseGrammarBackend,
BaseGrammarObject,
InvalidGrammarObject,
)
@@ -117,7 +117,7 @@ class ReasonerGrammarBackend(BaseGrammarBackend):
) -> Optional[BaseGrammarObject]:
ret = self.grammar_backend._init_value_dispatch(key, reasoning)
# avoid wrapping invalid grammar, so that the scheduler can detect it
if ret is None or ret is INVALID_GRAMMAR_OBJ:
if ret is None or isinstance(ret, InvalidGrammarObject):
return ret
obj = ReasonerGrammarObject(ret, self.think_end_id)
obj.maybe_init_reasoning(reasoning)

View File

@@ -29,10 +29,10 @@ from xgrammar import (
)
from sglang.srt.constrained.base_grammar_backend import (
INVALID_GRAMMAR_OBJ,
BaseGrammarBackend,
BaseGrammarObject,
GrammarStats,
InvalidGrammarObject,
)
from sglang.srt.constrained.utils import is_legacy_structural_tag
from sglang.srt.utils import is_hip
@@ -266,7 +266,7 @@ class XGrammarGrammarBackend(BaseGrammarBackend):
except (RuntimeError, json.decoder.JSONDecodeError, UnicodeDecodeError) as e:
logger.error(f"Hit invalid json_schema: {key_string=}, {e=}")
return INVALID_GRAMMAR_OBJ
return InvalidGrammarObject(str(e))
return self._from_context(ctx, key_string, GrammarStats(dispatch_type="json"))
def dispatch_ebnf(self, key_string: str) -> BaseGrammarObject:
@@ -274,7 +274,7 @@ class XGrammarGrammarBackend(BaseGrammarBackend):
ctx = self.grammar_compiler.compile_grammar(key_string)
except RuntimeError as e:
logger.error(f"Hit invalid ebnf: {key_string=}, {e=}")
return INVALID_GRAMMAR_OBJ
return InvalidGrammarObject(str(e))
return self._from_context(ctx, key_string, GrammarStats(dispatch_type="ebnf"))
def dispatch_regex(self, key_string: str) -> BaseGrammarObject:
@@ -282,7 +282,7 @@ class XGrammarGrammarBackend(BaseGrammarBackend):
ctx = self.grammar_compiler.compile_regex(key_string)
except RuntimeError as e:
logger.error(f"Hit invalid regex: {key_string=}, {e=}")
return INVALID_GRAMMAR_OBJ
return InvalidGrammarObject(str(e))
return self._from_context(ctx, key_string, GrammarStats(dispatch_type="regex"))
def dispatch_structural_tag(self, key_string: str) -> BaseGrammarObject:
@@ -311,7 +311,7 @@ class XGrammarGrammarBackend(BaseGrammarBackend):
ctx = self.grammar_compiler.compile_structural_tag(key_string)
except (RuntimeError, json.decoder.JSONDecodeError) as e:
logger.error(f"Hit invalid structural_tag: {key_string=}, {e=}")
return INVALID_GRAMMAR_OBJ
return InvalidGrammarObject(str(e))
return self._from_context(
ctx, key_string, GrammarStats(dispatch_type="structural_tag")
)