mirror of
https://github.com/imrayya/stable-diffusion-webui-Prompt_Generator.git
synced 2024-01-11 09:00:44 +01:00
Compare commits
15 Commits
b4833e90ce
..
master
| Author | SHA1 | Date | |
|---|---|---|---|
| 31ec08338b | |||
| 89215f828d | |||
| d0ce3a73b6 | |||
| a361185ad5 | |||
| 961237ea06 | |||
| 275ca5511e | |||
| 4fd1f9fed2 | |||
| 0d4de20f95 | |||
| fb17e6d8e2 | |||
| 22af16370d | |||
| 399fb22585 | |||
| 5072034da4 | |||
| 05ae6610c2 | |||
| 1e7a6f9221 | |||
| 952f3fee41 |
@@ -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,7 +2,9 @@
|
||||
|
||||
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)
|
||||
|
||||

|
||||
|
||||
|
||||

|
||||
|
||||
|
||||
## Installation
|
||||
@@ -13,7 +15,8 @@ Adds a tab to the webui that allows the user to generate a prompt from a small b
|
||||
## Usage
|
||||
|
||||
1. Write in the prompt in the *Start of the prompt* text box
|
||||
2. Click Generate and wait
|
||||
2. Select which model you want to use
|
||||
3. Click Generate and wait
|
||||
|
||||
The initial use of the model may take longer as it needs to be downloaded to your machine for offline use. The model will be used on your device and will be stored in the default location of `*username*/.cache/huggingface/hub/models`. The entire process of generating results will be done on your local machine and not require internet access.
|
||||
|
||||
@@ -25,11 +28,12 @@ The initial use of the model may take longer as it needs to be downloaded to you
|
||||
- **Max Length**: the maximum number of tokens for the output of the model
|
||||
- **Repetition Penalty**: The parameter for repetition penalty. 1.0 means no penalty. See [this paper](https://arxiv.org/pdf/1909.05858.pdf) for more details. Default setting is 1.2
|
||||
- **How Many To Generate**: The number of results to generate
|
||||
- **Use blacklist?** Using `.\extensions\stable-diffusion-webui-Prompt_Generator\blacklist.txt`. It will delete any matches to the generated result (case insensitive). Each item to be filtered out should be on a new line. *Be aware that it simply deletes it and doesn't generate more to make up for the lost words*
|
||||
- **Use blacklist?**: Using `.\extensions\stable-diffusion-webui-Prompt_Generator\blacklist.txt`. It will delete any matches to the generated result (case insensitive). Each item to be filtered out should be on a new line. *Be aware that it simply deletes it and doesn't generate more to make up for the lost words*
|
||||
- **Use punctuation**: Allows the use of commas in the output
|
||||
|
||||
## Models
|
||||
|
||||
There are two models provided:
|
||||
There are two 'default' models provided:
|
||||
|
||||
### FredZhang7
|
||||
|
||||
@@ -45,6 +49,12 @@ Useful to get more natural language prompts. Eg: "A cat sitting" -> "A cat sitti
|
||||
|
||||
*Be aware that sometimes the model fails to produce anything or less than the wanted amount, either try again or use a new prompt in that case*
|
||||
|
||||
## Install more models
|
||||
|
||||
To install more model to use, ensure that the models are hosted on [huggingface.co](https://huggingface.co) and edit the json file at `.\extensions\stable-diffusion-webui-Prompt_Generator\models.json` with the relevant information. Use the models in the file as an example
|
||||
|
||||
You might need to restart the extension/reload the UI if new items are added onto the list
|
||||
|
||||
## Credits
|
||||
|
||||
Credits to both [FredZhang7](https://huggingface.co/FredZhang7) and [Gustavosta](https://huggingface.co/Gustavosta)
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
//Basically copied this from AUTOMATIC1111 implementation of the main UI
|
||||
//Basically copied and adapted from AUTOMATIC1111 implementation of the main UI
|
||||
// mouseover tooltips for various UI elements in the form of "UI element label"="Tooltip text".
|
||||
|
||||
titles = {
|
||||
prompt_generator_titles = {
|
||||
"Temperature": "A higher temperature will produce more diverse results, but with a higher risk of less coherent text",
|
||||
"Max Length": "The maximum number of tokens for the output of the model",
|
||||
"Top K": "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.",
|
||||
@@ -11,16 +12,16 @@ titles = {
|
||||
|
||||
onUiUpdate(function(){
|
||||
gradioApp().querySelectorAll('span, button, select, p').forEach(function(span){
|
||||
tooltip = titles[span.textContent];
|
||||
tooltip = prompt_generator_titles[span.textContent];
|
||||
|
||||
if(!tooltip){
|
||||
tooltip = titles[span.value];
|
||||
tooltip = prompt_generator_titles[span.value];
|
||||
}
|
||||
|
||||
if(!tooltip){
|
||||
for (const c of span.classList) {
|
||||
if (c in titles) {
|
||||
tooltip = titles[c];
|
||||
if (c in prompt_generator_titles) {
|
||||
tooltip = prompt_generator_titles[c];
|
||||
break;
|
||||
}
|
||||
}
|
||||
@@ -35,7 +36,7 @@ onUiUpdate(function(){
|
||||
if (select.onchange != null) return;
|
||||
|
||||
select.onchange = function(){
|
||||
select.title = titles[select.value] || "";
|
||||
select.title = prompt_generator_titles[select.value] || "";
|
||||
}
|
||||
})
|
||||
})
|
||||
+12
@@ -0,0 +1,12 @@
|
||||
[
|
||||
{
|
||||
"Title":"Gustavosta",
|
||||
"Tokenizer":"gpt2",
|
||||
"Model":"Gustavosta/MagicPrompt-Dalle"
|
||||
},
|
||||
{
|
||||
"Title":"FredZhang7",
|
||||
"Tokenizer":"distilgpt2",
|
||||
"Model":"FredZhang7/distilgpt2-stable-diffusion-v2"
|
||||
}
|
||||
]
|
||||
+175
-136
@@ -10,23 +10,51 @@ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLI
|
||||
"""
|
||||
|
||||
|
||||
import json
|
||||
import re
|
||||
|
||||
import gradio as gr
|
||||
import modules
|
||||
from pathlib import Path
|
||||
from modules import script_callbacks
|
||||
from transformers import GPT2Tokenizer, GPT2LMHeadModel
|
||||
import re
|
||||
import math
|
||||
import modules.scripts as scripts
|
||||
from transformers import GPT2LMHeadModel, GPT2Tokenizer
|
||||
|
||||
result_prompt = ""
|
||||
models = {}
|
||||
max_no_results = 20 # TODO move to setting panel
|
||||
base_dir = scripts.basedir()
|
||||
model_file = Path(base_dir, "models.json")
|
||||
|
||||
|
||||
def add_to_prompt(num): # A function that determines which prompt to pass
|
||||
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}"
|
||||
class Model:
|
||||
'''
|
||||
Small strut to hold data for the text generator
|
||||
'''
|
||||
|
||||
def __init__(self, name, model, tokenizer) -> None:
|
||||
self.name = name
|
||||
self.model = model
|
||||
self.tokenizer = tokenizer
|
||||
pass
|
||||
|
||||
|
||||
def populate_models():
|
||||
"""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:
|
||||
data = json.load(f)
|
||||
for item in data:
|
||||
name = item["Title"]
|
||||
model = item["Model"]
|
||||
tokenizer = item["Tokenizer"]
|
||||
models[name] = Model(name, model, tokenizer)
|
||||
|
||||
|
||||
def add_to_prompt(prompt): # A holder TODO figure out how to get rid of it
|
||||
return prompt
|
||||
|
||||
|
||||
def get_list_blacklist():
|
||||
@@ -43,91 +71,68 @@ def get_list_blacklist():
|
||||
|
||||
def on_ui_tabs():
|
||||
# Method to create the extended prompt
|
||||
def generate_longer_generic(prompt, temperature, top_k,
|
||||
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
|
||||
|
||||
def generate_longer_prompt_gustavosta(prompt, temperature, top_k,
|
||||
max_length, repetition_penalty, num_return_sequences, use_blacklist=False, use_early_stop=True):
|
||||
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:
|
||||
tokenizer = GPT2Tokenizer.from_pretrained('gpt2')
|
||||
print("[Prompt_Generator]:","Loading Tokenizer")
|
||||
tokenizer = GPT2Tokenizer.from_pretrained(models[name].tokenizer)
|
||||
tokenizer.add_special_tokens({'pad_token': '[PAD]'})
|
||||
# Full credits for the model to Gustavosta (https://huggingface.co/Gustavosta). Under the MIT license
|
||||
model = GPT2LMHeadModel.from_pretrained(
|
||||
'Gustavosta/MagicPrompt-Dalle')
|
||||
print("[Prompt_Generator]:","Loading Model")
|
||||
model = GPT2LMHeadModel.from_pretrained(models[name].model)
|
||||
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}"
|
||||
try:
|
||||
min = len(prompt)
|
||||
print(f"Generate new prompt from: \"{prompt}\"")
|
||||
input_ids = tokenizer(prompt, return_tensors='pt').input_ids
|
||||
output = model.generate(input_ids, do_sample=True, temperature=temperature,
|
||||
top_k=round(top_k), max_length=max_length,
|
||||
num_return_sequences=num_return_sequences*4,
|
||||
repetition_penalty=float(repetition_penalty),
|
||||
penalty_alpha=0.6, no_repeat_ngram_size=1,
|
||||
early_stopping=use_early_stop)
|
||||
print("Generation complete!")
|
||||
tempString = ""
|
||||
if (use_blacklist):
|
||||
blacklist = get_list_blacklist()
|
||||
j = 0
|
||||
for i in range(len(output)):
|
||||
tempt_of_temp_String = tokenizer.decode(
|
||||
output[i], skip_special_tokens=True)
|
||||
if (len(tempt_of_temp_String) > min + 4):
|
||||
tempString += str(j+1) + ": " + tempt_of_temp_String
|
||||
j += 1
|
||||
else:
|
||||
continue
|
||||
if (use_blacklist):
|
||||
for to_check in blacklist:
|
||||
tempString = re.sub(
|
||||
to_check, "", tempString, flags=re.IGNORECASE)
|
||||
if (j == num_return_sequences):
|
||||
break
|
||||
|
||||
global result_prompt
|
||||
result_prompt = tempString
|
||||
|
||||
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:
|
||||
print(
|
||||
f"Exception encountered while attempting to generate prompt: {e}")
|
||||
return gr.update(), f"Error: {e}"
|
||||
|
||||
def generate_longer_prompt_FredZhang7(prompt, temperature, top_k,
|
||||
max_length, repetition_penalty, num_return_sequences, use_blacklist=False):
|
||||
try:
|
||||
tokenizer = GPT2Tokenizer.from_pretrained('distilgpt2')
|
||||
tokenizer.add_special_tokens({'pad_token': '[PAD]'})
|
||||
# Full credits for the model to FredZhang7 (https://huggingface.co/FredZhang7). Under creativeml-openrail-m license.
|
||||
model = GPT2LMHeadModel.from_pretrained(
|
||||
'FredZhang7/distilgpt2-stable-diffusion-v2')
|
||||
except Exception as e:
|
||||
print(f"Exception encountered while attempting to install tokenizer")
|
||||
return gr.update(), f"Error: {e}"
|
||||
try:
|
||||
print(f"Generate new prompt from: \"{prompt}\"")
|
||||
print("[Prompt_Generator]:",f"Generate new prompt from: \"{prompt}\" with {name}")
|
||||
input_ids = tokenizer(prompt, return_tensors='pt').input_ids
|
||||
if (use_punctuation):
|
||||
output = model.generate(input_ids, do_sample=True, temperature=temperature,
|
||||
top_k=round(top_k), max_length=max_length,
|
||||
num_return_sequences=num_return_sequences,
|
||||
repetition_penalty=float(repetition_penalty),
|
||||
repetition_penalty=float(
|
||||
repetition_penalty),
|
||||
early_stopping=True)
|
||||
else:
|
||||
output = model.generate(input_ids, do_sample=True, temperature=temperature,
|
||||
top_k=round(top_k), max_length=max_length,
|
||||
num_return_sequences=num_return_sequences,
|
||||
repetition_penalty=float(
|
||||
repetition_penalty),
|
||||
penalty_alpha=0.6, no_repeat_ngram_size=1,
|
||||
early_stopping=True)
|
||||
print("Generation complete!")
|
||||
print("[Prompt_Generator]:","Generation complete!")
|
||||
tempString = ""
|
||||
if (use_blacklist):
|
||||
blacklist = get_list_blacklist()
|
||||
for i in range(len(output)):
|
||||
|
||||
tempString += str(i+1)+": "+tokenizer.decode(
|
||||
tempString += tokenizer.decode(
|
||||
output[i], skip_special_tokens=True) + "\n"
|
||||
|
||||
if (use_blacklist):
|
||||
@@ -139,92 +144,126 @@ def on_ui_tabs():
|
||||
|
||||
result_prompt = tempString
|
||||
# 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:
|
||||
print(
|
||||
print("[Prompt_Generator]:",
|
||||
f"Exception encountered while attempting to generate prompt: {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
|
||||
txt2img_prompt = modules.ui.txt2img_paste_fields[0][0]
|
||||
img2img_prompt = modules.ui.img2img_paste_fields[0][0]
|
||||
|
||||
with gr.Blocks(analytics_enabled=False) as prompt_generator:
|
||||
# Handles UI for prompt creation
|
||||
with gr.Column():
|
||||
with gr.Row():
|
||||
promptTxt = gr.Textbox(
|
||||
prompt_textbox = gr.Textbox(
|
||||
lines=2, elem_id="promptTxt", label="Start of the prompt")
|
||||
with gr.Column():
|
||||
gr.HTML(
|
||||
"Mouse over the labels to access tooltips that provide explanations for the parameters.")
|
||||
with gr.Row():
|
||||
temp_slider = gr.Slider(
|
||||
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)
|
||||
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)
|
||||
with gr.Column():
|
||||
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)
|
||||
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)
|
||||
numReturnSequences_slider = gr.Slider(
|
||||
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.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>")
|
||||
with gr.Column():
|
||||
with gr.Row():
|
||||
generateButton_fred = gr.Button(
|
||||
value="Generate Using FredZhang7", elem_id="generate_button_FredZhang7")
|
||||
generateButton_magic = gr.Button(
|
||||
value="Generate Using Magic Prompt", elem_id="generate_button_MagicPrompt")
|
||||
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)
|
||||
populate_models()
|
||||
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?")
|
||||
generate_button = gr.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
|
||||
|
||||
# events
|
||||
generateButton_fred.click(fn=generate_longer_prompt_FredZhang7, inputs=[
|
||||
promptTxt, temp_slider, top_k_slider, max_length_slider,
|
||||
repetition_penalty_slider, num_return_sequences_slider,
|
||||
use_blacklist_checkbox],
|
||||
outputs=[results, send_to_img2img, send_to_txt2img, send_to_text,
|
||||
results_col, warning, promptNum_col])
|
||||
generateButton_magic.click(fn=generate_longer_prompt_gustavosta, inputs=[
|
||||
promptTxt, temp_slider, top_k_slider, max_length_slider,
|
||||
repetition_penalty_slider, num_return_sequences_slider,
|
||||
use_blacklist_checkbox],
|
||||
outputs=[results, send_to_img2img, send_to_txt2img, send_to_text,
|
||||
results_col, warning, promptNum_col])
|
||||
send_to_img2img.click(add_to_prompt, inputs=[
|
||||
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',
|
||||
# Handles UI for results
|
||||
results_vis = []
|
||||
results_txt_list = []
|
||||
with gr.Tab("Results"):
|
||||
with gr.Column():
|
||||
for i in range(max_no_results):
|
||||
with gr.Row(visible=False) as row:
|
||||
# Doesn't seem to do anything
|
||||
row.style(equal_height=True)
|
||||
with gr.Column(scale=3): # Guessing at the scale
|
||||
textBox = gr.Textbox(label="", lines=3)
|
||||
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)
|
||||
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)
|
||||
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"),
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user