fix: propagate grammar errors and improve llguidance backend (#20467)
This commit is contained in:
@@ -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:
|
||||
"""
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user