diff --git a/prompttester.py b/prompttester.py index 4d67830..8180b7b 100644 --- a/prompttester.py +++ b/prompttester.py @@ -2,6 +2,7 @@ import sys, os import random import uuid import re +from superprompter.superprompter import * from datetime import datetime sys.path.append(os.path.abspath("..")) @@ -27,6 +28,11 @@ def generateprompts(amount = 1,insanitylevel="5",subject="all", artist="all", im else: result = build_dynamic_prompt(insanitylevel,subject,artist,imagetype, onlyartists,antistring,prefixprompt,suffixprompt,promptcompounderlevel, seperator,givensubject,smartsubject,giventypeofimage,imagemodechance, gender, subtypeobject, subtypehumanoid, subtypeconcept, advancedprompting, hardturnoffemojis, seed, overrideoutfit, prompt_g_and_l, base_model, OBP_preset) + + load_models() + result = answer(input_text=result) + unload_models() + print("") print("loop " + str(steps)) print("") @@ -95,13 +101,14 @@ def generateprompts(amount = 1,insanitylevel="5",subject="all", artist="all", im print("All done!") if __name__ == "__main__": - generateprompts(10,5 + generateprompts(1,5 ,"all" # subject ,"all" # artists - ,"all" # image type "only other types", "only templates mode", "art blaster mode", "quality vomit mode", "color cannon mode", "unique art mode", "massive madness mode", "photo fantasy mode", "subject only mode", "fixed styles mode", "dynamic templates mode", "artify mode" + ,"subject only mode" # image type "only other types", "only templates mode", "art blaster mode", "quality vomit mode", "color cannon mode", "unique art mode", "massive madness mode", "photo fantasy mode", "subject only mode", "fixed styles mode", "dynamic templates mode", "artify mode" , False # only artists - ,"","","PREFIXPROMPT" - ,"SUFFIXPROMPT" + ,"","" + ,"" #prefix prompt + ,"" #suffix prompt ,"",1,"" ,"" # subject override ,True, # smart subject @@ -115,6 +122,6 @@ if __name__ == "__main__": , 0 # seed , "" #outfit override , False #prompt_g_and_l - , "Stable Cascade" #base model + , "SDXL" #base model , "" #preset "All (random)..." ) \ No newline at end of file diff --git a/superprompter/download_models.py b/superprompter/download_models.py new file mode 100644 index 0000000..4921ce1 --- /dev/null +++ b/superprompter/download_models.py @@ -0,0 +1,17 @@ +from transformers import T5Tokenizer, T5ForConditionalGeneration +import torch +import os + +def download_models(): + model_name = "roborovski/superprompt-v1" + tokenizer = T5Tokenizer.from_pretrained(model_name) + model = T5ForConditionalGeneration.from_pretrained(model_name, torch_dtype=torch.float16) + modelDir = os.path.expanduser("~") + "/.superprompter/model_files" + os.makedirs(modelDir, exist_ok=True) + tokenizer.save_pretrained(modelDir) + model.save_pretrained(modelDir) + print("Downloaded SuperPrompt-v1 model files to", modelDir) + return modelDir + +if __name__ == '__main__': + download_models() \ No newline at end of file diff --git a/superprompter/superprompter.py b/superprompter/superprompter.py new file mode 100644 index 0000000..cb3ce44 --- /dev/null +++ b/superprompter/superprompter.py @@ -0,0 +1,67 @@ +#!/usr/bin/env python +import datetime +import os +import random +import tkinter as tk +from tkinter import scrolledtext, ttk +import torch +from transformers import T5Tokenizer, T5ForConditionalGeneration +from superprompter.download_models import download_models + +global tokenizer, model +modelDir = os.path.expanduser("~") + "/.superprompter/model_files" + +def load_models(): + + if not all(os.path.exists(modelDir) for file in modelDir): + print("Model files not found. Downloading...\n") + download_models() + else: + print("Model files found. Skipping download.\n") + + print("Loading SuperPrompt-v1 model...\n") + + + global tokenizer, model + tokenizer = T5Tokenizer.from_pretrained(modelDir) + model = T5ForConditionalGeneration.from_pretrained(modelDir, torch_dtype=torch.float16) + + print("SuperPrompt-v1 model loaded successfully.\n") + +def unload_models(): + global tokenizer, model + del tokenizer + del model + + for file in os.listdir(modelDir): + os.remove(os.path.join(modelDir, file)) + os.rmdir(modelDir) + + + +def answer(input_text="", max_new_tokens=512, repetition_penalty=1.2, temperature=0.5, top_p=1, top_k = 1 , seed=-1): + + # if the seed is "0", generate a random seed and log it to the output + if seed == -1: + seed = random.randint(1, 1000000) + + torch.manual_seed(seed) + + if torch.cuda.is_available(): + device = 'cuda' + else: + device = 'cpu' + + input_ids = tokenizer(input_text, return_tensors="pt").input_ids.to(device) + if torch.cuda.is_available(): + model.to('cuda') + + outputs = model.generate(input_ids, max_new_tokens=max_new_tokens, repetition_penalty=repetition_penalty, + do_sample=True, temperature=temperature, top_p=top_p, top_k=top_k) + + dirty_text = tokenizer.decode(outputs[0]) + text = dirty_text.replace("", "").replace("", "").strip() + print("Temperature: {temperature}\nTop P: {top_p}\nTop K: {top_k}\nSeed: {seed}\nOutput:\n\n") + + print(text) + return text