token healing is added

This commit is contained in:
alpertunga-bile
2024-11-02 15:04:52 +03:00
parent d92d9dc967
commit a7b97943f3
2 changed files with 23 additions and 11 deletions
+22 -10
View File
@@ -1,5 +1,3 @@
from dataclasses import dataclass
from comfy.model_management import get_torch_device from comfy.model_management import get_torch_device
from generator.model import get_model_tokenizer from generator.model import get_model_tokenizer
@@ -7,10 +5,12 @@ from generator.utility import (
get_accelerator_type, get_accelerator_type,
get_variable_dictionary, get_variable_dictionary,
str_to_quant_type, str_to_quant_type,
check_transformers_version,
ModelType,
) )
from generator.preprocess import preprocess from generator.preprocess import preprocess
from .utility import ModelType from dataclasses import dataclass
@dataclass @dataclass
@@ -37,6 +37,7 @@ class Generator:
model = None model = None
tokenizer = None tokenizer = None
dev = None dev = None
extra_params = {}
def __init__( def __init__(
self, model_path: str, is_accelerate: bool, model_quant_type: str self, model_path: str, is_accelerate: bool, model_quant_type: str
@@ -54,6 +55,15 @@ class Generator:
model_path, accelerator_type, quantize_type model_path, accelerator_type, quantize_type
) )
self.tokenizer.padding_side = "left"
self.extra_params["renormalize_logits"] = True
self.extra_params["pad_token_id"] = self.tokenizer.eos_token_id
if check_transformers_version(4, 6):
self.extra_params["token_healing"] = True
self.extra_params["tokenizer"] = self.tokenizer
self.dev = get_torch_device() self.dev = get_torch_device()
# generate single output # generate single output
@@ -69,12 +79,13 @@ class Generator:
args.num_return_sequences = 1 args.num_return_sequences = 1
inputs = self.tokenizer(input, padding=True, return_tensors="pt").to(self.dev) inputs = self.tokenizer(
input.strip(), padding=True, truncation=True, return_tensors="pt"
).to(self.dev)
generated_ids = self.model.generate( generated_ids = self.model.generate(
**inputs, **inputs,
**get_variable_dictionary(args), **get_variable_dictionary(args),
pad_token_id=self.tokenizer.eos_token_id, **self.extra_params,
renormalize_logits=True,
) )
output = self.tokenizer.decode( output = self.tokenizer.decode(
@@ -96,15 +107,16 @@ class Generator:
args.num_return_sequences = 5 args.num_return_sequences = 5
inputs = self.tokenizer(input, padding=True, return_tensors="pt").to(self.dev) inputs = self.tokenizer(
input.strip(), padding=True, truncation=True, return_tensors="pt"
).to(self.dev)
generated_ids = self.model.generate( generated_ids = self.model.generate(
**inputs, **inputs,
**get_variable_dictionary(args), **get_variable_dictionary(args),
pad_token_id=self.tokenizer.eos_token_id, **self.extra_params,
renormalize_logits=True,
) )
outputs = self.tokenizer.batch_decode( outputs = self.tokenizer.batch_decode(
generated_ids, skip_special_tokens=True, cleanup_tokenization_spaces=True generated_ids, skip_special_tokens=True, clean_up_tokenization_spaces=True
) )
return outputs return outputs