diff --git a/__init__.py b/__init__.py index 9c41863..9110b76 100644 --- a/__init__.py +++ b/__init__.py @@ -69,7 +69,9 @@ if find_spec("onnxruntime-gpu") is None: # Import PromptGenerator node to ComfyUI -NODE_CLASS_MAPPINGS = {"Prompt Generator": PromptGenerator,} +NODE_CLASS_MAPPINGS = { + "Prompt Generator": PromptGenerator, +} NODE_DISPLAY_NAME_MAPPINGS = { "Prompt Generator": "Prompt Generator", diff --git a/__pycache__/__init__.cpython-311.pyc b/__pycache__/__init__.cpython-311.pyc new file mode 100644 index 0000000..525c667 Binary files /dev/null and b/__pycache__/__init__.cpython-311.pyc differ diff --git a/__pycache__/preprocess.cpython-311.pyc b/__pycache__/preprocess.cpython-311.pyc new file mode 100644 index 0000000..8783977 Binary files /dev/null and b/__pycache__/preprocess.cpython-311.pyc differ diff --git a/__pycache__/prompt_generator.cpython-311.pyc b/__pycache__/prompt_generator.cpython-311.pyc new file mode 100644 index 0000000..f5e66de Binary files /dev/null and b/__pycache__/prompt_generator.cpython-311.pyc differ diff --git a/generator/__pycache__/generate.cpython-311.pyc b/generator/__pycache__/generate.cpython-311.pyc new file mode 100644 index 0000000..73c9058 Binary files /dev/null and b/generator/__pycache__/generate.cpython-311.pyc differ diff --git a/generator/__pycache__/model.cpython-311.pyc b/generator/__pycache__/model.cpython-311.pyc new file mode 100644 index 0000000..e981fc8 Binary files /dev/null and b/generator/__pycache__/model.cpython-311.pyc differ diff --git a/generator/__pycache__/utility.cpython-311.pyc b/generator/__pycache__/utility.cpython-311.pyc new file mode 100644 index 0000000..8edd9fc Binary files /dev/null and b/generator/__pycache__/utility.cpython-311.pyc differ diff --git a/generator/generate.py b/generator/generate.py index 4ef09ce..e475e6a 100644 --- a/generator/generate.py +++ b/generator/generate.py @@ -1,5 +1,9 @@ from dataclasses import dataclass -from generator.model import get_onnx_pipeline, get_default_pipeline +from generator.model import ( + get_default_pipeline, + get_onnx_pipeline, + get_bettertransformer_pipeline, +) from generator.utility import get_accelerator_type, get_variable_dictionary from transformers import Pipeline @@ -27,13 +31,17 @@ class GenerateArgs: class Generator: pipe: Pipeline = None - def __init__(self, model_path: str) -> None: + def __init__(self, model_path: str, is_accelerate: bool) -> None: + if is_accelerate is False: + self.pipe = get_default_pipeline(model_path) + return + 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) + self.pipe = get_bettertransformer_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" diff --git a/generator/model.py b/generator/model.py index ec70748..14bb0bb 100644 --- a/generator/model.py +++ b/generator/model.py @@ -1,25 +1,34 @@ -from transformers import AutoModelForCausalLM, AutoTokenizer, Pipeline +from transformers import AutoModelForCausalLM, AutoTokenizer, Pipeline, pipeline -from optimum.pipelines import pipeline +import optimum.pipelines 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( + pipe = pipeline(task="text-generation", model=model, tokenizer=tokenizer) + + return pipe + + +def get_onnx_pipeline(model_name: str) -> Pipeline: + model = ORTModelForCausalLM.from_pretrained(model_name) + tokenizer = AutoTokenizer.from_pretrained(model_name) + + pipe = optimum.pipelines.pipeline( + task="text-generation", model=model, tokenizer=tokenizer, accelerator="ort" + ) + + return pipe + + +def get_bettertransformer_pipeline(model_name: str) -> Pipeline: + model = AutoModelForCausalLM.from_pretrained(model_name, device_map="auto") + tokenizer = AutoTokenizer.from_pretrained(model_name) + + pipe = optimum.pipelines.pipeline( task="text-generation", model=model, tokenizer=tokenizer, diff --git a/hires.fixWithPromptGenerator.json b/hires.fixWithPromptGenerator.json index 2f28bc3..80e6bf6 100644 --- a/hires.fixWithPromptGenerator.json +++ b/hires.fixWithPromptGenerator.json @@ -1,6 +1,6 @@ { - "last_node_id": 41, - "last_link_id": 56, + "last_node_id": 42, + "last_link_id": 59, "nodes": [ { "id": 3, @@ -14,7 +14,7 @@ "1": 262 }, "flags": {}, - "order": 10, + "order": 11, "mode": 0, "inputs": [ { @@ -25,7 +25,7 @@ { "name": "positive", "type": "CONDITIONING", - "link": 53 + "link": 58 }, { "name": "negative", @@ -52,7 +52,7 @@ "Node name for S&R": "KSampler" }, "widgets_values": [ - 1091892191454277, + 884569790411861, "randomize", 20, 7, @@ -195,7 +195,7 @@ "Node name for S&R": "KSampler" }, "widgets_values": [ - 154049367201419, + 887115133007299, "randomize", 13, 8, @@ -216,7 +216,7 @@ 26 ], "flags": {}, - "order": 8, + "order": 4, "mode": 0, "inputs": [ { @@ -253,7 +253,7 @@ 26 ], "flags": {}, - "order": 12, + "order": 9, "mode": 0, "inputs": [ { @@ -290,7 +290,7 @@ 26 ], "flags": {}, - "order": 9, + "order": 10, "mode": 0, "inputs": [ { @@ -326,7 +326,7 @@ 26 ], "flags": {}, - "order": 4, + "order": 5, "mode": 0, "inputs": [ { @@ -362,7 +362,7 @@ 26 ], "flags": {}, - "order": 5, + "order": 6, "mode": 0, "inputs": [ { @@ -542,13 +542,13 @@ 26 ], "flags": {}, - "order": 11, + "order": 12, "mode": 0, "inputs": [ { "name": "", "type": "*", - "link": 54 + "link": 59 } ], "outputs": [ @@ -566,53 +566,6 @@ "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", @@ -625,7 +578,7 @@ "1": 58 }, "flags": {}, - "order": 3, + "order": 2, "mode": 0, "outputs": [ { @@ -659,7 +612,7 @@ "flags": { "collapsed": true }, - "order": 6, + "order": 7, "mode": 0, "inputs": [ { @@ -687,63 +640,6 @@ "(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", @@ -876,6 +772,111 @@ "widgets_values": [ "ComfyUI" ] + }, + { + "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, + 57 + ], + "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": 42, + "type": "Prompt Generator", + "pos": [ + -140, + -362 + ], + "size": [ + 480.43060302734375, + 580.0348632812501 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "clip", + "type": "CLIP", + "link": 57 + } + ], + "outputs": [ + { + "name": "CONDITIONING", + "type": "CONDITIONING", + "links": [ + 58, + 59 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "Prompt Generator" + }, + "widgets_values": [ + "female_positive-gpt2-141972-5-1", + "enable", + "((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" + ] } ], "links": [ @@ -1063,30 +1064,6 @@ 0, "*" ], - [ - 52, - 4, - 1, - 39, - 0, - "CLIP" - ], - [ - 53, - 39, - 0, - 3, - 1, - "CONDITIONING" - ], - [ - 54, - 39, - 0, - 38, - 0, - "*" - ], [ 55, 8, @@ -1102,6 +1079,30 @@ 41, 0, "IMAGE" + ], + [ + 57, + 4, + 1, + 42, + 0, + "CLIP" + ], + [ + 58, + 42, + 0, + 3, + 1, + "CONDITIONING" + ], + [ + 59, + 42, + 0, + 38, + 0, + "*" ] ], "groups": [ @@ -1120,8 +1121,8 @@ "bounding": [ 926, -441, - 1499, - 1561 + 1468, + 1275 ], "color": "#a1309b" } diff --git a/prompt_generator.py b/prompt_generator.py index 36722ca..72e6990 100644 --- a/prompt_generator.py +++ b/prompt_generator.py @@ -17,6 +17,7 @@ class PromptGenerator: if isdir(join(join("models", "prompt_generators"), file)) ], ), + "accelerator": (["enable", "disable"],), "prompt": ( "STRING", { @@ -109,7 +110,7 @@ class PromptGenerator: ) -> None: from datetime import datetime - print_string = "{' PROMPT GENERATOR OUTPUT '.center(200, '#')}\n" + print_string = f"{' PROMPT GENERATOR OUTPUT '.center(200, '#')}\n" print_string += f"{generated_text}\n" print_string += f"{'#'*200}\n" @@ -144,6 +145,7 @@ class PromptGenerator: self, clip, model_name, + accelerator, prompt, cfg, min_length, @@ -174,49 +176,47 @@ class PromptGenerator: file.close() if exists(real_path) is False: - print(f"{real_path} is not exists") - generated_text = prompt - else: - is_self_recursive = True if self_recursive == "enable" else False + raise ValueError(f"{real_path} is not exists") - generator = Generator(real_path) + is_self_recursive = True if self_recursive == "enable" else False + is_accelerate = True if accelerator == "enable" else False - 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, - repetition_penalty=repetition_penalty, - no_repeat_ngram_size=no_repeat_ngram_size, - remove_invalid_values=True - if remove_invalid_values == "enable" - else False, - ) + generator = Generator(real_path, is_accelerate) - generated_text = self.get_generated_text( - generator, - gen_settings, - prompt, - is_self_recursive, - recursive_level, - preprocess_mode, - ) + 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, + repetition_penalty=repetition_penalty, + no_repeat_ngram_size=no_repeat_ngram_size, + remove_invalid_values=True if remove_invalid_values == "enable" else False, + ) - self.log_outputs( - prompt, - generated_text, - self_recursive, - recursive_level, - preprocess_mode, - gen_settings, - prompt_log_filename, - ) + generated_text = self.get_generated_text( + generator, + gen_settings, + prompt, + is_self_recursive, + recursive_level, + preprocess_mode, + ) + + 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)