preprocess algo is optimized and some code parts are refactored
This commit is contained in:
+6
-3
@@ -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.
@@ -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
@@ -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
@@ -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
@@ -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"],),
|
||||||
|
|||||||
Reference in New Issue
Block a user