rework prompt paste

This commit is contained in:
Vladimir Mandic
2023-11-08 11:33:40 -05:00
parent 8bcb14d0ac
commit a68f0bcbe4
5 changed files with 38 additions and 57 deletions
+23 -46
View File
@@ -3,20 +3,20 @@ import io
import os
import re
import json
from PIL import Image
import gradio as gr
from modules.paths import data_path
from modules import shared, ui_tempdir, script_callbacks, images
re_param_code = r'\s*([\w ]+):\s*("(?:\\"[^,]|\\"|\\|[^\"])+"|[^,]*)(?:,|$)'
re_param = re.compile(re_param_code)
re_imagesize = re.compile(r"^(\d+)x(\d+)$")
re_hypernet_hash = re.compile("\(([0-9a-f]+)\)$") # pylint: disable=anomalous-backslash-in-string
type_of_gr_update = type(gr.update())
paste_fields = {}
registered_param_bindings = []
debug = shared.log.info if os.environ.get('SD_PASTE_DEBUG', None) is not None else lambda *args, **kwargs: None
class ParamBinding:
@@ -203,37 +203,19 @@ def find_hypernetwork_key(hypernet_name, hypernet_hash=None):
def parse_generation_parameters(x: str):
"""parses generation parameters string, the one you see in text field under the picture in UI:
```
girl with an artist's beret, determined, blue eyes, desert scene, computer monitors, heavy makeup, by Alphonse Mucha and Charlie Bowater, ((eyeshadow)), (coquettish), detailed, intricate
Negative prompt: ugly, fat, obese, chubby, (((deformed))), [blurry], bad anatomy, disfigured, poorly drawn face, mutation, mutated, (extra_limb), (ugly), (poorly drawn hands), messy drawing
Steps: 20, Sampler: Euler a, CFG scale: 7, Seed: 965400086, Size: 512x512, Model hash: 45dee52b
```
returns a dict with field values
"""
if x is None:
return {}
res = {}
prompt = ""
negative_prompt = ""
done_with_prompt = False
*lines, lastline = x.strip().split("\n")
if len(re_param.findall(lastline)) < 3:
lines.append(lastline)
lastline = ''
for line in lines:
line = line.strip()
if line.startswith("Negative prompt:"):
done_with_prompt = True
line = line[16:].strip()
if done_with_prompt:
negative_prompt += ("" if negative_prompt == "" else "\n") + line
else:
prompt += ("" if prompt == "" else "\n") + line
res["Prompt"] = prompt
res["Negative prompt"] = negative_prompt
for k, v in re_param.findall(lastline):
if x is None:
return res
remaining = x.strip()
if len(remaining) == 0:
return res
remaining = x[7:] if x.startswith('Prompt: ') else x
res["Prompt"], remaining = remaining.split(' Negative prompt: ', maxsplit=1) if ' Negative prompt: ' in remaining else (remaining, '')
res["Negative prompt"], remaining = remaining.split(' Steps: ', maxsplit=1) if ' Steps: ' in remaining else (remaining, None)
if remaining is None:
return res
remaining = f'Steps: {remaining}'
for k, v in re_param.findall(remaining):
try:
if v[0] == '"' and v[-1] == '"':
v = unquote(v)
@@ -245,16 +227,9 @@ Steps: 20, Sampler: Euler a, CFG scale: 7, Seed: 965400086, Size: 512x512, Model
res[k] = v
except Exception:
pass
# Missing CLIP skip means it was set to 1 (the default)
if "Clip skip" not in res:
res["Clip skip"] = "1"
hypernet = res.get("Hypernet", None)
if hypernet is not None:
res["Prompt"] += f"""<hypernet:{hypernet}:{res.get("Hypernet strength", "1.0")}>"""
if "Hires resize-1" not in res:
res["Hires resize-1"] = 0
res["Hires resize-2"] = 0
res["Full quality"] = res.get('VAE', None) != 'TAESD'
for k, v in res.items():
debug(f"Parse prompt: '{k}'={v}")
return res
@@ -328,16 +303,16 @@ def create_override_settings_dict(text_pairs):
def connect_paste(button, local_paste_fields, input_comp, override_settings_component, tabname):
def paste_func(prompt):
if prompt is not None and 'Negative prompt' not in prompt and 'Steps' not in prompt:
prompt = None
if not prompt and not shared.cmd_opts.hide_ui_dir_config:
if prompt is None or len(prompt.strip()) == 0 and not shared.cmd_opts.hide_ui_dir_config:
filename = os.path.join(data_path, "params.txt")
if os.path.exists(filename):
with open(filename, "r", encoding="utf8") as file:
prompt = file.read()
shared.log.debug(f'Paste prompt last: {prompt}')
else:
prompt = ''
shared.log.debug(f'Paste prompt: {prompt}')
else:
shared.log.debug(f'Paste prompt current: {prompt}')
params = parse_generation_parameters(prompt)
script_callbacks.infotext_pasted_callback(prompt, params)
res = []
@@ -346,6 +321,8 @@ def connect_paste(button, local_paste_fields, input_comp, override_settings_comp
v = key(params)
else:
v = params.get(key, None)
if v is not None:
debug(f"Parse apply: '{key}'={v}")
if v is None:
res.append(gr.update())
elif isinstance(v, type_of_gr_update):
+9 -6
View File
@@ -588,7 +588,7 @@ def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=No
args["Denoising strength"] = p.denoising_strength
args["Latent sampler"] = p.latent_sampler
args["Image CFG scale"] = p.image_cfg_scale
args["CFG rescale"] = p.diffusers_guidance_rescale if shared.backend == shared.Backend.DIFFUSERS else None
args["CFG rescale"] = p.diffusers_guidance_rescale
if 'refine' in p.ops:
args["Second pass"] = p.enable_hr
args["Refiner"] = None if (not shared.opts.add_model_name_to_info) or (not shared.sd_refiner) or (not shared.sd_refiner.sd_checkpoint_info.model_name) else shared.sd_refiner.sd_checkpoint_info.model_name.replace(',', '').replace(':', '')
@@ -597,7 +597,7 @@ def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=No
args['Refiner start'] = p.refiner_start
args["Hires steps"] = p.hr_second_pass_steps
args["Latent sampler"] = p.latent_sampler
args["CFG rescale"] = p.diffusers_guidance_rescale if shared.backend == shared.Backend.DIFFUSERS else None
args["CFG rescale"] = p.diffusers_guidance_rescale
if 'img2img' in p.ops or 'inpaint' in p.ops:
args["Init image size"] = f"{getattr(p, 'init_img_width', 0)}x{getattr(p, 'init_img_height', 0)}"
args["Init image hash"] = getattr(p, 'init_img_hash', None)
@@ -756,7 +756,13 @@ def process_images(p: StableDiffusionProcessing) -> Processed:
return res
def validate_sample(sample):
def validate_sample(tensor):
if tensor.dtype == torch.bfloat16: # numpy does not support bf16
tensor = tensor.to(torch.float16)
if shared.backend == shared.Backend.ORIGINAL:
sample = 255.0 * np.moveaxis(tensor.cpu().numpy(), 0, 2)
else:
sample = 255. * tensor
with warnings.catch_warnings(record=True) as w:
cast = sample.astype(np.uint8)
if len(w) > 0:
@@ -914,7 +920,6 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed:
image = x_sample
x_sample = np.array(x_sample)
else:
x_sample = 255. * (np.moveaxis(x_sample.cpu().numpy(), 0, 2) if shared.backend == shared.Backend.ORIGINAL else x_sample)
x_sample = validate_sample(x_sample)
image = Image.fromarray(x_sample)
if p.restore_faces:
@@ -1118,7 +1123,6 @@ class StableDiffusionProcessingTxt2Img(StableDiffusionProcessing):
decoded_samples = decode_first_stage(self.sd_model, samples.to(dtype=devices.dtype_vae), self.full_quality)
decoded_samples = torch.clamp((decoded_samples + 1.0) / 2.0, min=0.0, max=1.0)
for i, x_sample in enumerate(decoded_samples):
x_sample = 255. * np.moveaxis(x_sample.cpu().numpy(), 0, 2)
x_sample = validate_sample(x_sample)
image = Image.fromarray(x_sample)
bak_extra_generation_params, bak_restore_faces = self.extra_generation_params, self.restore_faces
@@ -1134,7 +1138,6 @@ class StableDiffusionProcessingTxt2Img(StableDiffusionProcessing):
decoded_samples = torch.clamp((decoded_samples + 1.0) / 2.0, min=0.0, max=1.0)
batch_images = []
for _i, x_sample in enumerate(decoded_samples):
x_sample = 255. * np.moveaxis(x_sample.cpu().numpy(), 0, 2)
x_sample = validate_sample(x_sample)
image = Image.fromarray(x_sample)
image = images.resize_image(1, image, target_width, target_height, upscaler_name=self.hr_upscaler)
+1 -1
View File
@@ -268,7 +268,7 @@ def select_checkpoint(op='model'):
shared.log.info(f'Select: {op}="{checkpoint_info.title if checkpoint_info is not None else None}"')
return checkpoint_info
if len(checkpoints_list) == 0 and not shared.cmd_opts.no_download:
shared.log.error("Cannot generate without a checkpoint")
shared.log.warning("Cannot generate without a checkpoint")
shared.log.info("Set system paths to use existing folders in a different location")
shared.log.info("Or use --ckpt <path-to-checkpoint> to force using existing checkpoint")
return None
+4 -3
View File
@@ -513,7 +513,7 @@ def create_ui(startup_timer = None):
txt2img_paste_fields = [
(txt2img_prompt, "Prompt"),
(txt2img_negative_prompt, "Negative prompt"),
# (txt2img_prompt_styles, "Styles"),
(txt2img_prompt_styles, "Styles"),
(steps, "Steps"),
(seed, "Seed"),
(sampler_index, "Sampler"),
@@ -530,8 +530,8 @@ def create_ui(startup_timer = None):
(refiner_start, "Refiner start"),
(full_quality, "Full quality"),
(restore_faces, "Face restoration"),
(batch_size, "Batch size"),
(batch_count, "Batch count"),
(batch_count, "Batch-1"),
(batch_size, "Batch-2"),
(seed_resize_from_w, "Seed resize from-1"),
(seed_resize_from_h, "Seed resize from-2"),
(enable_hr, "Second pass"),
@@ -548,6 +548,7 @@ def create_ui(startup_timer = None):
(tiling, "Tiling"),
(refiner_negative, "Negative2"),
(refiner_prompt, "Prompt2"),
# TODO restore params complete list
*modules.scripts.scripts_txt2img.infotext_fields
]
parameters_copypaste.add_paste_fields("txt2img", None, txt2img_paste_fields, override_settings)