[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
This commit is contained in:
parent
9e89d5d3d7
commit
ad7194cac6
1 changed files with 17 additions and 11 deletions
|
|
@ -54,6 +54,7 @@ ws ::= ([
|
|||
] ws)?
|
||||
'''
|
||||
|
||||
|
||||
def generate_with_grammar(
|
||||
model,
|
||||
tokenizer,
|
||||
|
|
@ -67,12 +68,12 @@ def generate_with_grammar(
|
|||
do_sample: bool = False,
|
||||
repetition_penalty: float = 1.1,
|
||||
num_return_sequences: int = 1,
|
||||
**kwargs
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Generate text with grammar constraints using transformers-cfg.
|
||||
Automatically handles model-specific parameter compatibility.
|
||||
|
||||
|
||||
Args:
|
||||
model: The model to generate with.
|
||||
tokenizer: The tokenizer associated with the model.
|
||||
|
|
@ -90,7 +91,9 @@ def generate_with_grammar(
|
|||
"""
|
||||
try:
|
||||
from transformers_cfg.grammar_utils import IncrementalGrammarConstraint
|
||||
from transformers_cfg.generation.logits_process import GrammarConstrainedLogitsProcessor
|
||||
from transformers_cfg.generation.logits_process import (
|
||||
GrammarConstrainedLogitsProcessor,
|
||||
)
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"Unsloth: Please install transformers-cfg to use grammar-constrained generation: "
|
||||
|
|
@ -103,9 +106,7 @@ def generate_with_grammar(
|
|||
|
||||
# Create grammar constraint
|
||||
grammar = IncrementalGrammarConstraint(
|
||||
grammar_str,
|
||||
start_rule_name=start_rule,
|
||||
tokenizer=tokenizer
|
||||
grammar_str, start_rule_name = start_rule, tokenizer = tokenizer
|
||||
)
|
||||
grammar_processor = GrammarConstrainedLogitsProcessor(grammar)
|
||||
|
||||
|
|
@ -117,14 +118,17 @@ def generate_with_grammar(
|
|||
"repetition_penalty": repetition_penalty,
|
||||
"num_return_sequences": num_return_sequences,
|
||||
"logits_processor": [grammar_processor],
|
||||
**kwargs
|
||||
**kwargs,
|
||||
}
|
||||
|
||||
# Handle sampling parameters
|
||||
if do_sample:
|
||||
if temperature is not None: generation_kwargs["temperature"] = temperature
|
||||
if top_p is not None: generation_kwargs["top_p"] = top_p
|
||||
if top_k is not None: generation_kwargs["top_k"] = top_k
|
||||
if temperature is not None:
|
||||
generation_kwargs["temperature"] = temperature
|
||||
if top_p is not None:
|
||||
generation_kwargs["top_p"] = top_p
|
||||
if top_k is not None:
|
||||
generation_kwargs["top_k"] = top_k
|
||||
else:
|
||||
generation_kwargs["temperature"] = None
|
||||
generation_kwargs["top_p"] = None
|
||||
|
|
@ -136,7 +140,9 @@ def generate_with_grammar(
|
|||
except ValueError as e:
|
||||
# Some models fail if specific kwargs like 'sliding_window' or 'num_logits_to_keep' are passed
|
||||
if "model_kwargs" in str(e):
|
||||
logger.warning("Unsloth: Generation failed with model_kwargs error. Retrying with minimal parameters...")
|
||||
logger.warning(
|
||||
"Unsloth: Generation failed with model_kwargs error. Retrying with minimal parameters..."
|
||||
)
|
||||
minimal_kwargs = {
|
||||
"input_ids": input_ids,
|
||||
"max_new_tokens": max_new_tokens,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue