testing with superprompter

This commit is contained in:
AIrjen
2024-03-24 10:34:55 +01:00
parent f958f2781a
commit e5544bc0e8
3 changed files with 96 additions and 5 deletions
+12 -5
View File
@@ -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)..."
)
+17
View File
@@ -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()
+67
View File
@@ -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("<pad>", "").replace("</s>", "").strip()
print("Temperature: {temperature}\nTop P: {top_p}\nTop K: {top_k}\nSeed: {seed}\nOutput:\n\n")
print(text)
return text