Skip to content

Feature Request: allow same GrammarConstrainedLogitsProcessor to be reused across multiple generations #49

Description

@Saibo-creator

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.

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions