added optimizations and repo is refactored
This commit is contained in:
+47
-21
@@ -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"]
|
||||
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
@@ -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
@@ -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
|
||||
],
|
||||
|
||||
@@ -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
@@ -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