From 53f831691a245fe857f599325ea2b4ae45d490c1 Mon Sep 17 00:00:00 2001 From: jellysnack <158609015+jellysnack@users.noreply.github.com> Date: Mon, 16 Mar 2026 02:11:18 +0300 Subject: [PATCH] fix: propagate grammar errors and improve llguidance backend (#20467) --- .../srt/constrained/base_grammar_backend.py | 12 +++++- .../sglang/srt/constrained/grammar_manager.py | 18 +++++--- .../srt/constrained/llguidance_backend.py | 42 ++++++++++++++----- .../srt/constrained/outlines_backend.py | 6 +-- .../constrained/reasoner_grammar_backend.py | 4 +- .../srt/constrained/xgrammar_backend.py | 10 ++--- 6 files changed, 64 insertions(+), 28 deletions(-) diff --git a/python/sglang/srt/constrained/base_grammar_backend.py b/python/sglang/srt/constrained/base_grammar_backend.py index 4f9dc0fb3..e0738db9c 100644 --- a/python/sglang/srt/constrained/base_grammar_backend.py +++ b/python/sglang/srt/constrained/base_grammar_backend.py @@ -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: """ diff --git a/python/sglang/srt/constrained/grammar_manager.py b/python/sglang/srt/constrained/grammar_manager.py index a1ea6d73f..829675ec5 100644 --- a/python/sglang/srt/constrained/grammar_manager.py +++ b/python/sglang/srt/constrained/grammar_manager.py @@ -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) diff --git a/python/sglang/srt/constrained/llguidance_backend.py b/python/sglang/srt/constrained/llguidance_backend.py index 6029fb909..b3d32301c 100644 --- a/python/sglang/srt/constrained/llguidance_backend.py +++ b/python/sglang/srt/constrained/llguidance_backend.py @@ -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)) diff --git a/python/sglang/srt/constrained/outlines_backend.py b/python/sglang/srt/constrained/outlines_backend.py index eaed4bafe..881749633 100644 --- a/python/sglang/srt/constrained/outlines_backend.py +++ b/python/sglang/srt/constrained/outlines_backend.py @@ -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): diff --git a/python/sglang/srt/constrained/reasoner_grammar_backend.py b/python/sglang/srt/constrained/reasoner_grammar_backend.py index e2ae8405e..d204bdd9e 100644 --- a/python/sglang/srt/constrained/reasoner_grammar_backend.py +++ b/python/sglang/srt/constrained/reasoner_grammar_backend.py @@ -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) diff --git a/python/sglang/srt/constrained/xgrammar_backend.py b/python/sglang/srt/constrained/xgrammar_backend.py index 56b053fa0..0012229d2 100644 --- a/python/sglang/srt/constrained/xgrammar_backend.py +++ b/python/sglang/srt/constrained/xgrammar_backend.py @@ -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") )