Reproduce
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
from transformers_cfg.grammar_utils import IncrementalGrammarConstraint
from transformers_cfg.generation.logits_process import GrammarConstrainedLogitsProcessor
if __name__ == "__main__":
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Using device: {device}")
model_id = "mistralai/Mistral-7B-v0.1"
# Load model and tokenizer
tokenizer = AutoTokenizer.from_pretrained(model_id)
tokenizer.pad_token = tokenizer.eos_token
model = AutoModelForCausalLM.from_pretrained(model_id, load_in_4bit=True, device_map="auto")
model.generation_config.pad_token_id = model.generation_config.eos_token_id
grammar_str = """
# Grammar for subset of JSON
# String doesn't support unicode and escape yet
# If you don't need to generate unicode and escape, you can use this grammar
# We are working to support unicode and escape
root ::= object
object ::= "{" ws ( string ":" ws value ("," ws string ":" ws value)* )? "}"
value ::= object | array | string | number | ("true" | "false" | "null") ws
array ::= "[" ws ( value ("," ws value)* )? "]" ws
string ::= "\"" [ \t!#-\[\]-~]* "\"" ws
number ::= ("-"? ([0-9] | [1-9] [0-9]*)) ("." [0-9]+)? ([eE] [-+]? [0-9]+)? ws
ws ::= ([ \t\n] ws)?
"""
grammar = IncrementalGrammarConstraint(grammar_str, "root", tokenizer)
grammar_processor = GrammarConstrainedLogitsProcessor(grammar)
# Generate
prefix1 = "This is a valid json string for http request:"
prefix2 = "This is a valid json string for shopping cart:"
for prefix in [prefix1, prefix2]:
input_ids = tokenizer(
[prefix], add_special_tokens=False, return_tensors="pt", padding=True
)["input_ids"]
output = model.generate(
input_ids,
do_sample=False,
max_new_tokens=60,
logits_processor=[grammar_processor],
repetition_penalty=1.1,
num_return_sequences=1,
)
# decode output
generations = tokenizer.batch_decode(output, skip_special_tokens=True)
print(generations)
"""
'This is a valid json string for http request:{ "request": { "method": "GET", "headers": [], "content": "Content","type": "application" }}
'This is a valid json string for shopping cart:This is a valid json string for shopping cart:{ "name": "MyCart", "price": 0, "value": 1 }
"""
Error message
Traceback (most recent call last):
File "/home/saibo/Dev/SGCD-new/scripts/reproduce_tcfg_bug1.py", line 54, in <module>
output = model.generate(
File "/home/saibo/.virtualenvs/python3.10/SGCD-new/lib/python3.10/site-packages/torch/utils/_contextlib.py", line 115, in decorate_context
return func(*args, **kwargs)
File "/home/saibo/.virtualenvs/python3.10/SGCD-new/lib/python3.10/site-packages/transformers/generation/utils.py", line 1736, in generate
result = self._sample(
File "/home/saibo/.virtualenvs/python3.10/SGCD-new/lib/python3.10/site-packages/transformers/generation/utils.py", line 2388, in _sample
next_token_scores = logits_processor(input_ids, next_token_logits)
File "/home/saibo/.virtualenvs/python3.10/SGCD-new/lib/python3.10/site-packages/transformers/generation/logits_process.py", line 98, in __call__
scores = processor(input_ids, scores)
File "/home/saibo/.virtualenvs/python3.10/SGCD-new/lib/python3.10/site-packages/transformers_cfg/generation/logits_process.py", line 106, in __call__
return self.process_logits(input_ids, scores)
File "/home/saibo/.virtualenvs/python3.10/SGCD-new/lib/python3.10/site-packages/transformers_cfg/generation/logits_process.py", line 93, in process_logits
self.batch_accept_states = self.grammar_constraint.consume_token_ids(
File "/home/saibo/.virtualenvs/python3.10/SGCD-new/lib/python3.10/site-packages/transformers_cfg/token_grammar_recognizer.py", line 211, in consume_token_ids
raise RuntimeError(
RuntimeError: Input ID's length is inconsistent with the current state of the GrammarConstrainedLogitsProcessor. If you want to process another input sequence, please instantiate a new GrammarConstrainedLogitsProcessor.
Reproduce
Error message