added optimizations and repo is refactored
This commit is contained in:
+161
-153
@@ -1,5 +1,8 @@
|
||||
from os import listdir, mkdir
|
||||
from os import listdir
|
||||
from os.path import join, isdir, exists
|
||||
from preprocess import preprocess
|
||||
from generator.generate import GenerateArgs, Generator
|
||||
|
||||
|
||||
class PromptGenerator:
|
||||
@classmethod
|
||||
@@ -7,175 +10,163 @@ class PromptGenerator:
|
||||
return {
|
||||
"required": {
|
||||
"clip": ("CLIP",),
|
||||
"model_type": ("STRING", {
|
||||
"multiline" : False,
|
||||
"default" : "gpt2"
|
||||
}),
|
||||
"model_name":([file for file in listdir(join("models", "prompt_generators")) if isdir(join(join("models", "prompt_generators"), file))],),
|
||||
"seed": ("STRING", {
|
||||
"multiline" : True,
|
||||
"default" : "((masterpiece, best quality, ultra detailed)), illustration, digital art, 1girl, solo, ((stunningly beautiful))"
|
||||
}),
|
||||
"min_length": ("INT", {
|
||||
"default": 20,
|
||||
"min":0,
|
||||
"max":100,
|
||||
"step":1
|
||||
}),
|
||||
"max_length": ("INT", {
|
||||
"default": 50,
|
||||
"min":35,
|
||||
"max":200,
|
||||
"step":1
|
||||
}),
|
||||
"model_name": (
|
||||
[
|
||||
file
|
||||
for file in listdir(join("models", "prompt_generators"))
|
||||
if isdir(join(join("models", "prompt_generators"), file))
|
||||
],
|
||||
),
|
||||
"prompt": (
|
||||
"STRING",
|
||||
{
|
||||
"multiline": True,
|
||||
"default": "((masterpiece, best quality, ultra detailed)), illustration, digital art, 1girl, solo, ((stunningly beautiful))",
|
||||
},
|
||||
),
|
||||
"cfg": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.1},
|
||||
),
|
||||
"min_length": ("INT", {"default": 20, "min": 0, "max": 100, "step": 1}),
|
||||
"max_length": (
|
||||
"INT",
|
||||
{"default": 50, "min": 35, "max": 200, "step": 1},
|
||||
),
|
||||
"do_sample": (["disable", "enable"],),
|
||||
"early_stopping": (["disable", "enable"],),
|
||||
"num_beams": ("INT", {
|
||||
"default": 1,
|
||||
"min":1,
|
||||
"max":50,
|
||||
"step":1
|
||||
}),
|
||||
"temperature": ("FLOAT", {
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 2.0,
|
||||
"step": 0.1
|
||||
}),
|
||||
"top_k": ("INT", {
|
||||
"default": 50,
|
||||
"min":0,
|
||||
"max":150,
|
||||
"step":1
|
||||
}),
|
||||
"top_p": ("FLOAT", {
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 2.0,
|
||||
"step": 0.1
|
||||
}),
|
||||
"no_repeat_ngram_size": ("INT", {
|
||||
"default": 0,
|
||||
"min":0,
|
||||
"max":50,
|
||||
"step":1
|
||||
}),
|
||||
"num_beams": ("INT", {"default": 1, "min": 1, "max": 50, "step": 1}),
|
||||
"num_beam_groups": (
|
||||
"INT",
|
||||
{"default": 1, "min": 1, "max": 50, "step": 1},
|
||||
),
|
||||
"temperature": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.1},
|
||||
),
|
||||
"top_k": ("INT", {"default": 50, "min": 0, "max": 150, "step": 1}),
|
||||
"top_p": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.1},
|
||||
),
|
||||
"repetition_penalty": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 1.0, "max": 2.0, "step": 0.1},
|
||||
),
|
||||
"no_repeat_ngram_size": (
|
||||
"INT",
|
||||
{"default": 0, "min": 0, "max": 50, "step": 1},
|
||||
),
|
||||
"remove_invalid_values": (["disable", "enable"],),
|
||||
"self_recursive": (["disable", "enable"],),
|
||||
"recursive_level": ("INT", {
|
||||
"default": 0,
|
||||
"min":0,
|
||||
"max":50,
|
||||
"step":1
|
||||
}),
|
||||
"recursive_level": (
|
||||
"INT",
|
||||
{"default": 0, "min": 0, "max": 50, "step": 1},
|
||||
),
|
||||
"preprocess_mode": (["exact_keyword", "exact_prompt", "none"],),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING", )
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
FUNCTION = "generate"
|
||||
|
||||
CATEGORY = "Prompt Generator"
|
||||
|
||||
def GetUniqueList(self, sequence : list) -> list:
|
||||
seen = set()
|
||||
return [x for x in sequence if not (x in seen or seen.add(x))]
|
||||
|
||||
def RemoveDuplicates(self, line : str) -> list[str]:
|
||||
from string import punctuation
|
||||
|
||||
char_blacklist = set(f"{punctuation}0123456789")
|
||||
|
||||
# remove exact prompts
|
||||
prompts = self.GetUniqueList(line.split(","))
|
||||
pure_prompts = []
|
||||
|
||||
# remove exact keyword
|
||||
for prompt in prompts:
|
||||
can_add = True
|
||||
# extract the keyword
|
||||
keyword = "".join(c for c in prompt if c not in char_blacklist).lstrip()
|
||||
|
||||
if keyword == "":
|
||||
continue
|
||||
|
||||
for pure_prompt in pure_prompts:
|
||||
if prompt == pure_prompt:
|
||||
can_add = False
|
||||
break
|
||||
|
||||
extracted_pure_prompt = "".join(c for c in pure_prompt if c not in char_blacklist).lstrip()
|
||||
if keyword == extracted_pure_prompt:
|
||||
can_add = False
|
||||
break
|
||||
|
||||
if can_add:
|
||||
pure_prompts.append(prompt)
|
||||
|
||||
return self.GetUniqueList(pure_prompts)
|
||||
|
||||
def Preprocess(self, line : str, preprocess_mode : str) -> str:
|
||||
from re import sub, compile
|
||||
|
||||
pattern = compile(r'(,\s){2,}')
|
||||
|
||||
temp_line = line.replace(u'\xa0', u' ')
|
||||
temp_line = temp_line.replace("\n", ", ")
|
||||
temp_line = temp_line.replace("\t", " ")
|
||||
temp_line = temp_line.replace("|", ",")
|
||||
temp_line = temp_line.replace(" ", " ")
|
||||
temp_line = sub(pattern, ', ', temp_line)
|
||||
|
||||
if preprocess_mode == "exact_keyword":
|
||||
temp_line = ','.join(self.RemoveDuplicates(temp_line))
|
||||
elif preprocess_mode == "exact_prompt":
|
||||
temp_line = ','.join(self.GetUniqueList(temp_line.split(",")))
|
||||
|
||||
return temp_line
|
||||
|
||||
def GetGeneratedText(self, generator, gen_args, seed : str, is_self_recursive : bool, recursive_level : int, preprocess_mode : str) -> str:
|
||||
result = generator.generate_text(seed, gen_args)
|
||||
generated_text = self.Preprocess(seed + result.text, preprocess_mode)
|
||||
def get_generated_text(
|
||||
self,
|
||||
generator: Generator,
|
||||
gen_args: GenerateArgs,
|
||||
prompt: str,
|
||||
is_self_recursive: bool,
|
||||
recursive_level: int,
|
||||
preprocess_mode: str,
|
||||
) -> str:
|
||||
result = generator.generate_text(prompt, gen_args)
|
||||
generated_text = preprocess(prompt + result, preprocess_mode)
|
||||
|
||||
if is_self_recursive:
|
||||
for _ in range(0, recursive_level):
|
||||
result = generator.generate_text(generated_text, gen_args)
|
||||
generated_text = self.Preprocess(result.text, preprocess_mode)
|
||||
generated_text = self.Preprocess(seed + generated_text, preprocess_mode)
|
||||
generated_text = preprocess(result, preprocess_mode)
|
||||
generated_text = preprocess(prompt + generated_text, preprocess_mode)
|
||||
else:
|
||||
for _ in range(0, recursive_level):
|
||||
result = generator.generate_text(generated_text, gen_args)
|
||||
generated_text += result.text
|
||||
generated_text = self.Preprocess(generated_text, preprocess_mode)
|
||||
|
||||
generated_text += result
|
||||
generated_text = preprocess(generated_text, preprocess_mode)
|
||||
|
||||
return generated_text
|
||||
|
||||
def LogOutputs(self, seed : str, generated_text : str, self_recursive : str, recursive_level : int, preprocess_mode : str, gen_settings, log_filename : str) -> None:
|
||||
|
||||
def log_outputs(
|
||||
self,
|
||||
prompt: str,
|
||||
generated_text: str,
|
||||
self_recursive: str,
|
||||
recursive_level: int,
|
||||
preprocess_mode: str,
|
||||
gen_settings: GenerateArgs,
|
||||
log_filename: str,
|
||||
) -> None:
|
||||
from datetime import datetime
|
||||
|
||||
print_string = f"{' PROMPT GENERATOR OUTPUT '.center(200, '#')}\n{generated_text}\n{'#'*200}\n"
|
||||
print_string = "{' PROMPT GENERATOR OUTPUT '.center(200, '#')}\n"
|
||||
print_string += f"{generated_text}\n"
|
||||
print_string += f"{'#'*200}\n"
|
||||
|
||||
print(print_string)
|
||||
|
||||
log_string = f"{'#'*200}\nDate & Time : {datetime.now()}\nSeed : {seed}\nPrompt : {generated_text}\n"
|
||||
log_string += f"min_length : {gen_settings.min_length}\n"
|
||||
log_string += f"max_length : {gen_settings.max_length}\n"
|
||||
log_string += f"do_sample : {gen_settings.do_sample}\n"
|
||||
log_string += f"early_stopping : {gen_settings.early_stopping}\n"
|
||||
log_string += f"num_beams : {gen_settings.num_beams}\n"
|
||||
log_string += f"temperature : {gen_settings.temperature}\n"
|
||||
log_string += f"top_k : {gen_settings.top_k}\n"
|
||||
log_string += f"top_p : {gen_settings.top_p}\n"
|
||||
log_string += f"no_repeat_ngram_size : {gen_settings.no_repeat_ngram_size}\n"
|
||||
log_string += f"self_recursive : {self_recursive}\nrecursive_level : {recursive_level}\npreprocess_mode : {preprocess_mode}\n"
|
||||
with open(log_filename, "a") as file:
|
||||
file.write(log_string)
|
||||
file.write(f"{'#'*200}\n")
|
||||
file.write(f"Date & Time : {datetime.now()}\n")
|
||||
file.write(f"Prompt : {prompt}\n")
|
||||
file.write(f"Generated Prompt : {generated_text}\n")
|
||||
file.write(f"cfg : {gen_settings.guidance_scale}\n")
|
||||
file.write(f"min_length : {gen_settings.min_length}\n")
|
||||
file.write(f"max_length : {gen_settings.max_length}\n")
|
||||
file.write(f"do_sample : {gen_settings.do_sample}\n")
|
||||
file.write(f"early_stopping : {gen_settings.early_stopping}\n")
|
||||
file.write(f"early_stopping : {gen_settings.early_stopping}\n")
|
||||
file.write(f"num_beams : {gen_settings.num_beams}\n")
|
||||
file.write(f"num_beam_groups : {gen_settings.num_beam_groups}\n")
|
||||
file.write(f"temperature : {gen_settings.temperature}\n")
|
||||
file.write(f"top_k : {gen_settings.top_k}\n")
|
||||
file.write(f"top_p : {gen_settings.top_p}\n")
|
||||
file.write(f"repetition_penalty : {gen_settings.repetition_penalty}\n")
|
||||
file.write(f"no_repeat_ngram_size : {gen_settings.no_repeat_ngram_size}\n")
|
||||
file.write(
|
||||
f"remove_invalid_values : {gen_settings.remove_invalid_values}\n"
|
||||
)
|
||||
file.write(f"self_recursive : {self_recursive}\n")
|
||||
file.write(f"recursive_level : {recursive_level}\n")
|
||||
file.write(f"preprocess_mode : {preprocess_mode}\n")
|
||||
|
||||
def generate(self, clip, model_type, model_name, seed, min_length, max_length, do_sample, early_stopping, num_beams, temperature, top_k, top_p, no_repeat_ngram_size, self_recursive, recursive_level, preprocess_mode):
|
||||
from happytransformer import HappyGeneration, GENSettings
|
||||
def generate(
|
||||
self,
|
||||
clip,
|
||||
model_name,
|
||||
prompt,
|
||||
cfg,
|
||||
min_length,
|
||||
max_length,
|
||||
do_sample,
|
||||
early_stopping,
|
||||
num_beams,
|
||||
num_beam_groups,
|
||||
temperature,
|
||||
top_k,
|
||||
top_p,
|
||||
repetition_penalty,
|
||||
no_repeat_ngram_size,
|
||||
remove_invalid_values,
|
||||
self_recursive,
|
||||
recursive_level,
|
||||
preprocess_mode,
|
||||
):
|
||||
from datetime import date
|
||||
|
||||
root = join("models", "prompt_generators")
|
||||
real_path = join(root, model_name)
|
||||
prompt_log_filename = join("generated_prompts", str(date.today()))
|
||||
prompt_log_filename = join("generated_prompts", str(date.today())) + ".txt"
|
||||
generated_text = ""
|
||||
|
||||
if exists(prompt_log_filename) is False:
|
||||
@@ -184,32 +175,49 @@ class PromptGenerator:
|
||||
|
||||
if exists(real_path) is False:
|
||||
print(f"{real_path} is not exists")
|
||||
generated_text = seed
|
||||
generated_text = prompt
|
||||
else:
|
||||
is_self_recursive = True if self_recursive == "enable" else False
|
||||
|
||||
upper_model_type = model_type.upper()
|
||||
if model_type.find("/") != -1:
|
||||
upper_model_type = model_type.split("/")[1].upper()
|
||||
generator = Generator(real_path)
|
||||
|
||||
generator = HappyGeneration(model_type=upper_model_type, model_name=model_type, load_path=real_path)
|
||||
|
||||
gen_settings = GENSettings(
|
||||
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,
|
||||
no_repeat_ngram_size=no_repeat_ngram_size
|
||||
repetition_penalty=repetition_penalty,
|
||||
no_repeat_ngram_size=no_repeat_ngram_size,
|
||||
remove_invalid_values=True
|
||||
if remove_invalid_values == "enable"
|
||||
else False,
|
||||
)
|
||||
|
||||
generated_text = self.GetGeneratedText(generator, gen_settings, seed, is_self_recursive, recursive_level, preprocess_mode)
|
||||
generated_text = self.get_generated_text(
|
||||
generator,
|
||||
gen_settings,
|
||||
prompt,
|
||||
is_self_recursive,
|
||||
recursive_level,
|
||||
preprocess_mode,
|
||||
)
|
||||
|
||||
self.LogOutputs(seed, generated_text, self_recursive, recursive_level, preprocess_mode, gen_settings, prompt_log_filename)
|
||||
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)
|
||||
return ([[cond, {"pooled_output": pooled}]], )
|
||||
return ([[cond, {"pooled_output": pooled}]],)
|
||||
|
||||
Reference in New Issue
Block a user