Skip to content

Commit b255847

Browse files
authored
Merge pull request #105 from DocShotgun/main
Add support for regex pattern constraints
2 parents 5752521 + abe411c commit b255847

3 files changed

Lines changed: 55 additions & 0 deletions

File tree

backends/exllamav2/grammar.py

Lines changed: 43 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -97,6 +97,49 @@ def add_json_schema_filter(
9797
gen_settings.filters.extend([lmfilter, prefix_filter])
9898
gen_settings.filter_prefer_eos = True
9999

100+
def add_regex_filter(
101+
self,
102+
pattern: str,
103+
gen_settings: ExLlamaV2Sampler.Settings,
104+
tokenizer: ExLlamaV2Tokenizer,
105+
):
106+
"""Adds an ExllamaV2 filter based on regular expressions."""
107+
108+
# Import optional dependencies
109+
try:
110+
from lmformatenforcer import RegexParser
111+
from lmformatenforcer.integrations.exllamav2 import (
112+
ExLlamaV2TokenEnforcerFilter,
113+
)
114+
except ImportError:
115+
logger.error(
116+
"Skipping regex parsing because "
117+
"lm-format-enforcer is not installed.\n"
118+
"Please run the following command in your environment "
119+
"to reinstall dependencies:\n"
120+
"pip install -U ."
121+
)
122+
123+
return
124+
125+
# Create the parser
126+
try:
127+
pattern_parser = RegexParser(pattern)
128+
except Exception:
129+
traceback.print_exc()
130+
logger.error(
131+
"Skipping because the regex pattern couldn't be parsed. "
132+
"Please read the above error for more information."
133+
)
134+
135+
return
136+
137+
lmfilter = ExLlamaV2TokenEnforcerFilter(pattern_parser, tokenizer)
138+
139+
# Append the filters
140+
gen_settings.filters.extend([lmfilter])
141+
gen_settings.filter_prefer_eos = True
142+
100143
def add_ebnf_filter(
101144
self,
102145
ebnf_string: str,

backends/exllamav2/model.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -850,6 +850,13 @@ def generate_gen_sync(
850850
json_schema, gen_settings, self.model, self.tokenizer
851851
)
852852

853+
# Add regex filter if it exists
854+
regex_pattern = unwrap(kwargs.get("regex_pattern"))
855+
if regex_pattern:
856+
grammar_handler.add_regex_filter(
857+
regex_pattern, gen_settings, self.tokenizer
858+
)
859+
853860
# Add EBNF filter if it exists
854861
grammar_string = unwrap(kwargs.get("grammar_string"))
855862
if grammar_string:

common/sampling.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -138,6 +138,10 @@ class BaseSamplerRequest(BaseModel):
138138
default_factory=lambda: get_default_sampler_value("json_schema"),
139139
)
140140

141+
regex_pattern: Optional[str] = Field(
142+
default_factory=lambda: get_default_sampler_value("regex_pattern"),
143+
)
144+
141145
grammar_string: Optional[str] = Field(
142146
default_factory=lambda: get_default_sampler_value("grammar_string"),
143147
)
@@ -312,6 +316,7 @@ def to_gen_params(self, **kwargs):
312316
"cfg_scale": self.cfg_scale,
313317
"negative_prompt": self.negative_prompt,
314318
"json_schema": self.json_schema,
319+
"regex_pattern": self.regex_pattern,
315320
"grammar_string": self.grammar_string,
316321
"speculative_ngram": self.speculative_ngram,
317322
}

0 commit comments

Comments
 (0)