Compare commits

..

7 Commits

5 changed files with 215 additions and 85 deletions
+30
View File
@@ -0,0 +1,30 @@
---
name: Bug report
about: Create a report to help us improve
title: ''
labels: ''
assignees: ''
---
**Describe the bug**
A clear and concise description of what the bug is.
**To Reproduce**
Steps to reproduce the behavior:
1. Go to '...'
2. Click on '....'
3. Scroll down to '....'
4. See error
**Expected behavior**
A clear and concise description of what you expected to happen.
**Screenshots**
If applicable, add screenshots to help explain your problem.
**What fork of Webui are you using (Eg: Automatic1111, vladmandic):**
**Additional context**
Add any other context about the problem here.
+10
View File
@@ -0,0 +1,10 @@
---
name: Custom issue template
about: Describe this issue template's purpose here.
title: ''
labels: ''
assignees: ''
---
+20
View File
@@ -0,0 +1,20 @@
---
name: Feature request
about: Suggest an idea for this project
title: ''
labels: ''
assignees: ''
---
**Is your feature request related to a problem? Please describe.**
A clear and concise description of what the problem is. Ex. I'm always frustrated when [...]
**Describe the solution you'd like**
A clear and concise description of what you want to happen.
**Describe alternatives you've considered**
A clear and concise description of any alternative solutions or features you've considered.
**Additional context**
Add any other context or screenshots about the feature request here.
+2 -1
View File
@@ -2,9 +2,10 @@
Adds a tab to the webui that allows the user to generate a prompt from a small base prompt. Based on [FredZhang7/distilgpt2-stable-diffusion-v2](https://huggingface.co/FredZhang7/distilgpt2-stable-diffusion-v2) and [Gustavosta/MagicPrompt-Stable-Diffusion](https://huggingface.co/Gustavosta/MagicPrompt-Stable-Diffusion). I did nothing apart from porting it to [AUTOMATIC1111 WebUI](https://github.com/AUTOMATIC1111/stable-diffusion-webui) Adds a tab to the webui that allows the user to generate a prompt from a small base prompt. Based on [FredZhang7/distilgpt2-stable-diffusion-v2](https://huggingface.co/FredZhang7/distilgpt2-stable-diffusion-v2) and [Gustavosta/MagicPrompt-Stable-Diffusion](https://huggingface.co/Gustavosta/MagicPrompt-Stable-Diffusion). I did nothing apart from porting it to [AUTOMATIC1111 WebUI](https://github.com/AUTOMATIC1111/stable-diffusion-webui)
<img width="928" alt="image" src="https://user-images.githubusercontent.com/8998556/218254854-aa59f924-53b1-4514-95bb-20077b7c5aab.png">
![Screenshot 2023-04-29 000027](https://user-images.githubusercontent.com/8998556/235261664-2c92689d-9915-4543-8d6a-57a8ecd0f484.png)
## Installation ## Installation
+153 -84
View File
@@ -10,18 +10,28 @@ 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 pathlib import Path
from modules import script_callbacks from modules import script_callbacks
from transformers import GPT2Tokenizer, GPT2LMHeadModel import modules.scripts as scripts
import re from transformers import GPT2LMHeadModel, GPT2Tokenizer
import json
result_prompt = "" result_prompt = ""
models = {} models = {}
max_no_results = 20 # TODO move to setting panel
base_dir = scripts.basedir()
model_file = Path(base_dir, "models.json")
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,7 +40,10 @@ class Model:
def populate_models(): def populate_models():
path = "./extensions/stable-diffusion-webui-Prompt_Generator/models.json" """Get the models that this extension can use via models.json
"""
# TODO add button to refresh and update model list
path = model_file
with open(path, 'r') as f: with open(path, 'r') as f:
data = json.load(f) data = json.load(f)
for item in data: for item in data:
@@ -40,15 +53,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,42 +71,68 @@ 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): # 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
"""
try: try:
print("[Prompt_Generator]:","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. print("[Prompt_Generator]:","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("[Prompt_Generator]:",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("[Prompt_Generator]:",f"Generate new prompt from: \"{prompt}\" with {name}")
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,
top_k=round(top_k), max_length=max_length, top_k=round(top_k), max_length=max_length,
num_return_sequences=num_return_sequences, num_return_sequences=num_return_sequences,
repetition_penalty=float( repetition_penalty=float(
repetition_penalty), repetition_penalty),
early_stopping=True) early_stopping=True)
else: else:
output = model.generate(input_ids, do_sample=True, temperature=temperature, output = model.generate(input_ids, do_sample=True, temperature=temperature,
top_k=round(top_k), max_length=max_length, top_k=round(top_k), max_length=max_length,
num_return_sequences=num_return_sequences, num_return_sequences=num_return_sequences,
repetition_penalty=float( repetition_penalty=float(
repetition_penalty), repetition_penalty),
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("[Prompt_Generator]:","Generation complete!")
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,28 +144,50 @@ 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("[Prompt_Generator]:",
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
def ui_dynamic_result_batch():
return result_prompt
def save_prompt_to_file(path, append: bool):
if len(result_prompt) == 0:
print("[Prompt_Generator]:","Prompt is empty")
return
with open(path, encoding="utf-8", mode="a" if append else "w") as f:
f.write(result_prompt)
print("[Prompt_Generator]:","Prompt written to: ", path)
# ----------------------------------------------------------------------------
# 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]
with gr.Blocks(analytics_enabled=False) as prompt_generator: with gr.Blocks(analytics_enabled=False) as prompt_generator:
# Handles UI for prompt creation
with gr.Column(): with gr.Column():
with gr.Row(): with gr.Row():
promptTxt = gr.Textbox( prompt_textbox = gr.Textbox(
lines=2, elem_id="promptTxt", label="Start of the prompt") lines=2, elem_id="promptTxt", label="Start of the prompt")
with gr.Column(): with gr.Column():
gr.HTML( gr.HTML(
@@ -141,60 +195,75 @@ def on_ui_tabs():
with gr.Row(): with gr.Row():
temp_slider = gr.Slider( temp_slider = gr.Slider(
elem_id="temp_slider", label="Temperature", interactive=True, minimum=0, maximum=1, value=0.9) elem_id="temp_slider", label="Temperature", interactive=True, minimum=0, maximum=1, value=0.9)
max_length_slider = gr.Slider( maxLength_slider = gr.Slider(
elem_id="max_length_slider", label="Max Length", interactive=True, minimum=1, maximum=200, step=1, value=90) elem_id="max_length_slider", label="Max Length", interactive=True, minimum=1, maximum=200, step=1, value=90)
top_k_slider = gr.Slider( topK_slider = gr.Slider(
elem_id="top_k_slider", label="Top K", value=8, minimum=1, maximum=20, step=1, interactive=True) elem_id="top_k_slider", label="Top K", value=8, minimum=1, maximum=20, step=1, interactive=True)
with gr.Column(): with gr.Column():
with gr.Row(): with gr.Row():
repetition_penalty_slider = gr.Slider( repetitionPenalty_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( numReturnSequences_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?") useBlacklist_checkbox = gr.Checkbox(label="Use blacklist?")
gr.HTML(value="<center>Using <code>\".\extensions\stable-diffusion-webui-Prompt_Generator\\blacklist.txt</code>\".<br>It will delete any matches to the generated result (case insensitive).</center>") gr.HTML(value="<center>Using <code>\".\extensions\stable-diffusion-webui-Prompt_Generator\\blacklist.txt</code>\".<br>It will delete any matches to the generated result (case insensitive).</center>")
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( generate_button = 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 UI for results
generateButton_fred.click(fn=generate_longer_generic, inputs=[ results_vis = []
promptTxt, temp_slider, top_k_slider, max_length_slider, results_txt_list = []
repetition_penalty_slider, num_return_sequences_slider, with gr.Tab("Results"):
generate_dropdown,use_punctuation_check, use_blacklist_checkbox], with gr.Column():
outputs=[results, send_to_img2img, send_to_txt2img, send_to_text, for i in range(max_no_results):
results_col, warning, promptNum_col]) with gr.Row(visible=False) as row:
send_to_img2img.click(add_to_prompt, inputs=[ # Doesn't seem to do anything
promptNum], outputs=[img2img_prompt]) row.style(equal_height=True)
send_to_txt2img.click(add_to_prompt, inputs=[ with gr.Column(scale=3): # Guessing at the scale
promptNum], outputs=[txt2img_prompt]) textBox = gr.Textbox(label="", lines=3)
send_to_text.click(add_to_prompt, inputs=[ with gr.Column(scale=1):
promptNum], outputs=[promptTxt]) txt2img = gr.Button("send to txt2img")
send_to_txt2img.click(None, _js='switch_to_txt2img', img2img = gr.Button("send to img2img")
inputs=None, outputs=None) # Handles ___2img buttons
send_to_img2img.click(None, _js="switch_to_img2img", txt2img.click(add_to_prompt, inputs=[
inputs=None, outputs=None) 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)
with gr.Tab("Batch"):
with gr.Column():
batch_texbox = gr.Textbox("", label="Results")
with gr.Row():
with gr.Column(scale=4):
savePathText = gr.Textbox(
Path(base_dir, "batch_prompt.txt"), label="Path", interactive=True)
with gr.Column(scale=1):
append_checkBox = gr.Checkbox(label="Append")
save_button = gr.Button("Save To file")
# ----------------------------------------------------------------------------------
# Handle buttons
save_button.click(fn=save_prompt_to_file, inputs=[
savePathText, append_checkBox])
# Please note that we use `.then()` to run other ui elements after the generation is done
generate_button.click(fn=generate_longer_generic, inputs=[
prompt_textbox, temp_slider, topK_slider, maxLength_slider,
repetitionPenalty_slider, numReturnSequences_slider,
generate_dropdown, use_punctuation_check, useBlacklist_checkbox]).then(
fn=ui_dynamic_result_visible, inputs=numReturnSequences_slider,
outputs=results_vis).then(
fn=ui_dynamic_result_prompts, outputs=results_txt_list).then(fn=ui_dynamic_result_batch, outputs=batch_texbox)
return (prompt_generator, "Prompt Generator", "Prompt Generator"), return (prompt_generator, "Prompt Generator", "Prompt Generator"),