preprocess algo is optimized and some code parts are refactored

This commit is contained in:
alpertunga-bile
2023-11-29 23:43:38 +03:00
parent 22bdd0f276
commit b82f11460b
11 changed files with 39 additions and 33 deletions
+6 -3
View File
@@ -12,18 +12,18 @@ from prompt_generator import PromptGenerator
print("/_\ Loading Prompt Generator") print("/_\ Loading Prompt Generator")
# Check prompt_generators folder under the models folder # check prompt_generators folder under the models folder
root = join(models_dir, "prompt_generators") root = join(models_dir, "prompt_generators")
if exists(root) is False: if exists(root) is False:
print(f"/_\ {root} is created. Please add your prompt generators to {root} folder")
mkdir(root) mkdir(root)
print(f"/_\ {root} is created. Please add your prompt generators to {root} folder")
prompts_file = join(base_path, "generated_prompts") prompts_file = join(base_path, "generated_prompts")
if exists(prompts_file) is False: if exists(prompts_file) is False:
mkdir(prompts_file) mkdir(prompts_file)
# Import PromptGenerator node to ComfyUI # import PromptGenerator node to ComfyUI
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
"Prompt Generator": PromptGenerator, "Prompt Generator": PromptGenerator,
@@ -34,3 +34,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
} }
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
print("/_\ Loaded Successfully")
print("-" * 100)
Binary file not shown.
Binary file not shown.
Binary file not shown.
+5 -2
View File
@@ -15,7 +15,7 @@ def check_package(package_name: str, install_name: str) -> None:
print(f"/_\ Installing {package_name}") print(f"/_\ Installing {package_name}")
command = "" # check if portable version or manual
if exists("python_embeded"): if exists("python_embeded"):
command = f".\\python_embeded\\python.exe -s -m pip install {install_name}" command = f".\\python_embeded\\python.exe -s -m pip install {install_name}"
else: else:
@@ -24,17 +24,20 @@ def check_package(package_name: str, install_name: str) -> None:
process = run(command, shell=True, check=True, capture_output=True) process = run(command, shell=True, check=True, capture_output=True)
print(" Prompt Generator ComfyUI Node ".center(100, "-"))
# Check required packages # Check required packages
print("/_\ Checking packages") print("/_\ Checking packages")
check_package("transformers", "transformers") check_package("transformers", "transformers")
check_package("accelerate", "accelerate") check_package("accelerate", "accelerate")
# triton package exists only in Linux
if os_name == "Linux": if os_name == "Linux":
check_package("triton", "triton") check_package("triton", "triton")
check_package("optimum", "optimum") check_package("optimum", "optimum")
check_package("onnxruntime-gpu", "optimum[onnxruntime-gpu]") check_package("onnxruntime", "optimum[onnxruntime-gpu]")
""" """
# This package is for onnx models that run on CPU # This package is for onnx models that run on CPU
Binary file not shown.
Binary file not shown.
Binary file not shown.
+7 -5
View File
@@ -1,6 +1,8 @@
from transformers import AutoModelForCausalLM, AutoTokenizer, Pipeline, pipeline from transformers import AutoModelForCausalLM, AutoTokenizer, Pipeline
from transformers import pipeline as tf_pipe
from optimum.pipelines import pipeline as opt_pipe
import optimum.pipelines
from optimum.onnxruntime import ORTModelForCausalLM from optimum.onnxruntime import ORTModelForCausalLM
@@ -8,7 +10,7 @@ def get_default_pipeline(model_name: str) -> Pipeline:
model = AutoModelForCausalLM.from_pretrained(model_name, device_map="auto") model = AutoModelForCausalLM.from_pretrained(model_name, device_map="auto")
tokenizer = AutoTokenizer.from_pretrained(model_name) tokenizer = AutoTokenizer.from_pretrained(model_name)
pipe = pipeline(task="text-generation", model=model, tokenizer=tokenizer) pipe = tf_pipe(task="text-generation", model=model, tokenizer=tokenizer)
return pipe return pipe
@@ -17,7 +19,7 @@ def get_onnx_pipeline(model_name: str) -> Pipeline:
model = ORTModelForCausalLM.from_pretrained(model_name) model = ORTModelForCausalLM.from_pretrained(model_name)
tokenizer = AutoTokenizer.from_pretrained(model_name) tokenizer = AutoTokenizer.from_pretrained(model_name)
pipe = optimum.pipelines.pipeline( pipe = opt_pipe(
task="text-generation", model=model, tokenizer=tokenizer, accelerator="ort" task="text-generation", model=model, tokenizer=tokenizer, accelerator="ort"
) )
@@ -28,7 +30,7 @@ def get_bettertransformer_pipeline(model_name: str) -> Pipeline:
model = AutoModelForCausalLM.from_pretrained(model_name, device_map="auto") model = AutoModelForCausalLM.from_pretrained(model_name, device_map="auto")
tokenizer = AutoTokenizer.from_pretrained(model_name) tokenizer = AutoTokenizer.from_pretrained(model_name)
pipe = optimum.pipelines.pipeline( pipe = opt_pipe(
task="text-generation", task="text-generation",
model=model, model=model,
tokenizer=tokenizer, tokenizer=tokenizer,
+20 -22
View File
@@ -1,54 +1,52 @@
from string import punctuation from string import punctuation
from re import sub, compile from re import sub, compile
from collections import OrderedDict
def get_unique_list(sequence : list) -> list:
def get_unique_list(sequence: list) -> list:
seen = set() seen = set()
return [x for x in sequence if not (x in seen or seen.add(x))] return [x for x in sequence if not (x in seen or seen.add(x))]
def remove_exact_keywords(line : str) -> list[str]:
def remove_exact_keywords(line: str) -> list[str]:
char_blacklist = set(f"{punctuation}0123456789") char_blacklist = set(f"{punctuation}0123456789")
# remove exact prompts # remove exact prompts
prompts = get_unique_list(line.split(",")) prompts = get_unique_list(line.split(","))
pure_prompts = []
pure_prompts = OrderedDict() # order matters
extracted_pure_prompts = {} # order isn't important
# remove exact keyword # remove exact keyword
for prompt in prompts: for prompt in prompts:
can_add = True
# extract the keyword # extract the keyword
keyword = "".join(c for c in prompt if c not in char_blacklist).lstrip() keyword = "".join(c for c in prompt if c not in char_blacklist).lstrip()
if keyword == "": if keyword == "":
continue continue
for pure_prompt in pure_prompts: if prompt in pure_prompts or keyword in extracted_pure_prompts:
if prompt == pure_prompt: continue
can_add = False
break
extracted_pure_prompt = "".join(c for c in pure_prompt if c not in char_blacklist).lstrip() pure_prompts[prompt] = True
if keyword == extracted_pure_prompt: extracted_pure_prompts[keyword] = True
can_add = False
break
if can_add: return pure_prompts.keys()
pure_prompts.append(prompt)
return pure_prompts
def preprocess(line : str, preprocess_mode : str) -> str: def preprocess(line: str, preprocess_mode: str) -> str:
pattern = compile(r'(,\s){2,}') pattern = compile(r"(,\s){2,}")
temp_line = line.replace(u'\xa0', u' ') temp_line = line.replace("\xa0", " ")
temp_line = temp_line.replace("\n", ", ") temp_line = temp_line.replace("\n", ", ")
temp_line = temp_line.replace("\t", " ") temp_line = temp_line.replace("\t", " ")
temp_line = temp_line.replace("|", ",") temp_line = temp_line.replace("|", ",")
temp_line = temp_line.replace(" ", " ") temp_line = temp_line.replace(" ", " ")
temp_line = sub(pattern, ', ', temp_line) temp_line = sub(pattern, ", ", temp_line)
if preprocess_mode == "exact_keyword": if preprocess_mode == "exact_keyword":
temp_line = ','.join(remove_exact_keywords(temp_line)) temp_line = ",".join(remove_exact_keywords(temp_line))
elif preprocess_mode == "exact_prompt": elif preprocess_mode == "exact_prompt":
temp_line = ','.join(get_unique_list(temp_line.split(","))) temp_line = ",".join(get_unique_list(temp_line.split(",")))
return temp_line return temp_line
+1 -1
View File
@@ -16,7 +16,7 @@ class PromptGenerator:
[ [
file file
for file in listdir(join(models_dir, "prompt_generators")) for file in listdir(join(models_dir, "prompt_generators"))
if isdir(join(join(models_dir, "prompt_generators"), file)) if isdir(join(models_dir, "prompt_generators", file))
], ],
), ),
"accelerate": (["enable", "disable"],), "accelerate": (["enable", "disable"],),