Compare commits

..

2 Commits

Author SHA1 Message Date
Imrayya 31ec08338b Updated print statement so its prettier 2023-05-08 13:29:45 +02:00
Imrayya 89215f828d Enabled support for export to file #23. 2023-05-05 10:52:55 +02:00
+53 -32
View File
@@ -12,7 +12,6 @@ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLI
import json import json
import re import re
import sys
import gradio as gr import gradio as gr
import modules import modules
@@ -24,7 +23,9 @@ from transformers import GPT2LMHeadModel, GPT2Tokenizer
result_prompt = "" result_prompt = ""
models = {} models = {}
max_no_results = 20 # TODO move to setting panel max_no_results = 20 # TODO move to setting panel
model_file = Path(scripts.basedir(), "models.json") base_dir = scripts.basedir()
model_file = Path(base_dir, "models.json")
class Model: class Model:
''' '''
@@ -73,7 +74,7 @@ def on_ui_tabs():
def generate_longer_generic(prompt, temperature, top_k, def generate_longer_generic(prompt, temperature, top_k,
max_length, repetition_penalty, max_length, repetition_penalty,
num_return_sequences, name, use_punctuation=False, num_return_sequences, name, use_punctuation=False,
use_blacklist=False, progress=gr.Progress()): # TODO make the progress bar work use_blacklist=False): # TODO make the progress bar work
"""Generates a longer string from the input """Generates a longer string from the input
Args: Args:
@@ -98,21 +99,17 @@ def on_ui_tabs():
Returns: Returns:
Returns only an error otherwise saves it in result_prompt Returns only an error otherwise saves it in result_prompt
""" """
progress(0, "Starting")
try: try:
progress(0.25) print("[Prompt_Generator]:","Loading Tokenizer")
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]'})
progress(0.5) print("[Prompt_Generator]:","Loading Model")
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("[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}")
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,
@@ -129,8 +126,7 @@ 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!")
progress(1, "Done!")
tempString = "" tempString = ""
if (use_blacklist): if (use_blacklist):
blacklist = get_list_blacklist() blacklist = get_list_blacklist()
@@ -149,7 +145,7 @@ def on_ui_tabs():
result_prompt = tempString result_prompt = tempString
# print(result_prompt) # print(result_prompt)
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}"
@@ -171,15 +167,27 @@ def on_ui_tabs():
result_list.append("") result_list.append("")
return result_list 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(
@@ -187,19 +195,19 @@ 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=max_no_results, 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():
@@ -207,22 +215,23 @@ def on_ui_tabs():
generate_dropdown = gr.Dropdown(choices=list(models.keys()), value=list(models.keys())[ 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 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 = gr.Button( 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 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
# Handles Dynamic results # Handles UI for results
results_vis = [] results_vis = []
results_txt_list = [] results_txt_list = []
with gr.Column() as results_col: with gr.Tab("Results"):
with gr.Column():
for i in range(max_no_results): for i in range(max_no_results):
with gr.Row(visible=False) as row: with gr.Row(visible=False) as row:
row.style(equal_height=True) # Doesn't seem to do anything # Doesn't seem to do anything
row.style(equal_height=True)
with gr.Column(scale=3): # Guessing at the scale with gr.Column(scale=3): # Guessing at the scale
textBox = gr.Textbox(label="") textBox = gr.Textbox(label="", lines=3)
with gr.Column(scale=1): with gr.Column(scale=1):
txt2img = gr.Button("send to txt2img") txt2img = gr.Button("send to txt2img")
img2img = gr.Button("send to img2img") img2img = gr.Button("send to img2img")
# Handles ___2img buttons # Handles ___2img buttons
txt2img.click(add_to_prompt, inputs=[ txt2img.click(add_to_prompt, inputs=[
textBox], outputs=[txt2img_prompt]).then(None, _js='switch_to_txt2img', textBox], outputs=[txt2img_prompt]).then(None, _js='switch_to_txt2img',
@@ -232,17 +241,29 @@ def on_ui_tabs():
inputs=None, outputs=None) inputs=None, outputs=None)
results_txt_list.append(textBox) results_txt_list.append(textBox)
results_vis.append(row) 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 # 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 # Please note that we use `.then()` to run other ui elements after the generation is done
generateButton.click(fn=generate_longer_generic, inputs=[ generate_button.click(fn=generate_longer_generic, inputs=[
promptTxt, temp_slider, top_k_slider, max_length_slider, prompt_textbox, temp_slider, topK_slider, maxLength_slider,
repetition_penalty_slider, num_return_sequences_slider, repetitionPenalty_slider, numReturnSequences_slider,
generate_dropdown, use_punctuation_check, use_blacklist_checkbox]).then( generate_dropdown, use_punctuation_check, useBlacklist_checkbox]).then(
fn=ui_dynamic_result_visible, inputs=num_return_sequences_slider, fn=ui_dynamic_result_visible, inputs=numReturnSequences_slider,
outputs=results_vis).then( outputs=results_vis).then(
fn=ui_dynamic_result_prompts, outputs=results_txt_list) 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"),