mirror of
https://github.com/imrayya/stable-diffusion-webui-Prompt_Generator.git
synced 2024-01-11 09:00:44 +01:00
Compare commits
7 Commits
0d4de20f95
...
master
| Author | SHA1 | Date | |
|---|---|---|---|
| 31ec08338b | |||
| 89215f828d | |||
| d0ce3a73b6 | |||
| a361185ad5 | |||
| 961237ea06 | |||
| 275ca5511e | |||
| 4fd1f9fed2 |
@@ -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.
|
||||||
@@ -0,0 +1,10 @@
|
|||||||
|
---
|
||||||
|
name: Custom issue template
|
||||||
|
about: Describe this issue template's purpose here.
|
||||||
|
title: ''
|
||||||
|
labels: ''
|
||||||
|
assignees: ''
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
|
||||||
@@ -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,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">
|
|
||||||
|
|
||||||
|
|
||||||
|

|
||||||
|
|
||||||
|
|
||||||
## Installation
|
## Installation
|
||||||
|
|
||||||
|
|||||||
+139
-70
@@ -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,19 +71,45 @@ 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,
|
||||||
@@ -94,13 +126,13 @@ def on_ui_tabs():
|
|||||||
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")
|
||||||
|
# Handles ___2img buttons
|
||||||
|
txt2img.click(add_to_prompt, inputs=[
|
||||||
|
textBox], outputs=[txt2img_prompt]).then(None, _js='switch_to_txt2img',
|
||||||
inputs=None, outputs=None)
|
inputs=None, outputs=None)
|
||||||
send_to_img2img.click(None, _js="switch_to_img2img",
|
img2img.click(add_to_prompt, inputs=[
|
||||||
|
textBox], outputs=[img2img_prompt]).then(None, _js='switch_to_img2img',
|
||||||
inputs=None, outputs=None)
|
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"),
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user