added optimizations and repo is refactored

This commit is contained in:
alpertunga-bile
2023-08-12 12:58:30 +03:00
committed by GitHub
parent f00b490979
commit 3d665a61ef
7 changed files with 754 additions and 556 deletions
+47 -21
View File
@@ -1,11 +1,25 @@
from sys import path
from os.path import dirname, exists, join
from subprocess import run
from os import remove, mkdir
from os import mkdir
from importlib.util import find_spec
from platform import system
os_name = system()
path.append(dirname(__file__))
from prompt_generator import PromptGenerator
from aspect_node import AspectNode
def check_package(package_name: str) -> None:
if find_spec(package_name) is None:
print(f"/_\ Installing {package_name}")
process = run(
f"pip install {package_name}", shell=True, check=True, capture_output=True
)
print("/_\ Loading Prompt Generator")
@@ -19,36 +33,48 @@ if exists(root) is False:
if exists("generated_prompts") is False:
mkdir("generated_prompts")
# Check happytranformer package
# Check required packages
temp_requirements_file = "temp_requirements.txt"
# Installing from git because there is a error with normal transformers package
if find_spec("transformers") is None:
print(f"/_\ Installing transformers")
process = run(
f"pip install git+https://github.com/huggingface/transformers",
shell=True,
check=True,
capture_output=True,
)
process = run(f"pip freeze > {temp_requirements_file}", shell=True, check=True, capture_output=True)
need_to_install = True
packages = set()
check_package("accelerate")
check_package("xformers")
with open(temp_requirements_file, "r") as file:
packages = set(file.readlines())
if os_name == "Linux":
check_package("triton")
for package in packages:
if "happytransformer" in package:
need_to_install = False
break
check_package("optimum")
remove(temp_requirements_file)
if find_spec("onnxruntime") is None:
print(f"/_\ Installing onnxruntime")
process = run(
f"pip install optimum[onnxruntime]", shell=True, check=True, capture_output=True
)
if need_to_install:
print("/_\ Installing happytransformer")
process = run("pip install happytransformer", shell=True, check=True, capture_output=True)
if find_spec("onnxruntime-gpu") is None:
print(f"/_\ Installing onnxruntime-gpu")
process = run(
f"pip install optimum[onnxruntime-gpu]",
shell=True,
check=True,
capture_output=True,
)
# Import PromptGenerator node to ComfyUI
NODE_CLASS_MAPPINGS = {
"Prompt Generator": PromptGenerator
}
NODE_CLASS_MAPPINGS = {"Prompt Generator": PromptGenerator, "Aspect": AspectNode}
NODE_DISPLAY_NAME_MAPPINGS = {
"Prompt Generator": "Prompt Generator"
"Prompt Generator": "Prompt Generator",
"Aspect": "Aspect Node",
}
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+53
View File
@@ -0,0 +1,53 @@
from dataclasses import dataclass
from generator.model import get_onnx_pipeline, get_default_pipeline
from generator.utility import get_accelerator_type, get_variable_dictionary
from transformers import Pipeline
@dataclass
class GenerateArgs:
num_return_sequences: int = 1
return_full_text: bool = False
min_length: int = 0
max_length: int = 50
early_stopping: bool = False
do_sample: bool = False
num_beams: int = 1
num_beam_groups: int = 1
temperature: float = 1.0
top_k: int = 50
top_p: float = 1.0
repetition_penalty: float = 1.2
no_repeat_ngram_size: int = 0
remove_invalid_values: bool = False
guidance_scale: float = 1.0
@dataclass
class Generator:
pipe: Pipeline = None
def __init__(self, model_path: str) -> None:
accelerator_type = get_accelerator_type(model_path)
if accelerator_type == "onnx":
self.pipe = get_onnx_pipeline(model_name=model_path)
elif accelerator_type == "bettertransformer":
self.pipe = get_default_pipeline(model_name=model_path)
else:
raise ValueError(
"Cant define accelerator type by folder. Can't find .onnx file for onnx, .bin for bettertransformer. Please check your model"
)
def generate_text(
self,
input: str,
args: GenerateArgs = GenerateArgs(),
) -> str:
if self.pipe is None:
raise RuntimeError("Pipeline is NONE. Please check your model path")
args = get_variable_dictionary(args)
output = self.pipe(input, **args)
return output[0]["generated_text"]
+29
View File
@@ -0,0 +1,29 @@
from transformers import AutoModelForCausalLM, AutoTokenizer, Pipeline
from optimum.pipelines import pipeline
from optimum.onnxruntime import ORTModelForCausalLM
def get_onnx_pipeline(model_name: str) -> Pipeline:
model = ORTModelForCausalLM.from_pretrained(model_name)
tokenizer = AutoTokenizer.from_pretrained(model_name)
pipe = pipeline(
task="text-generation", model=model, tokenizer=tokenizer, accelerator="ort"
)
return pipe
def get_default_pipeline(model_name: str) -> Pipeline:
model = AutoModelForCausalLM.from_pretrained(model_name, device_map="auto")
tokenizer = AutoTokenizer.from_pretrained(model_name)
pipe = pipeline(
task="text-generation",
model=model,
tokenizer=tokenizer,
accelerator="bettertransformer",
)
return pipe
+24
View File
@@ -0,0 +1,24 @@
from os import listdir
def get_variable_dictionary(given_class) -> dict:
return {
key: value
for key, value in given_class.__dict__.items()
if not key.startswith("__") and not callable(key)
}
def get_accelerator_type(path: str) -> str:
files = listdir(path)
accelerator_type = "none"
for file in files:
if file.endswith(".onnx"):
accelerator_type = "onnx"
break
if file.endswith(".bin"):
accelerator_type = "bettertransformer"
break
return accelerator_type
+386 -382
View File
@@ -1,6 +1,6 @@
{
"last_node_id": 38,
"last_link_id": 51,
"last_node_id": 41,
"last_link_id": 56,
"nodes": [
{
"id": 3,
@@ -14,7 +14,7 @@
"1": 262
},
"flags": {},
"order": 11,
"order": 10,
"mode": 0,
"inputs": [
{
@@ -25,7 +25,7 @@
{
"name": "positive",
"type": "CONDITIONING",
"link": 17
"link": 53
},
{
"name": "negative",
@@ -52,7 +52,7 @@
"Node name for S&R": "KSampler"
},
"widgets_values": [
462878846866285,
1091892191454277,
"randomize",
20,
7,
@@ -65,8 +65,8 @@
"id": 21,
"type": "ImageUpscaleWithModel",
"pos": [
984,
-104
962.4081015624997,
-106.3239873046875
],
"size": {
"0": 241.79998779296875,
@@ -106,8 +106,8 @@
"id": 24,
"type": "ImageScale",
"pos": [
1266,
-105
1244.4081015625006,
-107.3239873046875
],
"size": {
"0": 315,
@@ -144,53 +144,12 @@
"disabled"
]
},
{
"id": 8,
"type": "VAEDecode",
"pos": [
664.8721572265622,
163.6760126953125
],
"size": {
"0": 210,
"1": 46
},
"flags": {},
"order": 13,
"mode": 0,
"inputs": [
{
"name": "samples",
"type": "LATENT",
"link": 7
},
{
"name": "vae",
"type": "VAE",
"link": 37
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
20,
34
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VAEDecode"
}
},
{
"id": 26,
"type": "KSampler",
"pos": [
1912,
-84
1890.4081015625006,
-86.3239873046875
],
"size": {
"0": 315,
@@ -236,7 +195,7 @@
"Node name for S&R": "KSampler"
},
"widgets_values": [
502249068073047,
154049367201419,
"randomize",
13,
8,
@@ -245,38 +204,6 @@
0.55
]
},
{
"id": 19,
"type": "VAELoader",
"pos": [
-661.1278427734378,
250.6760126953125
],
"size": {
"0": 315,
"1": 58
},
"flags": {},
"order": 2,
"mode": 0,
"outputs": [
{
"name": "VAE",
"type": "VAE",
"links": [
36
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VAELoader"
},
"widgets_values": [
"kl-f8-anime2.ckpt"
]
},
{
"id": 32,
"type": "Reroute",
@@ -289,7 +216,7 @@
26
],
"flags": {},
"order": 4,
"order": 8,
"mode": 0,
"inputs": [
{
@@ -318,15 +245,15 @@
"id": 33,
"type": "Reroute",
"pos": [
1441,
277
1419.4081015625006,
274.67601269531247
],
"size": [
75,
26
],
"flags": {},
"order": 9,
"order": 12,
"mode": 0,
"inputs": [
{
@@ -351,61 +278,19 @@
"horizontal": false
}
},
{
"id": 7,
"type": "CLIPTextEncode",
"pos": [
120.87215722656248,
192.6760126953125
],
"size": {
"0": 441.3026428222656,
"1": 259.2375183105469
},
"flags": {
"collapsed": true
},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 5
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
6,
44
],
"slot_index": 0
}
],
"title": "Negative Prompt",
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"(by bad-artist-anime:0.8, by bad-artist:0.8, verybadimagenegative_v1.3, bad-hands-5, negative_hand, negative_hand-neg, bad_prompt_version2, FastNegativeV2, easynegative, ng_deepnegative_v1_75t), (worst quality:2.0, loli:1.4, old woman:1.4, low quality:2.0, blurry:1.4), (zombie, sketch, interlocked fingers, comic), ((text:2.0, title:2.0, logo:2.0, signature:2.0, watermark)), (crossed eyes, squint, unrealistic eyes, unrealistic pose), duplicate, mutilated, out of frame, extra fingers, mutated hands, poorly drawn hands, poorly drawn face, mutation, extra legs, extra limbs, ugly, bad anatomy, bad proportions, gross proportions, distorted face, poorly drawn body, poorly drawn lips, poorly drawn mouth, poorly drawn hair, unrealistic hair style, cloned face, deformed, extra arms, mutated hands, fused fingers, too many fingers, long neck, poorly drawn feet, bad visual effects, low quality visual effects"
]
},
{
"id": 35,
"type": "Reroute",
"pos": [
1359.222430419922,
144.03236694335936
1337.6305319824226,
141.7083796386719
],
"size": [
140.8,
26
],
"flags": {},
"order": 10,
"order": 9,
"mode": 0,
"inputs": [
{
@@ -441,7 +326,7 @@
26
],
"flags": {},
"order": 5,
"order": 4,
"mode": 0,
"inputs": [
{
@@ -465,66 +350,19 @@
"horizontal": false
}
},
{
"id": 4,
"type": "CheckpointLoaderSimple",
"pos": [
-654.1278427734378,
71.6760126953125
],
"size": {
"0": 315,
"1": 98
},
"flags": {},
"order": 3,
"mode": 0,
"outputs": [
{
"name": "MODEL",
"type": "MODEL",
"links": [
46,
48
],
"slot_index": 0
},
{
"name": "CLIP",
"type": "CLIP",
"links": [
5,
16
],
"slot_index": 1
},
{
"name": "VAE",
"type": "VAE",
"links": [],
"slot_index": 2
}
],
"properties": {
"Node name for S&R": "CheckpointLoaderSimple"
},
"widgets_values": [
"wololo-mix-v2.safetensors"
]
},
{
"id": 37,
"type": "Reroute",
"pos": [
1712,
-161
1690.4081015625006,
-163.32398730468748
],
"size": [
82,
26
],
"flags": {},
"order": 6,
"order": 5,
"mode": 0,
"inputs": [
{
@@ -552,8 +390,8 @@
"id": 25,
"type": "VAEEncode",
"pos": [
1653,
-27
1631.4081015625006,
-29.3239873046875
],
"size": {
"0": 210,
@@ -593,8 +431,8 @@
"id": 34,
"type": "Reroute",
"pos": [
1681,
-101
1659.4081015625006,
-103.3239873046875
],
"size": [
140.8,
@@ -625,59 +463,6 @@
"horizontal": false
}
},
{
"id": 18,
"type": "Prompt Generator",
"pos": [
-125.12784277343758,
-333.3239873046875
],
"size": {
"0": 445.82281494140625,
"1": 461.5046691894531
},
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 16
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
17,
50
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "Prompt Generator"
},
"widgets_values": [
"gpt2",
"female_positive-gpt2-141972-5-1",
"1girl, solo, koenigsegg jesko, koenigsegg, car, jesko, mature woman, white hair, short hair, driving high speed car, driving car",
20,
50,
"disable",
"disable",
1,
1,
50,
1,
0,
"disable",
0
]
},
{
"id": 5,
"type": "EmptyLatentImage",
@@ -713,79 +498,12 @@
1
]
},
{
"id": 31,
"type": "SaveImage",
"pos": [
1903,
351
],
"size": [
460.6601699218754,
738.5096362304689
],
"flags": {},
"order": 21,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 35
}
],
"properties": {},
"widgets_values": [
"ComfyUI"
]
},
{
"id": 27,
"type": "VAEDecode",
"pos": [
1951,
256
],
"size": {
"0": 210,
"1": 46
},
"flags": {},
"order": 20,
"mode": 0,
"inputs": [
{
"name": "samples",
"type": "LATENT",
"link": 29
},
{
"name": "vae",
"type": "VAE",
"link": 39
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
35
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VAEDecode"
}
},
{
"id": 22,
"type": "UpscaleModelLoader",
"pos": [
983,
-221
961.4081015624997,
-223.32398730468748
],
"size": {
"0": 315,
@@ -812,50 +530,25 @@
"RealESRGAN_x4plus_anime_6B.pth"
]
},
{
"id": 30,
"type": "PreviewImage",
"pos": [
562.8721572265622,
274.6760126953125
],
"size": [
327.5105535333803,
515.4847634055394
],
"flags": {},
"order": 16,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 34
}
],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 38,
"type": "Reroute",
"pos": [
1108,
-324
1086.4081015625006,
-326.32398730468753
],
"size": [
140.8,
26
],
"flags": {},
"order": 12,
"order": 11,
"mode": 0,
"inputs": [
{
"name": "",
"type": "*",
"link": 50
"link": 54
}
],
"outputs": [
@@ -872,6 +565,317 @@
"showOutputText": true,
"horizontal": false
}
},
{
"id": 4,
"type": "CheckpointLoaderSimple",
"pos": [
-654.1278427734378,
71.6760126953125
],
"size": {
"0": 315,
"1": 98
},
"flags": {},
"order": 2,
"mode": 0,
"outputs": [
{
"name": "MODEL",
"type": "MODEL",
"links": [
46,
48
],
"slot_index": 0
},
{
"name": "CLIP",
"type": "CLIP",
"links": [
5,
52
],
"slot_index": 1
},
{
"name": "VAE",
"type": "VAE",
"links": [],
"slot_index": 2
}
],
"properties": {
"Node name for S&R": "CheckpointLoaderSimple"
},
"widgets_values": [
"wololo-mix-v3.safetensors"
]
},
{
"id": 19,
"type": "VAELoader",
"pos": [
-661.1278427734378,
250.6760126953125
],
"size": {
"0": 315,
"1": 58
},
"flags": {},
"order": 3,
"mode": 0,
"outputs": [
{
"name": "VAE",
"type": "VAE",
"links": [
36
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VAELoader"
},
"widgets_values": [
"blessed2.vae.pt"
]
},
{
"id": 7,
"type": "CLIPTextEncode",
"pos": [
-7,
261
],
"size": {
"0": 441.3026428222656,
"1": 259.2375183105469
},
"flags": {
"collapsed": true
},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 5
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
6,
44
],
"slot_index": 0
}
],
"title": "Negative Prompt",
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"(by bad-artist-anime:0.8, by bad-artist:0.8, verybadimagenegative_v1.3, bad-hands-5, negative_hand, negative_hand-neg, bad_prompt_version2, FastNegativeV2, easynegative, ng_deepnegative_v1_75t), (worst quality:2.0, loli:1.4, old woman:1.4, low quality:2.0, blurry:1.4), (zombie, sketch, interlocked fingers, comic), ((text:2.0, title:2.0, logo:2.0, signature:2.0, watermark)), (crossed eyes, squint, unrealistic eyes, unrealistic pose), duplicate, mutilated, out of frame, extra fingers, mutated hands, poorly drawn hands, poorly drawn face, mutation, extra legs, extra limbs, ugly, bad anatomy, bad proportions, gross proportions, distorted face, poorly drawn body, poorly drawn lips, poorly drawn mouth, poorly drawn hair, unrealistic hair style, cloned face, deformed, extra arms, mutated hands, fused fingers, too many fingers, long neck, poorly drawn feet, bad visual effects, low quality visual effects"
]
},
{
"id": 39,
"type": "Prompt Generator",
"pos": [
-129,
-344
],
"size": {
"0": 436,
"1": 549
},
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 52
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
53,
54
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "Prompt Generator"
},
"widgets_values": [
"female_positive_generator",
"((masterpiece, best quality, ultra detailed)), illustration, digital art, 1girl, solo, ((stunningly beautiful))",
1,
20,
50,
"disable",
"disable",
1,
1,
1,
50,
1,
1,
0,
"disable",
"disable",
0,
"exact_keyword"
]
},
{
"id": 8,
"type": "VAEDecode",
"pos": [
664.8721572265622,
163.6760126953125
],
"size": {
"0": 210,
"1": 46
},
"flags": {},
"order": 13,
"mode": 0,
"inputs": [
{
"name": "samples",
"type": "LATENT",
"link": 7
},
{
"name": "vae",
"type": "VAE",
"link": 37
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
20,
55
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VAEDecode"
}
},
{
"id": 40,
"type": "PreviewImage",
"pos": [
679,
302
],
"size": {
"0": 210,
"1": 26
},
"flags": {},
"order": 16,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 55
}
],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 27,
"type": "VAEDecode",
"pos": [
1929.4081015625006,
253.67601269531252
],
"size": {
"0": 210,
"1": 46
},
"flags": {},
"order": 20,
"mode": 0,
"inputs": [
{
"name": "samples",
"type": "LATENT",
"link": 29
},
{
"name": "vae",
"type": "VAE",
"link": 39
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
56
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VAEDecode"
}
},
{
"id": 41,
"type": "SaveImage",
"pos": [
1937,
363
],
"size": {
"0": 315,
"1": 58
},
"flags": {},
"order": 21,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 56
}
],
"properties": {},
"widgets_values": [
"ComfyUI"
]
}
],
"links": [
@@ -907,22 +911,6 @@
0,
"LATENT"
],
[
16,
4,
1,
18,
0,
"CLIP"
],
[
17,
18,
0,
3,
1,
"CONDITIONING"
],
[
20,
8,
@@ -971,22 +959,6 @@
0,
"LATENT"
],
[
34,
8,
0,
30,
0,
"IMAGE"
],
[
35,
27,
0,
31,
0,
"IMAGE"
],
[
36,
19,
@@ -1083,14 +1055,6 @@
0,
"MODEL"
],
[
50,
18,
0,
38,
0,
"*"
],
[
51,
38,
@@ -1098,6 +1062,46 @@
34,
0,
"*"
],
[
52,
4,
1,
39,
0,
"CLIP"
],
[
53,
39,
0,
3,
1,
"CONDITIONING"
],
[
54,
39,
0,
38,
0,
"*"
],
[
55,
8,
0,
40,
0,
"IMAGE"
],
[
56,
27,
0,
41,
0,
"IMAGE"
]
],
"groups": [
@@ -1114,8 +1118,8 @@
{
"title": "Hires .fix",
"bounding": [
948,
-439,
926,
-441,
1499,
1561
],
+54
View File
@@ -0,0 +1,54 @@
from string import punctuation
from re import sub, compile
def get_unique_list(sequence : list) -> list:
seen = set()
return [x for x in sequence if not (x in seen or seen.add(x))]
def remove_exact_keywords(line : str) -> list[str]:
char_blacklist = set(f"{punctuation}0123456789")
# remove exact prompts
prompts = get_unique_list(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 pure_prompts
def preprocess(line : str, preprocess_mode : str) -> str:
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(remove_exact_keywords(temp_line))
elif preprocess_mode == "exact_prompt":
temp_line = ','.join(get_unique_list(temp_line.split(",")))
return temp_line
+161 -153
View File
@@ -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}]],)