accelerator option is added

This commit is contained in:
alpertunga-bile
2023-08-12 14:12:14 +03:00
parent 2f3320c6d0
commit ff56ccec2d
11 changed files with 223 additions and 203 deletions
+3 -1
View File
@@ -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",
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+11 -3
View File
@@ -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"
+23 -14
View File
@@ -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,
+146 -145
View File
@@ -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"
}
+40 -40
View File
@@ -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)