accelerator option is added

This commit is contained in:
alpertunga-bile
2023-08-12 14:12:14 +03:00
parent 2f3320c6d0
commit ff56ccec2d
11 changed files with 223 additions and 203 deletions
+40 -40
View File
@@ -17,6 +17,7 @@ class PromptGenerator:
if isdir(join(join("models", "prompt_generators"), file))
],
),
"accelerator": (["enable", "disable"],),
"prompt": (
"STRING",
{
@@ -109,7 +110,7 @@ class PromptGenerator:
) -> None:
from datetime import datetime
print_string = "{' PROMPT GENERATOR OUTPUT '.center(200, '#')}\n"
print_string = f"{' PROMPT GENERATOR OUTPUT '.center(200, '#')}\n"
print_string += f"{generated_text}\n"
print_string += f"{'#'*200}\n"
@@ -144,6 +145,7 @@ class PromptGenerator:
self,
clip,
model_name,
accelerator,
prompt,
cfg,
min_length,
@@ -174,49 +176,47 @@ class PromptGenerator:
file.close()
if exists(real_path) is False:
print(f"{real_path} is not exists")
generated_text = prompt
else:
is_self_recursive = True if self_recursive == "enable" else False
raise ValueError(f"{real_path} is not exists")
generator = Generator(real_path)
is_self_recursive = True if self_recursive == "enable" else False
is_accelerate = True if accelerator == "enable" else False
gen_settings = GenerateArgs(
guidance_scale=cfg,
min_length=min_length,
max_length=max_length,
do_sample=True if do_sample == "enable" else False,
early_stopping=True if early_stopping == "enable" else False,
num_beams=num_beams,
num_beam_groups=num_beam_groups,
temperature=temperature,
top_k=top_k,
top_p=top_p,
repetition_penalty=repetition_penalty,
no_repeat_ngram_size=no_repeat_ngram_size,
remove_invalid_values=True
if remove_invalid_values == "enable"
else False,
)
generator = Generator(real_path, is_accelerate)
generated_text = self.get_generated_text(
generator,
gen_settings,
prompt,
is_self_recursive,
recursive_level,
preprocess_mode,
)
gen_settings = GenerateArgs(
guidance_scale=cfg,
min_length=min_length,
max_length=max_length,
do_sample=True if do_sample == "enable" else False,
early_stopping=True if early_stopping == "enable" else False,
num_beams=num_beams,
num_beam_groups=num_beam_groups,
temperature=temperature,
top_k=top_k,
top_p=top_p,
repetition_penalty=repetition_penalty,
no_repeat_ngram_size=no_repeat_ngram_size,
remove_invalid_values=True if remove_invalid_values == "enable" else False,
)
self.log_outputs(
prompt,
generated_text,
self_recursive,
recursive_level,
preprocess_mode,
gen_settings,
prompt_log_filename,
)
generated_text = self.get_generated_text(
generator,
gen_settings,
prompt,
is_self_recursive,
recursive_level,
preprocess_mode,
)
self.log_outputs(
prompt,
generated_text,
self_recursive,
recursive_level,
preprocess_mode,
gen_settings,
prompt_log_filename,
)
tokens = clip.tokenize(generated_text)
cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True)