token_healing option is added

This commit is contained in:
alpertunga-bile
2025-01-12 18:24:14 +03:00
parent 4250f2e850
commit 4984cf73ff
2 changed files with 15 additions and 6 deletions
+10 -5
View File
@@ -5,6 +5,7 @@ from generator.utility import (
get_accelerator_type,
get_variable_dictionary,
str_to_quant_type,
check_transformers_version,
ModelType,
)
from generator.preprocess import preprocess
@@ -39,7 +40,11 @@ class Generator:
extra_params = {}
def __init__(
self, model_path: str, is_accelerate: bool, model_quant_type: str
self,
model_path: str,
is_accelerate: bool,
is_token_healing: bool,
model_quant_type: str,
) -> None:
quantize_type = str_to_quant_type(model_quant_type)
@@ -60,13 +65,13 @@ class Generator:
self.extra_params["pad_token_id"] = self.tokenizer.eos_token_id
"""
# token healing feature is breaking the generator
# in experiments so this is disabled for now
token healing feature is breaking the generator
in experiments so this is disabled for now
"""
if check_transformers_version(4, 6):
if check_transformers_version(4, 6) and is_token_healing is True:
self.extra_params["token_healing"] = True
self.extra_params["tokenizer"] = self.tokenizer
"""
self.dev = get_torch_device()
+5 -1
View File
@@ -38,6 +38,7 @@ class PromptGenerator:
"model_name": (model_names,),
"accelerate": (["enable", "disable"],),
"quantize": (quantize_sizes,),
"token_healing": (["disable", "enable"],),
"prompt": (
"STRING",
{
@@ -196,6 +197,7 @@ class PromptGenerator:
model_name: str,
accelerate: str,
quantize: str,
token_healing: str,
prompt: str,
seed: int,
lock: str,
@@ -263,6 +265,8 @@ class PromptGenerator:
is_self_recursive = True if self_recursive == "enable" else False
is_accelerate = True if accelerate == "enable" else False
is_token_healing = True if token_healing == "enable" else False
is_early_stopping = True if early_stopping == "enable" else False
is_remove_invalid_values = True if remove_invalid_values == "enable" else False
@@ -275,7 +279,7 @@ class PromptGenerator:
file = open(prompt_log_filename, "w")
file.close()
generator = Generator(model_path, is_accelerate, quantize)
generator = Generator(model_path, is_accelerate, is_token_healing, quantize)
self._gen_settings = GenerateArgs(
guidance_scale=cfg,