Made the results UI dynamic Feature Request: "Send to X" buttons for prompts. #20

This commit is contained in:
2023-04-29 00:03:15 +02:00
parent 0d4de20f95
commit 4fd1f9fed2
+103 -59
View File
@@ -10,18 +10,24 @@ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLI
""" """
import json
import re
import gradio as gr import gradio as gr
import modules import modules
from modules import script_callbacks from modules import script_callbacks
from transformers import GPT2Tokenizer, GPT2LMHeadModel from transformers import GPT2LMHeadModel, GPT2Tokenizer
import re
import json
result_prompt = "" result_prompt = ""
models = {} models = {}
max_no_results = 20 # TODO move to setting panel
class Model: class Model:
'''
Small strut to hold data for the text generator
'''
def __init__(self, name, model, tokenizer) -> None: def __init__(self, name, model, tokenizer) -> None:
self.name = name self.name = name
self.model = model self.model = model
@@ -30,6 +36,8 @@ class Model:
def populate_models(): def populate_models():
"""Get the models that this extension can use via models.json
"""
path = "./extensions/stable-diffusion-webui-Prompt_Generator/models.json" path = "./extensions/stable-diffusion-webui-Prompt_Generator/models.json"
with open(path, 'r') as f: with open(path, 'r') as f:
data = json.load(f) data = json.load(f)
@@ -40,15 +48,8 @@ def populate_models():
models[name] = Model(name, model, tokenizer) models[name] = Model(name, model, tokenizer)
def add_to_prompt(prompt): # A holder TODO figure out how to get rid of it
def add_to_prompt(num): # A function that determines which prompt to pass return prompt
hand_over_prompt_list = result_prompt.splitlines()
try:
return (hand_over_prompt_list[int(num)-1][3:])
except Exception as e:
print(
f"That line does not exist. Check number of prompts: {e}")
return gr.update(), f"Error: {e}"
def get_list_blacklist(): def get_list_blacklist():
@@ -65,19 +66,49 @@ def get_list_blacklist():
def on_ui_tabs(): def on_ui_tabs():
# Method to create the extended prompt # Method to create the extended prompt
def generate_longer_generic(prompt, temperature, top_k, def generate_longer_generic(prompt, temperature, top_k,
max_length, repetition_penalty, num_return_sequences, name, use_punctuation=False, use_blacklist=False): max_length, repetition_penalty,
num_return_sequences, name, use_punctuation=False,
use_blacklist=False, progress=gr.Progress()): # TODO make the progress bar work
"""Generates a longer string from the input
Args:
prompt (str): As the name suggests, the start of the prompt that the generator should start with.
temperature (float): A higher temperature will produce more diverse results, but with a higher risk of less coherent text
top_k (float): Strategy is to sample from a shortlist of the top K tokens. This approach allows the other high-scoring tokens a chance of being picked.
max_length (int): the maximum number of tokens for the output of the model
repetition_penalty (float): The parameter for repetition penalty. 1.0 means no penalty. Default setting is 1.2
num_return_sequences (int): The number of results to generate
name (str): Which Model to use
use_punctuation (bool): Allows the use of commas in the output. Defaults to False.
use_blacklist (bool): It will delete any matches to the generated result (case insensitive). Each item to be filtered out should be on a new line. Defaults to False.
Returns:
Returns only an error otherwise saves it in result_prompt
"""
progress(0, "Starting")
try: try:
progress(0.25)
print("Loading Tokenizer")
tokenizer = GPT2Tokenizer.from_pretrained(models[name].tokenizer) tokenizer = GPT2Tokenizer.from_pretrained(models[name].tokenizer)
tokenizer.add_special_tokens({'pad_token': '[PAD]'}) tokenizer.add_special_tokens({'pad_token': '[PAD]'})
# Full credits for the model to FredZhang7 (https://huggingface.co/FredZhang7). Under creativeml-openrail-m license. progress(0.5)
print("Loading Model")
model = GPT2LMHeadModel.from_pretrained(models[name].model) model = GPT2LMHeadModel.from_pretrained(models[name].model)
except Exception as e: except Exception as e:
print(f"Exception encountered while attempting to install tokenizer") print(f"Exception encountered while attempting to install tokenizer")
return gr.update(), f"Error: {e}" return gr.update(), f"Error: {e}"
try: try:
print(f"Generate new prompt from: \"{prompt}\" with {name}") print(f"Generate new prompt from: \"{prompt}\" with {name}")
progress(0.75)
input_ids = tokenizer(prompt, return_tensors='pt').input_ids input_ids = tokenizer(prompt, return_tensors='pt').input_ids
if (use_punctuation): if (use_punctuation):
output = model.generate(input_ids, do_sample=True, temperature=temperature, output = model.generate(input_ids, do_sample=True, temperature=temperature,
@@ -95,12 +126,13 @@ def on_ui_tabs():
penalty_alpha=0.6, no_repeat_ngram_size=1, penalty_alpha=0.6, no_repeat_ngram_size=1,
early_stopping=True) early_stopping=True)
print("Generation complete!") print("Generation complete!")
progress(1, "Done!")
tempString = "" tempString = ""
if (use_blacklist): if (use_blacklist):
blacklist = get_list_blacklist() blacklist = get_list_blacklist()
for i in range(len(output)): for i in range(len(output)):
tempString += str(i+1)+": "+tokenizer.decode( tempString += tokenizer.decode(
output[i], skip_special_tokens=True) + "\n" output[i], skip_special_tokens=True) + "\n"
if (use_blacklist): if (use_blacklist):
@@ -112,20 +144,30 @@ def on_ui_tabs():
result_prompt = tempString result_prompt = tempString
# print(result_prompt) # print(result_prompt)
return {results: tempString,
send_to_img2img: gr.update(visible=True),
send_to_txt2img: gr.update(visible=True),
send_to_text: gr.update(visible=True),
results_col: gr.update(visible=True),
warning: gr.update(visible=True),
promptNum_col: gr.update(visible=True)
}
except Exception as e: except Exception as e:
print( print(
f"Exception encountered while attempting to generate prompt: {e}") f"Exception encountered while attempting to generate prompt: {e}")
return gr.update(), f"Error: {e}" return gr.update(), f"Error: {e}"
def ui_dynamic_result_visible(num):
"""Makes the results visible"""
k = int(num)
return [gr.Row.update(visible=True)]*k + [gr.Row.update(visible=False)]*(max_no_results-k)
def ui_dynamic_result_prompts():
"""Populates the results with the prompts"""
lines = result_prompt.splitlines()
num = len(lines)
result_list = []
for i in range(int(max_no_results)):
if (i < num):
result_list.append(lines[i])
else:
result_list.append("")
return result_list
# ----------------------------------------------------------------------------
# UI structure # UI structure
txt2img_prompt = modules.ui.txt2img_paste_fields[0][0] txt2img_prompt = modules.ui.txt2img_paste_fields[0][0]
img2img_prompt = modules.ui.img2img_paste_fields[0][0] img2img_prompt = modules.ui.img2img_paste_fields[0][0]
@@ -150,7 +192,7 @@ def on_ui_tabs():
repetition_penalty_slider = gr.Slider( repetition_penalty_slider = gr.Slider(
elem_id="repetition_penalty_slider", label="Repetition Penalty", value=1.2, minimum=0.1, maximum=10, interactive=True) elem_id="repetition_penalty_slider", label="Repetition Penalty", value=1.2, minimum=0.1, maximum=10, interactive=True)
num_return_sequences_slider = gr.Slider( num_return_sequences_slider = gr.Slider(
elem_id="num_return_sequences_slider", label="How Many To Generate", value=5, minimum=1, maximum=20, interactive=True, step=1) elem_id="num_return_sequences_slider", label="How Many To Generate", value=5, minimum=1, maximum=max_no_results, interactive=True, step=1)
with gr.Column(): with gr.Column():
with gr.Row(): with gr.Row():
use_blacklist_checkbox = gr.Checkbox(label="Use blacklist?") use_blacklist_checkbox = gr.Checkbox(label="Use blacklist?")
@@ -158,43 +200,45 @@ def on_ui_tabs():
with gr.Column(): with gr.Column():
with gr.Row(): with gr.Row():
populate_models() populate_models()
generate_dropdown = gr.Dropdown(choices=list(models.keys()), value="FredZhang7", label = "Which model to use?",show_label=True) generate_dropdown = gr.Dropdown(choices=list(models.keys()), value=list(models.keys())[
1 if len(models) > 0 else 0], label="Which model to use?", show_label=True) # TODO Add default to setting page
use_punctuation_check = gr.Checkbox(label="Use punctuation?") use_punctuation_check = gr.Checkbox(label="Use punctuation?")
generateButton_fred = gr.Button( generateButton = gr.Button(
value="Generate", elem_id="generate_button") value="Generate", elem_id="generate_button") # TODO Add element to show that it is working in the background so users don't think nothing is happening
with gr.Column(visible=False) as results_col:
results = gr.Text(
label="Results", elem_id="Results_textBox", interactive=False)
with gr.Column(visible=False) as promptNum_col:
with gr.Row():
promptNum = gr.Textbox(
lines=1, elem_id="promptNum", label="Send which prompt")
with gr.Column():
warning = gr.HTML(
value="Select one number and send that prompt to txt2img or img2img", visible=False)
with gr.Row():
send_to_txt2img = gr.Button('Send to txt2img', visible=False)
send_to_img2img = gr.Button('Send to img2img', visible=False)
send_to_text = gr.Button(
'Send to back to prompter', visible=False)
# events # Handles Dynamic results
generateButton_fred.click(fn=generate_longer_generic, inputs=[ results_vis = []
results_txt_list = []
with gr.Column() as results_col:
for i in range(max_no_results):
with gr.Row(visible=False) as row:
row.style(equal_height=True) # Doesn't seem to do anything
with gr.Column(scale=3): # Guessing at the scale
textBox = gr.Textbox(label="")
with gr.Column(scale=1):
txt2img = gr.Button("send to txt2img")
img2img = gr.Button("send to img2img")
# Handles ___2img buttons
txt2img.click(add_to_prompt, inputs=[
textBox], outputs=[txt2img_prompt]).then(None, _js='switch_to_txt2img',
inputs=None, outputs=None)
img2img.click(add_to_prompt, inputs=[
textBox], outputs=[img2img_prompt]).then(None, _js='switch_to_img2img',
inputs=None, outputs=None)
results_txt_list.append(textBox)
results_vis.append(row)
# ----------------------------------------------------------------------------------
# Handle buttons
#Please note that we use `.then()` to run other ui elements after the generation is done
generateButton.click(fn=generate_longer_generic, inputs=[
promptTxt, temp_slider, top_k_slider, max_length_slider, promptTxt, temp_slider, top_k_slider, max_length_slider,
repetition_penalty_slider, num_return_sequences_slider, repetition_penalty_slider, num_return_sequences_slider,
generate_dropdown,use_punctuation_check, use_blacklist_checkbox], generate_dropdown, use_punctuation_check, use_blacklist_checkbox]).then(
outputs=[results, send_to_img2img, send_to_txt2img, send_to_text, fn=ui_dynamic_result_visible, inputs=num_return_sequences_slider,
results_col, warning, promptNum_col]) outputs=results_vis).then(
send_to_img2img.click(add_to_prompt, inputs=[ fn=ui_dynamic_result_prompts, outputs=results_txt_list)
promptNum], outputs=[img2img_prompt])
send_to_txt2img.click(add_to_prompt, inputs=[
promptNum], outputs=[txt2img_prompt])
send_to_text.click(add_to_prompt, inputs=[
promptNum], outputs=[promptTxt])
send_to_txt2img.click(None, _js='switch_to_txt2img',
inputs=None, outputs=None)
send_to_img2img.click(None, _js="switch_to_img2img",
inputs=None, outputs=None)
return (prompt_generator, "Prompt Generator", "Prompt Generator"), return (prompt_generator, "Prompt Generator", "Prompt Generator"),