From a22765274b7f921c48e4f78b00d144bdb1beec23 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 1 Feb 2024 14:00:45 -0500 Subject: [PATCH] major refactor: processing, taesd, scripts --- javascript/sdnext.css | 8 +- modules/control/run.py | 47 +- modules/processing.py | 1086 +-------------------------- modules/processing_class.py | 487 ++++++++++++ modules/processing_diffusers.py | 116 +-- modules/processing_helpers.py | 428 +++++++++++ modules/processing_info.py | 143 ++++ modules/processing_original.py | 163 ++++ modules/processing_vae.py | 3 +- modules/scripts.py | 32 +- modules/sd_samplers_common.py | 3 +- modules/{taesd => }/sd_vae_taesd.py | 95 ++- modules/taesd/taesd.py | 95 --- modules/txt2img.py | 7 +- requirements.txt | 2 +- 15 files changed, 1381 insertions(+), 1334 deletions(-) create mode 100644 modules/processing_class.py create mode 100644 modules/processing_helpers.py create mode 100644 modules/processing_info.py create mode 100644 modules/processing_original.py rename modules/{taesd => }/sd_vae_taesd.py (50%) delete mode 100644 modules/taesd/taesd.py diff --git a/javascript/sdnext.css b/javascript/sdnext.css index 8236450cc..e4865e03a 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -33,12 +33,10 @@ button { font-size: var(--text-lg) !important; } .gradio-container { max-width: unset !important; padding: var(--block-label-padding) !important; } .gradio-container .prose a, .gradio-container .prose a:visited{ color: unset; text-decoration: none; } .gradio-dropdown { margin-right: var(--spacing-sm) !important; min-width:160px; max-width:fit-content } -.gradio-dropdown ul.options { z-index: 1000; min-width: fit-content; max-height: 33vh !important; white-space: nowrap; } +.gradio-dropdown ul.options { z-index: 1000; min-width: fit-content; max-height: 50vh !important; white-space: nowrap; } .gradio-dropdown ul.options li.item { padding: var(--spacing-xs); } .gradio-dropdown ul.options li.item:not(:has(.hide)) { background-color: var(--primary-500); } -.gradio-dropdown .token { padding: var(--spacing-xs); } -.gradio-dropdown span { margin-bottom: 0 !important; font-size: var(--text-sm); } -.gradio-dropdown .reference { margin-bottom: var(--spacing-sm) !important; } +.gradio-dropdown .token { padding: var(--spacing-xs) !important; } .gradio-html { color: var(--body-text-color); } .gradio-html .min { min-height: 0; } .gradio-html div.wrap { height: 100%; } @@ -83,7 +81,7 @@ button.custom-button{ border-radius: var(--button-large-radius); padding: var(-- #txt2img_footer, #img2img_footer, #control_footer { height: fit-content; display: none; } #txt2img_generate_box, #img2img_generate_box, #control_general_box { gap: 0.5em; flex-wrap: wrap-reverse; height: fit-content; } #txt2img_actions_column, #img2img_actions_column, #control_actions_column { gap: 0.3em; height: fit-content; } -#txt2img_generate_box>button, #img2img_generate_box>button, #control_generate_box>button, #txt2img_enqueue, #img2img_enqueue { min-height: 42px; max-height: 42px; line-height: 1em; } +#txt2img_generate_box>button, #img2img_generate_box>button, #control_generate_box>button, #txt2img_enqueue, #img2img_enqueue { min-height: 44px; max-height: 44px; line-height: 1em; } #txt2img_generate_line2, #img2img_generate_line2, #txt2img_tools, #img2img_tools, #control_generate_line2, #control_tools { display: flex; } #txt2img_generate_line2>button, #img2img_generate_line2>button, #extras_generate_box>button, #control_generate_line2>button, #txt2img_tools>button, #img2img_tools>button, #control_tools>button { height: 2em; line-height: 0; font-size: var(--text-md); min-width: unset; display: block !important; } diff --git a/modules/control/run.py b/modules/control/run.py index 9a52576f2..12a301008 100644 --- a/modules/control/run.py +++ b/modules/control/run.py @@ -1,6 +1,5 @@ import os import time -import math from typing import List, Union import cv2 import numpy as np @@ -14,6 +13,7 @@ from modules.control.units import lite # Kohya ControlLLLite from modules.control.units import t2iadapter # TencentARC T2I-Adapter from modules.control.units import reference # ControlNet-Reference from modules import devices, shared, errors, processing, images, sd_models, scripts, masking +from modules.processing_class import StableDiffusionProcessingControl debug = shared.log.trace if os.environ.get('SD_CONTROL_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -22,49 +22,6 @@ pipe = None original_pipeline = None -class ControlProcessing(processing.StableDiffusionProcessingImg2Img): - def __init__(self, **kwargs): - super().__init__(**kwargs) - self.strength = None - self.adapter_conditioning_scale = None - self.adapter_conditioning_factor = None - self.guess_mode = None - self.controlnet_conditioning_scale = None - self.control_guidance_start = None - self.control_guidance_end = None - self.reference_attn = None - self.reference_adain = None - self.attention_auto_machine_weight = None - self.gn_auto_machine_weight = None - self.style_fidelity = None - self.ref_image = None - self.image = None - self.query_weight = None - self.adain_weight = None - self.adapter_conditioning_factor = 1.0 - self.attention = 'Attention' - self.fidelity = 0.5 - self.override = None - self.ip_adapter_name = None - self.ip_adapter_scale = 1.0 - self.ip_adapter_image = None - - def sample(self, conditioning, unconditional_conditioning, seeds, subseeds, subseed_strength, prompts): # abstract - pass - - def init_hr(self): - if self.resize_name == 'None' or self.scale_by == 1.0: - return - self.is_hr_pass = True - self.hr_force = True - self.hr_upscaler = self.resize_name - self.hr_upscale_to_x, self.hr_upscale_to_y = int(self.width * self.scale_by), int(self.height * self.scale_by) - self.hr_upscale_to_x, self.hr_upscale_to_y = 8 * math.ceil(self.hr_upscale_to_x / 8), 8 * math.ceil(self.hr_upscale_to_y / 8) - # hypertile_set(self, hr=True) - shared.state.job_count = 2 * self.n_iter - shared.log.debug(f'Control hires: upscaler="{self.hr_upscaler}" upscale={self.scale_by} size={self.hr_upscale_to_x}x{self.hr_upscale_to_y}') - - def restore_pipeline(): global pipe # pylint: disable=global-statement pipe = None @@ -100,7 +57,7 @@ def control_run(units: List[unit.Unit], inputs, inits, mask, unit_type: str, is_ if mask is not None and input_type == 0: input_type = 1 # inpaint always requires control_image - p = ControlProcessing( + p = StableDiffusionProcessingControl( prompt = prompt, negative_prompt = negative, styles = styles, diff --git a/modules/processing.py b/modules/processing.py index 7b1c2bb8b..12f6ff2a9 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -1,377 +1,32 @@ import os import json -import math import time -import hashlib -import random -import warnings from contextlib import nullcontext -from typing import Any, Dict, List -from dataclasses import dataclass, field -import torch import numpy as np -import cv2 -from PIL import Image, ImageOps -from skimage import exposure -from einops import repeat, rearrange -from blendmodes.blend import blendLayers, BlendType -from installer import git_commit -from modules import shared, devices, errors, images, scripts, memstats, lowvram, masking, prompt_parser, script_callbacks, extra_networks, face_restoration, sd_hijack_freeu, sd_samplers, sd_samplers_common, sd_models, sd_vae, generation_parameters_copypaste -from modules.taesd import sd_vae_taesd -from modules.sd_hijack_hypertile import context_hypertile_vae, context_hypertile_unet, hypertile_set +from PIL import Image +from modules import shared, devices, errors, images, scripts, memstats, lowvram, script_callbacks, extra_networks, face_restoration, sd_hijack_freeu, sd_models, sd_vae, processing_helpers +from modules.sd_hijack_hypertile import context_hypertile_vae, context_hypertile_unet +from modules.processing_class import StableDiffusionProcessing, StableDiffusionProcessingTxt2Img, StableDiffusionProcessingImg2Img # pylint: disable=unused-import +from modules.processing_info import create_infotext -if shared.backend == shared.Backend.ORIGINAL: - from modules import sd_hijack -else: - sd_hijack = None - opt_C = 4 opt_f = 8 debug = shared.log.trace if os.environ.get('SD_PROCESS_DEBUG', None) is not None else lambda *args, **kwargs: None debug('Trace: PROCESS') - - -def setup_color_correction(image): - debug("Calibrating color correction") - correction_target = cv2.cvtColor(np.asarray(image.copy()), cv2.COLOR_RGB2LAB) - return correction_target - - -def apply_color_correction(correction, original_image): - shared.log.debug(f"Applying color correction: correction={correction.shape} image={original_image}") - np_image = np.asarray(original_image) - np_recolor = cv2.cvtColor(np_image, cv2.COLOR_RGB2LAB) - np_match = exposure.match_histograms(np_recolor, correction, channel_axis=2) - np_output = cv2.cvtColor(np_match, cv2.COLOR_LAB2RGB) - image = Image.fromarray(np_output.astype("uint8")) - image = blendLayers(image, original_image, BlendType.LUMINOSITY) - return image - - -def apply_overlay(image: Image, paste_loc, index, overlays): - debug(f'Apply overlay: image={image} loc={paste_loc} index={index} overlays={overlays}') - if overlays is None or index >= len(overlays): - return image - overlay = overlays[index] - if paste_loc is not None: - x, y, w, h = paste_loc - if image.width != w or image.height != h or x != 0 or y != 0: - base_image = Image.new('RGBA', (overlay.width, overlay.height)) - image = images.resize_image(2, image, w, h) - base_image.paste(image, (x, y)) - image = base_image - image = image.convert('RGBA') - image.alpha_composite(overlay) - image = image.convert('RGB') - return image - - -def create_binary_mask(image): - if image.mode == 'RGBA' and image.getextrema()[-1] != (255, 255): - image = image.split()[-1].convert("L").point(lambda x: 255 if x > 128 else 0) - else: - image = image.convert('L') - return image - - -def images_tensor_to_samples(image, approximation=None, model=None): # pylint: disable=unused-argument - if model is None: - model = shared.sd_model - model.first_stage_model.to(devices.dtype_vae) - image = image.to(shared.device, dtype=devices.dtype_vae) - image = image * 2 - 1 - if len(image) > 1: - x_latent = torch.stack([ - model.get_first_stage_encoding(model.encode_first_stage(torch.unsqueeze(img, 0)))[0] - for img in image - ]) - else: - x_latent = model.get_first_stage_encoding(model.encode_first_stage(image)) - return x_latent - - -def txt2img_image_conditioning(sd_model, x, width, height): - if sd_model.model.conditioning_key in {'hybrid', 'concat'}: # Inpainting models - # The "masked-image" in this case will just be all zeros since the entire image is masked. - image_conditioning = torch.zeros(x.shape[0], 3, height, width, device=x.device) - image_conditioning = sd_model.get_first_stage_encoding(sd_model.encode_first_stage(image_conditioning)) - # Add the fake full 1s mask to the first dimension. - image_conditioning = torch.nn.functional.pad(image_conditioning, (0, 0, 0, 0, 1, 0), value=1.0) # pylint: disable=not-callable - image_conditioning = image_conditioning.to(x.dtype) - return image_conditioning - elif sd_model.model.conditioning_key == "crossattn-adm": # UnCLIP models - return x.new_zeros(x.shape[0], 2*sd_model.noise_augmentor.time_embed.dim, dtype=x.dtype, device=x.device) - else: - # Dummy zero conditioning if we're not using inpainting or unclip models. - # Still takes up a bit of memory, but no encoder call. - # Pretty sure we can just make this a 1x1 image since its not going to be used besides its batch size. - return x.new_zeros(x.shape[0], 5, 1, 1, dtype=x.dtype, device=x.device) - - -def get_sampler_name(sampler_index: int, img: bool = False) -> str: - sampler_index = sampler_index or 0 - if len(sd_samplers.samplers) > sampler_index: - sampler_name = sd_samplers.samplers[sampler_index].name - else: - sampler_name = "UniPC" - shared.log.warning(f'Sampler not found: index={sampler_index} available={[s.name for s in sd_samplers.samplers]} fallback={sampler_name}') - if img and sampler_name == "PLMS": - sampler_name = "UniPC" - shared.log.warning(f'Sampler not compatible: name=PLMS fallback={sampler_name}') - return sampler_name - - -@dataclass(repr=False) -class StableDiffusionProcessing: - """ - The first set of paramaters: sd_models -> do_not_reload_embeddings represent the minimum required to create a StableDiffusionProcessing - """ - def __init__(self, sd_model=None, outpath_samples=None, outpath_grids=None, prompt: str = "", styles: List[str] = None, seed: int = -1, subseed: int = -1, subseed_strength: float = 0, seed_resize_from_h: int = -1, seed_resize_from_w: int = -1, seed_enable_extras: bool = True, sampler_name: str = None, hr_sampler_name: str = None, batch_size: int = 1, n_iter: int = 1, steps: int = 50, cfg_scale: float = 7.0, image_cfg_scale: float = None, clip_skip: int = 1, width: int = 512, height: int = 512, full_quality: bool = True, restore_faces: bool = False, tiling: bool = False, do_not_save_samples: bool = False, do_not_save_grid: bool = False, extra_generation_params: Dict[Any, Any] = None, overlay_images: Any = None, negative_prompt: str = None, eta: float = None, do_not_reload_embeddings: bool = False, denoising_strength: float = 0, diffusers_guidance_rescale: float = 0.7, sag_scale: float = 0.0, resize_mode: int = 0, resize_name: str = 'None', scale_by: float = 0, selected_scale_tab: int = 0, hdr_clamp: bool = False, hdr_boundary: float = 4.0, hdr_threshold: float = 3.5, hdr_center: bool = False, hdr_channel_shift: float = 0.8, hdr_full_shift: float = 0.8, hdr_maximize: bool = False, hdr_max_center: float = 0.6, hdr_max_boundry: float = 1.0, override_settings: Dict[str, Any] = None, override_settings_restore_afterwards: bool = True, sampler_index: int = None, script_args: list = None): # pylint: disable=unused-argument - self.outpath_samples: str = outpath_samples - self.outpath_grids: str = outpath_grids - self.prompt: str = prompt - self.prompt_for_display: str = None - self.negative_prompt: str = (negative_prompt or "") - self.styles: list = styles or [] - self.seed: int = seed - self.subseed: int = subseed - self.subseed_strength: float = subseed_strength - self.seed_resize_from_h: int = seed_resize_from_h - self.seed_resize_from_w: int = seed_resize_from_w - self.sampler_name: str = sampler_name - self.hr_sampler_name: str = hr_sampler_name - self.batch_size: int = batch_size - self.n_iter: int = n_iter - self.steps: int = steps - self.hr_second_pass_steps = 0 - self.cfg_scale: float = cfg_scale - self.scale_by: float = scale_by - self.image_cfg_scale = image_cfg_scale - self.diffusers_guidance_rescale = diffusers_guidance_rescale - self.sag_scale = sag_scale - if devices.backend == "ipex" and width == 1024 and height == 1024 and not torch.xpu.has_fp64_dtype() and os.environ.get('DISABLE_IPEX_1024_WA', None) is None: - width = 1080 - height = 1080 - self.width: int = width - self.height: int = height - self.full_quality: bool = full_quality - self.restore_faces: bool = restore_faces - self.tiling: bool = tiling - self.do_not_save_samples: bool = do_not_save_samples - self.do_not_save_grid: bool = do_not_save_grid - self.extra_generation_params: dict = extra_generation_params or {} - self.overlay_images = overlay_images - self.eta = eta - self.do_not_reload_embeddings = do_not_reload_embeddings - self.paste_to = None - self.color_corrections = None - self.denoising_strength: float = denoising_strength - self.override_settings = {k: v for k, v in (override_settings or {}).items() if k not in shared.restricted_opts} - self.override_settings_restore_afterwards = override_settings_restore_afterwards - self.is_using_inpainting_conditioning = False - self.disable_extra_networks = False - self.token_merging_ratio = 0 - self.token_merging_ratio_hr = 0 - # self.scripts = scripts.ScriptRunner() # set via property - # self.script_args = script_args or [] # set via property - self.per_script_args = {} - self.all_prompts = None - self.all_negative_prompts = None - self.all_seeds = None - self.all_subseeds = None - self.clip_skip = clip_skip - self.iteration = 0 - self.is_control = False - self.is_hr_pass = False - self.is_refiner_pass = False - self.hr_force = False - self.enable_hr = None - self.hr_scale = None - self.hr_upscaler = None - self.hr_resize_x = 0 - self.hr_resize_y = 0 - self.hr_upscale_to_x = 0 - self.hr_upscale_to_y = 0 - self.truncate_x = 0 - self.truncate_y = 0 - self.applied_old_hires_behavior_to = None - self.refiner_steps = 5 - self.refiner_start = 0 - self.refiner_prompt = '' - self.refiner_negative = '' - self.ops = [] - self.resize_mode: int = resize_mode - self.resize_name: str = resize_name - self.ddim_discretize = shared.opts.ddim_discretize - self.s_min_uncond = shared.opts.s_min_uncond - self.s_churn = shared.opts.s_churn - self.s_noise = shared.opts.s_noise - self.s_min = shared.opts.s_min - self.s_max = shared.opts.s_max - self.s_tmin = shared.opts.s_tmin - self.s_tmax = float('inf') # not representable as a standard ui option - shared.opts.data['clip_skip'] = clip_skip - self.task_args = {} - # a1111 compatibility items - self.refiner_switch_at = 0 - self.hr_prompt = '' - self.all_hr_prompts = [] - self.hr_negative_prompt = '' - self.all_hr_negative_prompts = [] - self.comments = {} - self.is_api = False - self.scripts_value: scripts.ScriptRunner = field(default=None, init=False) - self.script_args_value: list = field(default=None, init=False) - self.scripts_setup_complete: bool = field(default=False, init=False) - # hdr - self.hdr_clamp = hdr_clamp - self.hdr_boundary = hdr_boundary - self.hdr_threshold = hdr_threshold - self.hdr_center = hdr_center - self.hdr_channel_shift = hdr_channel_shift - self.hdr_full_shift = hdr_full_shift - self.hdr_maximize = hdr_maximize - self.hdr_max_center = hdr_max_center - self.hdr_max_boundry = hdr_max_boundry - self.scheduled_prompt: bool = False - self.prompt_embeds = [] - self.positive_pooleds = [] - self.negative_embeds = [] - self.negative_pooleds = [] - - - @property - def sd_model(self): - return shared.sd_model - - @property - def scripts(self): - return self.scripts_value - - @scripts.setter - def scripts(self, value): - self.scripts_value = value - if self.scripts_value and self.script_args_value and not self.scripts_setup_complete: - self.setup_scripts() - - @property - def script_args(self): - return self.script_args_value - - @script_args.setter - def script_args(self, value): - self.script_args_value = value - if self.scripts_value and self.script_args_value and not self.scripts_setup_complete: - self.setup_scripts() - - def setup_scripts(self): - self.scripts_setup_complete = True - self.scripts.setup_scrips(self, is_ui=not self.is_api) - - def comment(self, text): - self.comments[text] = 1 - - def txt2img_image_conditioning(self, x, width=None, height=None): - self.is_using_inpainting_conditioning = self.sd_model.model.conditioning_key in {'hybrid', 'concat'} - return txt2img_image_conditioning(self.sd_model, x, width or self.width, height or self.height) - - def depth2img_image_conditioning(self, source_image): - # Use the AddMiDaS helper to Format our source image to suit the MiDaS model - from ldm.data.util import AddMiDaS - transformer = AddMiDaS(model_type="dpt_hybrid") - transformed = transformer({"jpg": rearrange(source_image[0], "c h w -> h w c")}) - midas_in = torch.from_numpy(transformed["midas_in"][None, ...]).to(device=shared.device) - midas_in = repeat(midas_in, "1 ... -> n ...", n=self.batch_size) - conditioning_image = self.sd_model.get_first_stage_encoding(self.sd_model.encode_first_stage(source_image)) - conditioning = torch.nn.functional.interpolate( - self.sd_model.depth_model(midas_in), - size=conditioning_image.shape[2:], - mode="bicubic", - align_corners=False, - ) - (depth_min, depth_max) = torch.aminmax(conditioning) - conditioning = 2. * (conditioning - depth_min) / (depth_max - depth_min) - 1. - return conditioning - - def edit_image_conditioning(self, source_image): - conditioning_image = self.sd_model.encode_first_stage(source_image).mode() - return conditioning_image - - def unclip_image_conditioning(self, source_image): - c_adm = self.sd_model.embedder(source_image) - if self.sd_model.noise_augmentor is not None: - noise_level = 0 - c_adm, noise_level_emb = self.sd_model.noise_augmentor(c_adm, noise_level=repeat(torch.tensor([noise_level]).to(c_adm.device), '1 -> b', b=c_adm.shape[0])) - c_adm = torch.cat((c_adm, noise_level_emb), 1) - return c_adm - - def inpainting_image_conditioning(self, source_image, latent_image, image_mask=None): - self.is_using_inpainting_conditioning = True - # Handle the different mask inputs - if image_mask is not None: - if torch.is_tensor(image_mask): - conditioning_mask = image_mask - else: - conditioning_mask = np.array(image_mask.convert("L")) - conditioning_mask = conditioning_mask.astype(np.float32) / 255.0 - conditioning_mask = torch.from_numpy(conditioning_mask[None, None]) - # Inpainting model uses a discretized mask as input, so we round to either 1.0 or 0.0 - conditioning_mask = torch.round(conditioning_mask) - else: - conditioning_mask = source_image.new_ones(1, 1, *source_image.shape[-2:]) - # Create another latent image, this time with a masked version of the original input. - # Smoothly interpolate between the masked and unmasked latent conditioning image using a parameter. - conditioning_mask = conditioning_mask.to(device=source_image.device, dtype=source_image.dtype) - conditioning_image = torch.lerp( - source_image, - source_image * (1.0 - conditioning_mask), - getattr(self, "inpainting_mask_weight", shared.opts.inpainting_mask_weight) - ) - # Encode the new masked image using first stage of network. - conditioning_image = self.sd_model.get_first_stage_encoding(self.sd_model.encode_first_stage(conditioning_image)) - # Create the concatenated conditioning tensor to be fed to `c_concat` - conditioning_mask = torch.nn.functional.interpolate(conditioning_mask, size=latent_image.shape[-2:]) - conditioning_mask = conditioning_mask.expand(conditioning_image.shape[0], -1, -1, -1) - image_conditioning = torch.cat([conditioning_mask, conditioning_image], dim=1) - image_conditioning = image_conditioning.to(device=shared.device, dtype=source_image.dtype) - return image_conditioning - - def diffusers_image_conditioning(self, _source_image, latent_image, _image_mask=None): - # shared.log.warning('Diffusers not implemented: img2img_image_conditioning') - return latent_image.new_zeros(latent_image.shape[0], 5, 1, 1) - - def img2img_image_conditioning(self, source_image, latent_image, image_mask=None): - from ldm.models.diffusion.ddpm import LatentDepth2ImageDiffusion - source_image = devices.cond_cast_float(source_image) - # HACK: Using introspection as the Depth2Image model doesn't appear to uniquely - # identify itself with a field common to all models. The conditioning_key is also hybrid. - if shared.backend == shared.Backend.DIFFUSERS: - return self.diffusers_image_conditioning(source_image, latent_image, image_mask) - if isinstance(self.sd_model, LatentDepth2ImageDiffusion): - return self.depth2img_image_conditioning(source_image) - if hasattr(self.sd_model, 'cond_stage_key') and self.sd_model.cond_stage_key == "edit": - return self.edit_image_conditioning(source_image) - if hasattr(self.sampler, 'conditioning_key') and self.sampler.conditioning_key in {'hybrid', 'concat'}: - return self.inpainting_image_conditioning(source_image, latent_image, image_mask=image_mask) - if hasattr(self.sampler, 'conditioning_key') and self.sampler.conditioning_key == "crossattn-adm": - return self.unclip_image_conditioning(source_image) - # Dummy zero conditioning if we're not using inpainting or depth model. - return latent_image.new_zeros(latent_image.shape[0], 5, 1, 1) - - def init(self, all_prompts, all_seeds, all_subseeds): - pass - - def sample(self, conditioning, unconditional_conditioning, seeds, subseeds, subseed_strength, prompts): - raise NotImplementedError - - def close(self): - self.sampler = None # pylint: disable=attribute-defined-outside-init - - def get_token_merging_ratio(self, for_hr=False): - if for_hr: - return self.token_merging_ratio_hr or shared.opts.token_merging_ratio_hr or self.token_merging_ratio or shared.opts.token_merging_ratio - return self.token_merging_ratio or shared.opts.token_merging_ratio +create_binary_mask = processing_helpers.create_binary_mask +apply_overlay = processing_helpers.apply_overlay +apply_color_correction = processing_helpers.apply_color_correction +setup_color_correction = processing_helpers.setup_color_correction +txt2img_image_conditioning = processing_helpers.txt2img_image_conditioning +img2img_image_conditioning = processing_helpers.img2img_image_conditioning +get_fixed_seed = processing_helpers.get_fixed_seed +create_random_tensors = processing_helpers.create_random_tensors +decode_first_stage = processing_helpers.decode_first_stage +old_hires_fix_first_pass_dimensions = processing_helpers.old_hires_fix_first_pass_dimensions +validate_sample = processing_helpers.validate_sample +get_sampler_name = processing_helpers.get_sampler_name +images_tensor_to_samples = processing_helpers.images_tensor_to_samples class Processed: @@ -451,7 +106,6 @@ class Processed: "styles": self.styles, "job_timestamp": self.job_timestamp, "clip_skip": self.clip_skip, - # "is_using_inpainting_conditioning": self.is_using_inpainting_conditioning, } return json.dumps(obj) @@ -462,241 +116,6 @@ class Processed: return self.token_merging_ratio_hr if for_hr else self.token_merging_ratio -def slerp(val, low, high): # from https://discuss.pytorch.org/t/help-regarding-slerp-function-for-generative-model-sampling/32475/3 - low_norm = low/torch.norm(low, dim=1, keepdim=True) - high_norm = high/torch.norm(high, dim=1, keepdim=True) - dot = (low_norm*high_norm).sum(1) - - if dot.mean() > 0.9995: - return low * val + high * (1 - val) - - omega = torch.acos(dot) - so = torch.sin(omega) - res = (torch.sin((1.0-val)*omega)/so).unsqueeze(1)*low + (torch.sin(val*omega)/so).unsqueeze(1) * high - return res - - -def create_random_tensors(shape, seeds, subseeds=None, subseed_strength=0.0, seed_resize_from_h=0, seed_resize_from_w=0, p=None): - eta_noise_seed_delta = shared.opts.eta_noise_seed_delta or 0 - xs = [] - # if we have multiple seeds, this means we are working with batch size>1; this then - # enables the generation of additional tensors with noise that the sampler will use during its processing. - # Using those pre-generated tensors instead of simple torch.randn allows a batch with seeds [100, 101] to - # produce the same images as with two batches [100], [101]. - if p is not None and p.sampler is not None and (len(seeds) > 1 and shared.opts.enable_batch_seeds or eta_noise_seed_delta > 0): - sampler_noises = [[] for _ in range(p.sampler.number_of_needed_noises(p))] - else: - sampler_noises = None - for i, seed in enumerate(seeds): - noise_shape = shape if seed_resize_from_h <= 0 or seed_resize_from_w <= 0 else (shape[0], seed_resize_from_h//8, seed_resize_from_w//8) - subnoise = None - if subseeds is not None: - subseed = 0 if i >= len(subseeds) else subseeds[i] - subnoise = devices.randn(subseed, noise_shape) - # randn results depend on device; gpu and cpu get different results for same seed; - # the way I see it, it's better to do this on CPU, so that everyone gets same result; - # but the original script had it like this, so I do not dare change it for now because - # it will break everyone's seeds. - noise = devices.randn(seed, noise_shape) - if subnoise is not None: - noise = slerp(subseed_strength, noise, subnoise) - if noise_shape != shape: - x = devices.randn(seed, shape) - dx = (shape[2] - noise_shape[2]) // 2 - dy = (shape[1] - noise_shape[1]) // 2 - w = noise_shape[2] if dx >= 0 else noise_shape[2] + 2 * dx - h = noise_shape[1] if dy >= 0 else noise_shape[1] + 2 * dy - tx = 0 if dx < 0 else dx - ty = 0 if dy < 0 else dy - dx = max(-dx, 0) - dy = max(-dy, 0) - x[:, ty:ty+h, tx:tx+w] = noise[:, dy:dy+h, dx:dx+w] - noise = x - if sampler_noises is not None: - cnt = p.sampler.number_of_needed_noises(p) - if eta_noise_seed_delta > 0: - torch.manual_seed(seed + eta_noise_seed_delta) - for j in range(cnt): - sampler_noises[j].append(devices.randn_without_seed(tuple(noise_shape))) - xs.append(noise) - if sampler_noises is not None: - p.sampler.sampler_noises = [torch.stack(n).to(shared.device) for n in sampler_noises] - x = torch.stack(xs).to(shared.device) - return x - - -def decode_first_stage(model, x, full_quality=True): - if not shared.opts.keep_incomplete and (shared.state.skipped or shared.state.interrupted): - shared.log.debug(f'Decode VAE: skipped={shared.state.skipped} interrupted={shared.state.interrupted}') - x_sample = torch.zeros((len(x), 3, x.shape[2] * 8, x.shape[3] * 8), dtype=devices.dtype_vae, device=devices.device) - return x_sample - prev_job = shared.state.job - shared.state.job = 'vae' - with devices.autocast(disable = x.dtype==devices.dtype_vae): - try: - if full_quality: - if hasattr(model, 'decode_first_stage'): - x_sample = model.decode_first_stage(x) - elif hasattr(model, 'vae'): - x_sample = model.vae(x) - else: - x_sample = x - shared.log.error('Decode VAE unknown model') - else: - x_sample = torch.zeros((len(x), 3, x.shape[2] * 8, x.shape[3] * 8), dtype=devices.dtype_vae, device=devices.device) - for i in range(len(x_sample)): - x_sample[i] = sd_vae_taesd.decode(x[i]) - except Exception as e: - x_sample = x - shared.log.error(f'Decode VAE: {e}') - shared.state.job = prev_job - return x_sample - - -def get_fixed_seed(seed): - if seed is None or seed == '' or seed == -1: - return int(random.randrange(4294967294)) - return seed - - -def fix_seed(p): - p.seed = get_fixed_seed(p.seed) - p.subseed = get_fixed_seed(p.subseed) - - -def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=None, all_subseeds=None, comments=None, iteration=0, position_in_batch=0, index=None, all_negative_prompts=None): - if p is None: - shared.log.warning('Processing info: no data') - return '' - if not hasattr(shared.sd_model, 'sd_checkpoint_info'): - return '' - if index is None: - index = position_in_batch + iteration * p.batch_size - if all_prompts is None: - all_prompts = p.all_prompts or [p.prompt] - if all_negative_prompts is None: - all_negative_prompts = p.all_negative_prompts or [p.negative_prompt] - if all_seeds is None: - all_seeds = p.all_seeds or [p.seed] - if all_subseeds is None: - all_subseeds = p.all_subseeds or [p.subseed] - while len(all_prompts) <= index: - all_prompts.append(all_prompts[-1]) - while len(all_seeds) <= index: - all_seeds.append(all_seeds[-1]) - while len(all_subseeds) <= index: - all_subseeds.append(all_subseeds[-1]) - while len(all_negative_prompts) <= index: - all_negative_prompts.append(all_negative_prompts[-1]) - comment = ', '.join(comments) if comments is not None and type(comments) is list else None - ops = list(set(p.ops)) - ops.reverse() - args = { - # basic - "Steps": p.steps, - "Seed": all_seeds[index], - "Sampler": p.sampler_name, - "CFG scale": p.cfg_scale, - "Size": f"{p.width}x{p.height}" if hasattr(p, 'width') and hasattr(p, 'height') else None, - "Batch": f'{p.n_iter}x{p.batch_size}' if p.n_iter > 1 or p.batch_size > 1 else None, - "Index": f'{p.iteration + 1}x{index + 1}' if (p.n_iter > 1 or p.batch_size > 1) and index >= 0 else None, - "Parser": shared.opts.prompt_attention, - "Model": None if (not shared.opts.add_model_name_to_info) or (not shared.sd_model.sd_checkpoint_info.model_name) else shared.sd_model.sd_checkpoint_info.model_name.replace(',', '').replace(':', ''), - "Model hash": getattr(p, 'sd_model_hash', None if (not shared.opts.add_model_hash_to_info) or (not shared.sd_model.sd_model_hash) else shared.sd_model.sd_model_hash), - "VAE": (None if not shared.opts.add_model_name_to_info or sd_vae.loaded_vae_file is None else os.path.splitext(os.path.basename(sd_vae.loaded_vae_file))[0]) if p.full_quality else 'TAESD', - "Seed resize from": None if p.seed_resize_from_w == 0 or p.seed_resize_from_h == 0 else f"{p.seed_resize_from_w}x{p.seed_resize_from_h}", - "Clip skip": p.clip_skip if p.clip_skip > 1 else None, - "Prompt2": p.refiner_prompt if len(p.refiner_prompt) > 0 else None, - "Negative2": p.refiner_negative if len(p.refiner_negative) > 0 else None, - "Styles": "; ".join(p.styles) if p.styles is not None and len(p.styles) > 0 else None, - "Tiling": p.tiling if p.tiling else None, - # sdnext - "Backend": 'Diffusers' if shared.backend == shared.Backend.DIFFUSERS else 'Original', - "App": 'SD.Next', - "Version": git_commit, - "Comment": comment, - "Operations": '; '.join(ops).replace('"', '') if len(p.ops) > 0 else 'none', - } - if 'txt2img' in p.ops: - pass - if shared.backend == shared.Backend.ORIGINAL: - args["Variation seed"] = all_subseeds[index] if p.subseed_strength > 0 else None - args["Variation strength"] = p.subseed_strength if p.subseed_strength > 0 else None - if 'hires' in p.ops or 'upscale' in p.ops: - args["Second pass"] = p.enable_hr - args["Hires force"] = p.hr_force - args["Hires steps"] = p.hr_second_pass_steps - args["Hires upscaler"] = p.hr_upscaler - args["Hires upscale"] = p.hr_scale - args["Hires resize"] = f"{p.hr_resize_x}x{p.hr_resize_y}" - args["Hires size"] = f"{p.hr_upscale_to_x}x{p.hr_upscale_to_y}" - args["Denoising strength"] = p.denoising_strength - args["Hires sampler"] = p.hr_sampler_name - args["Image CFG scale"] = p.image_cfg_scale - 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(':', '') - args['Image CFG scale'] = p.image_cfg_scale - args['Refiner steps'] = p.refiner_steps - args['Refiner start'] = p.refiner_start - args["Hires steps"] = p.hr_second_pass_steps - args["Hires sampler"] = p.hr_sampler_name - 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) - args['Resize scale'] = getattr(p, 'scale_by', None) - args["Mask weight"] = getattr(p, "inpainting_mask_weight", shared.opts.inpainting_mask_weight) if p.is_using_inpainting_conditioning else None - args["Denoising strength"] = getattr(p, 'denoising_strength', None) - if args["Size"] is None: - args["Size"] = args["Init image size"] - # lookup by index - if getattr(p, 'resize_mode', None) is not None: - args['Resize mode'] = shared.resize_modes[p.resize_mode] if shared.resize_modes[p.resize_mode] != 'None' else None - if 'face' in p.ops: - args["Face restoration"] = shared.opts.face_restoration_model - if 'color' in p.ops: - args["Color correction"] = True - # embeddings - if sd_hijack is not None and hasattr(sd_hijack.model_hijack, 'embedding_db') and len(sd_hijack.model_hijack.embedding_db.embeddings_used) > 0: # this is for original hijaacked models only, diffusers are handled separately - args["Embeddings"] = ', '.join(sd_hijack.model_hijack.embedding_db.embeddings_used) - # samplers - args["Sampler ENSD"] = shared.opts.eta_noise_seed_delta if shared.opts.eta_noise_seed_delta != 0 and sd_samplers_common.is_sampler_using_eta_noise_seed_delta(p) else None - args["Sampler ENSM"] = p.initial_noise_multiplier if getattr(p, 'initial_noise_multiplier', 1.0) != 1.0 else None - args['Sampler order'] = shared.opts.schedulers_solver_order if shared.opts.schedulers_solver_order != shared.opts.data_labels.get('schedulers_solver_order').default else None - if shared.backend == shared.Backend.DIFFUSERS: - args['Sampler beta schedule'] = shared.opts.schedulers_beta_schedule if shared.opts.schedulers_beta_schedule != shared.opts.data_labels.get('schedulers_beta_schedule').default else None - args['Sampler beta start'] = shared.opts.schedulers_beta_start if shared.opts.schedulers_beta_start != shared.opts.data_labels.get('schedulers_beta_start').default else None - args['Sampler beta end'] = shared.opts.schedulers_beta_end if shared.opts.schedulers_beta_end != shared.opts.data_labels.get('schedulers_beta_end').default else None - args['Sampler DPM solver'] = shared.opts.schedulers_dpm_solver if shared.opts.schedulers_dpm_solver != shared.opts.data_labels.get('schedulers_dpm_solver').default else None - if shared.backend == shared.Backend.ORIGINAL: - args['Sampler brownian'] = shared.opts.schedulers_brownian_noise if shared.opts.schedulers_brownian_noise != shared.opts.data_labels.get('schedulers_brownian_noise').default else None - args['Sampler discard'] = shared.opts.schedulers_discard_penultimate if shared.opts.schedulers_discard_penultimate != shared.opts.data_labels.get('schedulers_discard_penultimate').default else None - args['Sampler dyn threshold'] = shared.opts.schedulers_use_thresholding if shared.opts.schedulers_use_thresholding != shared.opts.data_labels.get('schedulers_use_thresholding').default else None - args['Sampler karras'] = shared.opts.schedulers_use_karras if shared.opts.schedulers_use_karras != shared.opts.data_labels.get('schedulers_use_karras').default else None - args['Sampler low order'] = shared.opts.schedulers_use_loworder if shared.opts.schedulers_use_loworder != shared.opts.data_labels.get('schedulers_use_loworder').default else None - args['Sampler quantization'] = shared.opts.enable_quantization if shared.opts.enable_quantization != shared.opts.data_labels.get('enable_quantization').default else None - args['Sampler sigma'] = shared.opts.schedulers_sigma if shared.opts.schedulers_sigma != shared.opts.data_labels.get('schedulers_sigma').default else None - args['Sampler sigma min'] = shared.opts.s_min if shared.opts.s_min != shared.opts.data_labels.get('s_min').default else None - args['Sampler sigma max'] = shared.opts.s_max if shared.opts.s_max != shared.opts.data_labels.get('s_max').default else None - args['Sampler sigma churn'] = shared.opts.s_churn if shared.opts.s_churn != shared.opts.data_labels.get('s_churn').default else None - args['Sampler sigma uncond'] = shared.opts.s_churn if shared.opts.s_churn != shared.opts.data_labels.get('s_churn').default else None - args['Sampler sigma noise'] = shared.opts.s_noise if shared.opts.s_noise != shared.opts.data_labels.get('s_noise').default else None - args['Sampler sigma tmin'] = shared.opts.s_tmin if shared.opts.s_tmin != shared.opts.data_labels.get('s_tmin').default else None - # tome - token_merging_ratio = p.get_token_merging_ratio() - token_merging_ratio_hr = p.get_token_merging_ratio(for_hr=True) if p.enable_hr else None - args['ToMe'] = token_merging_ratio if token_merging_ratio != 0 else None - args['ToMe hires'] = token_merging_ratio_hr if token_merging_ratio_hr != 0 else None - - args.update(p.extra_generation_params) - params_text = ", ".join([k if k == v else f'{k}: {generation_parameters_copypaste.quote(v)}' for k, v in args.items() if v is not None]) - negative_prompt_text = f"\nNegative prompt: {all_negative_prompts[index]}" if all_negative_prompts[index] else "" - infotext = f"{all_prompts[index]}{negative_prompt_text}\n{params_text}".strip() - return infotext - - def process_images(p: StableDiffusionProcessing) -> Processed: debug(f'Process images: {vars(p)}') if not hasattr(p.sd_model, 'sd_checkpoint_info'): @@ -712,7 +131,7 @@ def process_images(p: StableDiffusionProcessing) -> Processed: p.override_settings.pop(k, None) for k in p.override_settings.keys(): stored_opts[k] = shared.opts.data.get(k, None) or shared.opts.data_labels[k].default - res = None + processed = None try: # if no checkpoint override or the override checkpoint can't be found, remove override entry and load opts checkpoint if p.override_settings.get('sd_model_checkpoint', None) is not None and sd_models.checkpoint_aliases.get(p.override_settings.get('sd_model_checkpoint')) is None: @@ -759,12 +178,12 @@ def process_images(p: StableDiffusionProcessing) -> Processed: shared.profiler = torch.profiler.profile(activities=activities, profile_memory=True, with_modules=True) shared.profiler.start() shared.profiler.step() - res = process_images_inner(p) + processed = process_images_inner(p) errors.profile_torch(shared.profiler, 'Process') errors.profile(profile_python, 'Process') else: with context_hypertile_vae(p), context_hypertile_unet(p): - res = process_images_inner(p) + processed = process_images_inner(p) finally: if not shared.opts.cuda_compile: @@ -781,31 +200,7 @@ def process_images(p: StableDiffusionProcessing) -> Processed: sd_models.reload_model_weights() if k == 'sd_vae': sd_vae.reload_vae_weights() - return res - - -def validate_sample(tensor): - if not isinstance(tensor, np.ndarray) and not isinstance(tensor, torch.Tensor): - return tensor - if tensor.dtype == torch.bfloat16: # numpy does not support bf16 - tensor = tensor.to(torch.float16) - if isinstance(tensor, torch.Tensor) and hasattr(tensor, 'detach'): - sample = tensor.detach().cpu().numpy() - elif isinstance(tensor, np.ndarray): - sample = tensor - else: - shared.log.warning(f'Unknown sample type: {type(tensor)}') - sample = 255.0 * np.moveaxis(sample, 0, 2) if shared.backend == shared.Backend.ORIGINAL else 255.0 * sample - with warnings.catch_warnings(record=True) as w: - cast = sample.astype(np.uint8) - if len(w) > 0: - nans = np.isnan(sample).sum() - shared.log.error(f'Failed to validate samples: sample={sample.shape} invalid={nans}') - cast = np.nan_to_num(sample) - minimum, maximum, mean = np.min(cast), np.max(cast), np.mean(cast) - cast = cast.astype(np.uint8) - shared.log.warning(f'Attempted to correct samples: min={minimum:.2f} max={maximum:.2f} mean={mean:.2f}') - return cast + return processed def process_init(p: StableDiffusionProcessing): @@ -844,8 +239,6 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed: comments = {} infotexts = [] output_images = [] - cached_uc = [None, None] - cached_c = [None, None] process_init(p) if os.path.exists(shared.opts.embeddings_dir) and not p.do_not_reload_embeddings and shared.backend == shared.Backend.ORIGINAL: @@ -857,14 +250,6 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed: from modules import ipadapter ipadapter.apply(shared.sd_model, p) - def get_conds_with_caching(function, required_prompts, steps, cache): - if cache[0] is not None and (required_prompts, steps) == cache[0]: - return cache[1] - with devices.autocast(): - cache[1] = function(shared.sd_model, required_prompts, steps) - cache[0] = (required_prompts, steps) - return cache[1] - def infotext(_inxex=0): # dummy function overriden if there are iterations return '' @@ -899,42 +284,19 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed: extra_networks.activate(p, extra_network_data) if p.scripts is not None and isinstance(p.scripts, scripts.ScriptRunner): p.scripts.process_batch(p, batch_number=n, prompts=p.prompts, seeds=p.seeds, subseeds=p.subseeds) - step_multiplier = 1 - sampler_config = sd_samplers.find_sampler_config(p.sampler_name) - step_multiplier = 2 if sampler_config and sampler_config.options.get("second_order", False) else 1 - if shared.backend == shared.Backend.ORIGINAL: - uc = get_conds_with_caching(prompt_parser.get_learned_conditioning, p.negative_prompts, p.steps * step_multiplier, cached_uc) - c = get_conds_with_caching(prompt_parser.get_multicond_learned_conditioning, p.prompts, p.steps * step_multiplier, cached_c) - if len(modules.sd_hijack.model_hijack.comments) > 0: - for comment in modules.sd_hijack.model_hijack.comments: - comments[comment] = 1 - with devices.without_autocast() if devices.unet_needs_upcast else devices.autocast(): - samples_ddim = p.sample(conditioning=c, unconditional_conditioning=uc, seeds=p.seeds, subseeds=p.subseeds, subseed_strength=p.subseed_strength, prompts=p.prompts) - x_samples_ddim = [decode_first_stage(p.sd_model, samples_ddim[i:i+1].to(dtype=devices.dtype_vae), p.full_quality)[0].cpu() for i in range(samples_ddim.size(0))] - try: - for x in x_samples_ddim: - devices.test_for_nans(x, "vae") - except devices.NansException as e: - if not shared.opts.no_half and not shared.opts.no_half_vae and shared.cmd_opts.rollback_vae: - shared.log.warning('Tensor with all NaNs was produced in VAE') - devices.dtype_vae = torch.bfloat16 - vae_file, vae_source = modules.sd_vae.resolve_vae(p.sd_model.sd_model_checkpoint) - sd_vae.load_vae(p.sd_model, vae_file, vae_source) - x_samples_ddim = [decode_first_stage(p.sd_model, samples_ddim[i:i+1].to(dtype=devices.dtype_vae), p.full_quality)[0].cpu() for i in range(samples_ddim.size(0))] - for x in x_samples_ddim: - devices.test_for_nans(x, "vae") - else: - raise e - x_samples_ddim = torch.stack(x_samples_ddim).float() - x_samples_ddim = torch.clamp((x_samples_ddim + 1.0) / 2.0, min=0.0, max=1.0) - del samples_ddim - - elif shared.backend == shared.Backend.DIFFUSERS: - from modules.processing_diffusers import process_diffusers - x_samples_ddim = process_diffusers(p) - else: - raise ValueError(f"Unknown backend {shared.backend}") + x_samples_ddim = None + if p.scripts is not None and isinstance(p.scripts, scripts.ScriptRunner): + x_samples_ddim = p.scripts.process_images(p) + if x_samples_ddim is None: + if shared.backend == shared.Backend.ORIGINAL: + from modules.processing_original import process_original + x_samples_ddim = process_original(p) + elif shared.backend == shared.Backend.DIFFUSERS: + from modules.processing_diffusers import process_diffusers + x_samples_ddim = process_diffusers(p) + else: + raise ValueError(f"Unknown backend {shared.backend}") if not shared.opts.keep_incomplete and shared.state.interrupted: x_samples_ddim = [] @@ -1030,7 +392,7 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed: if not p.disable_extra_networks: extra_networks.deactivate(p, extra_network_data) - res = Processed( + processed = Processed( p, images_list=output_images, seed=p.all_seeds[0], @@ -1041,379 +403,5 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed: infotexts=infotexts, ) if p.scripts is not None and isinstance(p.scripts, scripts.ScriptRunner) and not (shared.state.interrupted or shared.state.skipped): - p.scripts.postprocess(p, res) - return res - - -def old_hires_fix_first_pass_dimensions(width, height): - """old algorithm for auto-calculating first pass size""" - desired_pixel_count = 512 * 512 - actual_pixel_count = width * height - scale = math.sqrt(desired_pixel_count / actual_pixel_count) - width = math.ceil(scale * width / 64) * 64 - height = math.ceil(scale * height / 64) * 64 - return width, height - - -class StableDiffusionProcessingTxt2Img(StableDiffusionProcessing): - - def __init__(self, enable_hr: bool = False, denoising_strength: float = 0.75, firstphase_width: int = 0, firstphase_height: int = 0, hr_scale: float = 2.0, hr_force: bool = False, hr_upscaler: str = None, hr_second_pass_steps: int = 0, hr_resize_x: int = 0, hr_resize_y: int = 0, refiner_steps: int = 5, refiner_start: float = 0, refiner_prompt: str = '', refiner_negative: str = '', **kwargs): - - super().__init__(**kwargs) - if devices.backend == "ipex" and not torch.xpu.has_fp64_dtype() and os.environ.get('DISABLE_IPEX_1024_WA', None) is None: - width_curse = bool(hr_resize_x == 1024 and self.height * (hr_resize_x / self.width) == 1024) - height_curse = bool(hr_resize_y == 1024 and self.width * (hr_resize_y / self.height) == 1024) - if (width_curse != height_curse) or (height_curse and width_curse): - if width_curse: - hr_resize_x = 1080 - if height_curse: - hr_resize_y = 1080 - if self.width * hr_scale == 1024 and self.height * hr_scale == 1024: - hr_scale = 1080 / self.width - if firstphase_width * hr_scale == 1024 and firstphase_height * hr_scale == 1024: - hr_scale = 1080 / firstphase_width - self.enable_hr = enable_hr - self.denoising_strength = denoising_strength - self.hr_scale = hr_scale - self.hr_upscaler = hr_upscaler - self.hr_force = hr_force - self.hr_second_pass_steps = hr_second_pass_steps - self.hr_resize_x = hr_resize_x - self.hr_resize_y = hr_resize_y - self.hr_upscale_to_x = hr_resize_x - self.hr_upscale_to_y = hr_resize_y - if firstphase_width != 0 or firstphase_height != 0: - self.hr_upscale_to_x = self.width - self.hr_upscale_to_y = self.height - self.width = firstphase_width - self.height = firstphase_height - self.truncate_x = 0 - self.truncate_y = 0 - self.applied_old_hires_behavior_to = None - self.refiner_steps = refiner_steps - self.refiner_start = refiner_start - self.refiner_prompt = refiner_prompt - self.refiner_negative = refiner_negative - self.sampler = None - self.scripts = None - self.script_args = [] - - def init(self, all_prompts, all_seeds, all_subseeds): - if shared.backend == shared.Backend.DIFFUSERS: - shared.sd_model = sd_models.set_diffuser_pipe(self.sd_model, sd_models.DiffusersTaskType.TEXT_2_IMAGE) - self.width = self.width or 512 - self.height = self.height or 512 - - def init_hr(self): - if self.hr_resize_x == 0 and self.hr_resize_y == 0: - self.hr_upscale_to_x = int(self.width * self.hr_scale) - self.hr_upscale_to_y = int(self.height * self.hr_scale) - else: - if self.hr_resize_y == 0: - self.hr_upscale_to_x = self.hr_resize_x - self.hr_upscale_to_y = self.hr_resize_x * self.height // self.width - elif self.hr_resize_x == 0: - self.hr_upscale_to_x = self.hr_resize_y * self.width // self.height - self.hr_upscale_to_y = self.hr_resize_y - else: - target_w = self.hr_resize_x - target_h = self.hr_resize_y - src_ratio = self.width / self.height - dst_ratio = self.hr_resize_x / self.hr_resize_y - if src_ratio < dst_ratio: - self.hr_upscale_to_x = self.hr_resize_x - self.hr_upscale_to_y = self.hr_resize_x * self.height // self.width - else: - self.hr_upscale_to_x = self.hr_resize_y * self.width // self.height - self.hr_upscale_to_y = self.hr_resize_y - self.truncate_x = (self.hr_upscale_to_x - target_w) // 8 - self.truncate_y = (self.hr_upscale_to_y - target_h) // 8 - # special case: the user has chosen to do nothing - if (self.hr_upscale_to_x == self.width and self.hr_upscale_to_y == self.height) or self.hr_upscaler is None or self.hr_upscaler == 'None': - self.is_hr_pass = False - return - self.is_hr_pass = True - hypertile_set(self, hr=True) - shared.state.job_count = 2 * self.n_iter - shared.log.debug(f'Init hires: upscaler="{self.hr_upscaler}" sampler="{self.hr_sampler_name}" resize={self.hr_resize_x}x{self.hr_resize_y} upscale={self.hr_upscale_to_x}x{self.hr_upscale_to_y}') - - def sample(self, conditioning, unconditional_conditioning, seeds, subseeds, subseed_strength, prompts): - - latent_scale_mode = shared.latent_upscale_modes.get(self.hr_upscaler, None) if self.hr_upscaler is not None else shared.latent_upscale_modes.get(shared.latent_upscale_default_mode, "None") - if latent_scale_mode is not None: - self.hr_force = False # no need to force anything - if self.enable_hr and (latent_scale_mode is None or self.hr_force): - if len([x for x in shared.sd_upscalers if x.name == self.hr_upscaler]) == 0: - shared.log.warning(f"Cannot find upscaler for hires: {self.hr_upscaler}") - self.enable_hr = False - - self.ops.append('txt2img') - hypertile_set(self) - self.sampler = sd_samplers.create_sampler(self.sampler_name, self.sd_model) - if hasattr(self.sampler, "initialize"): - self.sampler.initialize(self) - x = create_random_tensors([4, self.height // 8, self.width // 8], seeds=seeds, subseeds=subseeds, subseed_strength=self.subseed_strength, seed_resize_from_h=self.seed_resize_from_h, seed_resize_from_w=self.seed_resize_from_w, p=self) - samples = self.sampler.sample(self, x, conditioning, unconditional_conditioning, image_conditioning=self.txt2img_image_conditioning(x)) - shared.state.nextjob() - if not self.enable_hr or shared.state.interrupted or shared.state.skipped: - return samples - - self.init_hr() - if self.is_hr_pass: - prev_job = shared.state.job - target_width = self.hr_upscale_to_x - target_height = self.hr_upscale_to_y - decoded_samples = None - if shared.opts.save and shared.opts.save_images_before_highres_fix and not self.do_not_save_samples: - 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 = validate_sample(x_sample) - image = Image.fromarray(x_sample) - bak_extra_generation_params, bak_restore_faces = self.extra_generation_params, self.restore_faces - self.extra_generation_params = {} - self.restore_faces = False - info = create_infotext(self, self.all_prompts, self.all_seeds, self.all_subseeds, [], iteration=self.iteration, position_in_batch=i) - self.extra_generation_params, self.restore_faces = bak_extra_generation_params, bak_restore_faces - images.save_image(image, self.outpath_samples, "", seeds[i], prompts[i], shared.opts.samples_format, info=info, suffix="-before-hires") - if latent_scale_mode is None or self.hr_force: # non-latent upscaling - shared.state.job = 'upscale' - if decoded_samples is None: - 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) - batch_images = [] - for _i, x_sample in enumerate(decoded_samples): - 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) - image = np.array(image).astype(np.float32) / 255.0 - image = np.moveaxis(image, 2, 0) - batch_images.append(image) - resized_samples = torch.from_numpy(np.array(batch_images)) - resized_samples = resized_samples.to(device=shared.device, dtype=devices.dtype_vae) - resized_samples = 2.0 * resized_samples - 1.0 - if shared.opts.sd_vae_sliced_encode and len(decoded_samples) > 1: - samples = torch.stack([self.sd_model.get_first_stage_encoding(self.sd_model.encode_first_stage(torch.unsqueeze(resized_sample, 0)))[0] for resized_sample in resized_samples]) - else: - samples = self.sd_model.get_first_stage_encoding(self.sd_model.encode_first_stage(resized_samples)) - image_conditioning = self.img2img_image_conditioning(resized_samples, samples) - else: - samples = torch.nn.functional.interpolate(samples, size=(target_height // 8, target_width // 8), mode=latent_scale_mode["mode"], antialias=latent_scale_mode["antialias"]) - if getattr(self, "inpainting_mask_weight", shared.opts.inpainting_mask_weight) < 1.0: - image_conditioning = self.img2img_image_conditioning(decode_first_stage(self.sd_model, samples.to(dtype=devices.dtype_vae), self.full_quality), samples) - else: - image_conditioning = self.txt2img_image_conditioning(samples.to(dtype=devices.dtype_vae)) - if self.hr_sampler_name == "PLMS": - self.hr_sampler_name = 'UniPC' - if self.hr_force or latent_scale_mode is not None: - shared.state.job = 'hires' - if self.denoising_strength > 0: - self.ops.append('hires') - devices.torch_gc() # GC now before running the next img2img to prevent running out of memory - self.sampler = sd_samplers.create_sampler(self.hr_sampler_name or self.sampler_name, self.sd_model) - if hasattr(self.sampler, "initialize"): - self.sampler.initialize(self) - samples = samples[:, :, self.truncate_y//2:samples.shape[2]-(self.truncate_y+1)//2, self.truncate_x//2:samples.shape[3]-(self.truncate_x+1)//2] - noise = create_random_tensors(samples.shape[1:], seeds=seeds, subseeds=subseeds, subseed_strength=subseed_strength, p=self) - sd_models.apply_token_merging(self.sd_model, self.get_token_merging_ratio(for_hr=True)) - hypertile_set(self, hr=True) - samples = self.sampler.sample_img2img(self, samples, noise, conditioning, unconditional_conditioning, steps=self.hr_second_pass_steps or self.steps, image_conditioning=image_conditioning) - sd_models.apply_token_merging(self.sd_model, self.get_token_merging_ratio()) - else: - self.ops.append('upscale') - x = None - self.is_hr_pass = False - shared.state.job = prev_job - shared.state.nextjob() - - return samples - - -class StableDiffusionProcessingImg2Img(StableDiffusionProcessing): - - def __init__(self, init_images: list = None, resize_mode: int = 0, resize_name: str = 'None', denoising_strength: float = 0.3, image_cfg_scale: float = None, mask: Any = None, mask_blur: int = 4, inpainting_fill: int = 0, inpaint_full_res: bool = False, inpaint_full_res_padding: int = 0, inpainting_mask_invert: int = 0, initial_noise_multiplier: float = None, scale_by: float = 1, refiner_steps: int = 5, refiner_start: float = 0, refiner_prompt: str = '', refiner_negative: str = '', **kwargs): - super().__init__(**kwargs) - self.init_images = init_images - self.resize_mode: int = resize_mode - self.resize_name: str = resize_name - self.denoising_strength: float = denoising_strength - self.image_cfg_scale: float = image_cfg_scale - self.init_latent = None - self.image_mask = mask - self.latent_mask = None - self.mask_for_overlay = None - self.mask_blur_x = mask_blur # a1111 compatibility item - self.mask_blur_y = mask_blur # a1111 compatibility item - self.mask_blur = mask_blur - self.inpainting_fill = inpainting_fill - self.inpaint_full_res = inpaint_full_res - self.inpaint_full_res_padding = inpaint_full_res_padding - self.inpainting_mask_invert = inpainting_mask_invert - self.initial_noise_multiplier = shared.opts.initial_noise_multiplier if initial_noise_multiplier is None else initial_noise_multiplier - self.mask = None - self.nmask = None - self.image_conditioning = None - self.refiner_steps = refiner_steps - self.refiner_start = refiner_start - self.refiner_prompt = refiner_prompt - self.refiner_negative = refiner_negative - self.enable_hr = None - self.is_batch = False - self.scale_by = scale_by - self.sampler = None - self.scripts = None - self.script_args = [] - - def init(self, all_prompts, all_seeds, all_subseeds): - if shared.backend == shared.Backend.DIFFUSERS and self.image_mask is not None and not self.is_control: - shared.sd_model = sd_models.set_diffuser_pipe(self.sd_model, sd_models.DiffusersTaskType.INPAINTING) - elif shared.backend == shared.Backend.DIFFUSERS and self.image_mask is None and not self.is_control: - shared.sd_model = sd_models.set_diffuser_pipe(self.sd_model, sd_models.DiffusersTaskType.IMAGE_2_IMAGE) - - if self.sampler_name == "PLMS": - self.sampler_name = 'UniPC' - self.sampler = sd_samplers.create_sampler(self.sampler_name, self.sd_model) - if hasattr(self.sampler, "initialize"): - self.sampler.initialize(self) - - if self.image_mask is not None: - self.ops.append('inpaint') - else: - self.ops.append('img2img') - crop_region = None - - if self.image_mask is not None: - if type(self.image_mask) == list: - self.image_mask = self.image_mask[0] - if shared.backend == shared.Backend.ORIGINAL: # original way of processing mask - self.image_mask = create_binary_mask(self.image_mask) - if self.inpainting_mask_invert: - self.image_mask = ImageOps.invert(self.image_mask) - if self.mask_blur > 0: - np_mask = np.array(self.image_mask) - kernel_size = 2 * int(2.5 * self.mask_blur + 0.5) + 1 - np_mask = cv2.GaussianBlur(np_mask, (kernel_size, 1), self.mask_blur) - np_mask = cv2.GaussianBlur(np_mask, (1, kernel_size), self.mask_blur) - self.image_mask = Image.fromarray(np_mask) - else: - if hasattr(self, 'init_images'): - self.image_mask = masking.run_mask( - input_image=self.init_images, - input_mask=self.image_mask, - return_type='Grayscale', - mask_blur=self.mask_blur, - mask_padding=self.inpaint_full_res_padding, - segment_enable=False, - invert=self.inpainting_mask_invert, - ) - if self.inpaint_full_res: # mask only inpaint - self.mask_for_overlay = self.image_mask - mask = self.image_mask.convert('L') - crop_region = masking.get_crop_region(np.array(mask), self.inpaint_full_res_padding) - crop_region = masking.expand_crop_region(crop_region, self.width, self.height, mask.width, mask.height) - x1, y1, x2, y2 = crop_region - crop_mask = mask.crop(crop_region) - self.image_mask = images.resize_image(2, crop_mask, self.width, self.height) - self.paste_to = (x1, y1, x2-x1, y2-y1) - else: # full image inpaint - self.image_mask = images.resize_image(self.resize_mode, self.image_mask, self.width, self.height) - np_mask = np.array(self.image_mask) - np_mask = np.clip((np_mask.astype(np.float32)) * 2, 0, 255).astype(np.uint8) - self.mask_for_overlay = Image.fromarray(np_mask) - self.overlay_images = [] - - latent_mask = self.latent_mask if self.latent_mask is not None else self.image_mask - - add_color_corrections = shared.opts.img2img_color_correction and self.color_corrections is None - if add_color_corrections: - self.color_corrections = [] - processed = [] - if getattr(self, 'init_images', None) is None: - return - if not isinstance(self.init_images, list): - self.init_images = [self.init_images] - for img in self.init_images: - if img is None: - shared.log.warning(f"Skipping empty image: images={self.init_images}") - continue - self.init_img_hash = hashlib.sha256(img.tobytes()).hexdigest()[0:8] # pylint: disable=attribute-defined-outside-init - self.init_img_width = img.width # pylint: disable=attribute-defined-outside-init - self.init_img_height = img.height # pylint: disable=attribute-defined-outside-init - if shared.opts.save_init_img: - images.save_image(img, path=shared.opts.outdir_init_images, basename=None, forced_filename=self.init_img_hash, suffix="-init-image") - image = images.flatten(img, shared.opts.img2img_background_color) - if self.width is None or self.height is None or self.resize_mode == 0: - self.width, self.height = image.width, image.height - if crop_region is None and self.resize_mode != 4 and self.resize_mode > 0: - if image.width != self.width or image.height != self.height: - image = images.resize_image(self.resize_mode, image, self.width, self.height, self.resize_name) - self.width = image.width - self.height = image.height - if self.image_mask is not None: - try: - image_masked = Image.new('RGBa', (image.width, image.height)) - image_to_paste = image.convert("RGBA").convert("RGBa") - image_to_mask = ImageOps.invert(self.mask_for_overlay.convert('L')) if self.mask_for_overlay is not None else None - image_to_mask = image_to_mask.resize((image.width, image.height), Image.Resampling.BILINEAR) if image_to_mask is not None else None - image_masked.paste(image_to_paste, mask=image_to_mask) - self.overlay_images.append(image_masked.convert('RGBA')) - except Exception as e: - shared.log.error(f"Failed to apply mask to image: {e}") - if crop_region is not None: # crop_region is not None if we are doing inpaint full res - image = image.crop(crop_region) - if image.width != self.width or image.height != self.height: - image = images.resize_image(3, image, self.width, self.height, self.resize_name) - if self.image_mask is not None and self.inpainting_fill != 1: - image = masking.fill(image, latent_mask) - if add_color_corrections: - self.color_corrections.append(setup_color_correction(image)) - processed.append(image) - self.init_images = processed - self.batch_size = len(self.init_images) - if self.overlay_images is not None: - self.overlay_images = self.overlay_images * self.batch_size - if self.color_corrections is not None and len(self.color_corrections) == 1: - self.color_corrections = self.color_corrections * self.batch_size - if shared.backend == shared.Backend.DIFFUSERS: - return # we've already set self.init_images and self.mask and we dont need any more processing - - self.init_images = [np.moveaxis((np.array(image).astype(np.float32) / 255.0), 2, 0) for image in self.init_images] - if len(self.init_images) == 1: - batch_images = np.expand_dims(self.init_images[0], axis=0).repeat(self.batch_size, axis=0) - elif len(self.init_images) <= self.batch_size: - batch_images = np.array(self.init_images) - image = torch.from_numpy(batch_images) - image = 2. * image - 1. - image = image.to(device=shared.device, dtype=devices.dtype_vae) - self.init_latent = self.sd_model.get_first_stage_encoding(self.sd_model.encode_first_stage(image)) - if self.resize_mode == 4: - self.init_latent = torch.nn.functional.interpolate(self.init_latent, size=(self.height // 8, self.width // 8), mode="bilinear") - if self.image_mask is not None: - init_mask = latent_mask - latmask = init_mask.convert('RGB').resize((self.init_latent.shape[3], self.init_latent.shape[2])) - latmask = np.moveaxis(np.array(latmask, dtype=np.float32), 2, 0) / 255 - latmask = latmask[0] - latmask = np.tile(latmask[None], (4, 1, 1)) - latmask = np.around(latmask) - self.mask = torch.asarray(1.0 - latmask).to(device=shared.device, dtype=self.sd_model.dtype) - self.nmask = torch.asarray(latmask).to(device=shared.device, dtype=self.sd_model.dtype) - if self.inpainting_fill == 2: - self.init_latent = self.init_latent * self.mask + create_random_tensors(self.init_latent.shape[1:], all_seeds[0:self.init_latent.shape[0]]) * self.nmask - elif self.inpainting_fill == 3: - self.init_latent = self.init_latent * self.mask - self.image_conditioning = self.img2img_image_conditioning(image, self.init_latent, self.image_mask) - - def sample(self, conditioning, unconditional_conditioning, seeds, subseeds, subseed_strength, prompts): - hypertile_set(self) - x = create_random_tensors([4, self.height // 8, self.width // 8], seeds=seeds, subseeds=subseeds, subseed_strength=self.subseed_strength, seed_resize_from_h=self.seed_resize_from_h, seed_resize_from_w=self.seed_resize_from_w, p=self) - x *= self.initial_noise_multiplier - samples = self.sampler.sample_img2img(self, self.init_latent, x, conditioning, unconditional_conditioning, image_conditioning=self.image_conditioning) - if self.mask is not None: - samples = samples * self.nmask + self.init_latent * self.mask - del x - devices.torch_gc() - shared.state.nextjob() - return samples - - def get_token_merging_ratio(self, for_hr=False): - return self.token_merging_ratio or ("token_merging_ratio" in self.override_settings and shared.opts.token_merging_ratio) or shared.opts.token_merging_ratio_img2img or shared.opts.token_merging_ratio + p.scripts.postprocess(p, processed) + return processed diff --git a/modules/processing_class.py b/modules/processing_class.py new file mode 100644 index 000000000..db712dcaa --- /dev/null +++ b/modules/processing_class.py @@ -0,0 +1,487 @@ +import os +import math +import hashlib +from typing import Any, Dict, List +from dataclasses import dataclass, field +import torch +import numpy as np +import cv2 +from PIL import Image, ImageOps +from modules import shared, devices, images, scripts, masking, sd_samplers, sd_models, processing_helpers +from modules.sd_hijack_hypertile import hypertile_set + + +@dataclass(repr=False) +class StableDiffusionProcessing: + """ + The first set of paramaters: sd_models -> do_not_reload_embeddings represent the minimum required to create a StableDiffusionProcessing + """ + def __init__(self, sd_model=None, outpath_samples=None, outpath_grids=None, prompt: str = "", styles: List[str] = None, seed: int = -1, subseed: int = -1, subseed_strength: float = 0, seed_resize_from_h: int = -1, seed_resize_from_w: int = -1, seed_enable_extras: bool = True, sampler_name: str = None, hr_sampler_name: str = None, batch_size: int = 1, n_iter: int = 1, steps: int = 50, cfg_scale: float = 7.0, image_cfg_scale: float = None, clip_skip: int = 1, width: int = 512, height: int = 512, full_quality: bool = True, restore_faces: bool = False, tiling: bool = False, do_not_save_samples: bool = False, do_not_save_grid: bool = False, extra_generation_params: Dict[Any, Any] = None, overlay_images: Any = None, negative_prompt: str = None, eta: float = None, do_not_reload_embeddings: bool = False, denoising_strength: float = 0, diffusers_guidance_rescale: float = 0.7, sag_scale: float = 0.0, resize_mode: int = 0, resize_name: str = 'None', scale_by: float = 0, selected_scale_tab: int = 0, hdr_clamp: bool = False, hdr_boundary: float = 4.0, hdr_threshold: float = 3.5, hdr_center: bool = False, hdr_channel_shift: float = 0.8, hdr_full_shift: float = 0.8, hdr_maximize: bool = False, hdr_max_center: float = 0.6, hdr_max_boundry: float = 1.0, override_settings: Dict[str, Any] = None, override_settings_restore_afterwards: bool = True, sampler_index: int = None, script_args: list = None): # pylint: disable=unused-argument + self.outpath_samples: str = outpath_samples + self.outpath_grids: str = outpath_grids + self.prompt: str = prompt + self.prompt_for_display: str = None + self.negative_prompt: str = (negative_prompt or "") + self.styles: list = styles or [] + self.seed: int = seed + self.subseed: int = subseed + self.subseed_strength: float = subseed_strength + self.seed_resize_from_h: int = seed_resize_from_h + self.seed_resize_from_w: int = seed_resize_from_w + self.sampler_name: str = sampler_name + self.hr_sampler_name: str = hr_sampler_name + self.batch_size: int = batch_size + self.n_iter: int = n_iter + self.steps: int = steps + self.hr_second_pass_steps = 0 + self.cfg_scale: float = cfg_scale + self.scale_by: float = scale_by + self.image_cfg_scale = image_cfg_scale + self.diffusers_guidance_rescale = diffusers_guidance_rescale + self.sag_scale = sag_scale + if devices.backend == "ipex" and width == 1024 and height == 1024 and not torch.xpu.has_fp64_dtype() and os.environ.get('DISABLE_IPEX_1024_WA', None) is None: + width = 1080 + height = 1080 + self.width: int = width + self.height: int = height + self.full_quality: bool = full_quality + self.restore_faces: bool = restore_faces + self.tiling: bool = tiling + self.do_not_save_samples: bool = do_not_save_samples + self.do_not_save_grid: bool = do_not_save_grid + self.extra_generation_params: dict = extra_generation_params or {} + self.overlay_images = overlay_images + self.eta = eta + self.do_not_reload_embeddings = do_not_reload_embeddings + self.paste_to = None + self.color_corrections = None + self.denoising_strength: float = denoising_strength + self.override_settings = {k: v for k, v in (override_settings or {}).items() if k not in shared.restricted_opts} + self.override_settings_restore_afterwards = override_settings_restore_afterwards + self.is_using_inpainting_conditioning = False # a111 compatibility + self.disable_extra_networks = False + self.token_merging_ratio = 0 + self.token_merging_ratio_hr = 0 + # self.scripts = scripts.ScriptRunner() # set via property + # self.script_args = script_args or [] # set via property + self.per_script_args = {} + self.all_prompts = None + self.all_negative_prompts = None + self.all_seeds = None + self.all_subseeds = None + self.clip_skip = clip_skip + self.iteration = 0 + self.is_control = False + self.is_hr_pass = False + self.is_refiner_pass = False + self.hr_force = False + self.enable_hr = None + self.hr_scale = None + self.hr_upscaler = None + self.hr_resize_x = 0 + self.hr_resize_y = 0 + self.hr_upscale_to_x = 0 + self.hr_upscale_to_y = 0 + self.truncate_x = 0 + self.truncate_y = 0 + self.applied_old_hires_behavior_to = None + self.refiner_steps = 5 + self.refiner_start = 0 + self.refiner_prompt = '' + self.refiner_negative = '' + self.ops = [] + self.resize_mode: int = resize_mode + self.resize_name: str = resize_name + self.ddim_discretize = shared.opts.ddim_discretize + self.s_min_uncond = shared.opts.s_min_uncond + self.s_churn = shared.opts.s_churn + self.s_noise = shared.opts.s_noise + self.s_min = shared.opts.s_min + self.s_max = shared.opts.s_max + self.s_tmin = shared.opts.s_tmin + self.s_tmax = float('inf') # not representable as a standard ui option + shared.opts.data['clip_skip'] = clip_skip + self.task_args = {} + # a1111 compatibility items + self.refiner_switch_at = 0 + self.hr_prompt = '' + self.all_hr_prompts = [] + self.hr_negative_prompt = '' + self.all_hr_negative_prompts = [] + self.comments = {} + self.is_api = False + self.scripts_value: scripts.ScriptRunner = field(default=None, init=False) + self.script_args_value: list = field(default=None, init=False) + self.scripts_setup_complete: bool = field(default=False, init=False) + # hdr + self.hdr_clamp = hdr_clamp + self.hdr_boundary = hdr_boundary + self.hdr_threshold = hdr_threshold + self.hdr_center = hdr_center + self.hdr_channel_shift = hdr_channel_shift + self.hdr_full_shift = hdr_full_shift + self.hdr_maximize = hdr_maximize + self.hdr_max_center = hdr_max_center + self.hdr_max_boundry = hdr_max_boundry + self.scheduled_prompt: bool = False + self.prompt_embeds = [] + self.positive_pooleds = [] + self.negative_embeds = [] + self.negative_pooleds = [] + + + @property + def sd_model(self): + return shared.sd_model + + @property + def scripts(self): + return self.scripts_value + + @scripts.setter + def scripts(self, value): + self.scripts_value = value + if self.scripts_value and self.script_args_value and not self.scripts_setup_complete: + self.setup_scripts() + + @property + def script_args(self): + return self.script_args_value + + @script_args.setter + def script_args(self, value): + self.script_args_value = value + if self.scripts_value and self.script_args_value and not self.scripts_setup_complete: + self.setup_scripts() + + def setup_scripts(self): + self.scripts_setup_complete = True + self.scripts.setup_scrips(self, is_ui=not self.is_api) + + def comment(self, text): + self.comments[text] = 1 + + def init(self, all_prompts, all_seeds, all_subseeds): + pass + + def sample(self, conditioning, unconditional_conditioning, seeds, subseeds, subseed_strength, prompts): + raise NotImplementedError + + def close(self): + self.sampler = None # pylint: disable=attribute-defined-outside-init + + def get_token_merging_ratio(self, for_hr=False): + if for_hr: + return self.token_merging_ratio_hr or shared.opts.token_merging_ratio_hr or self.token_merging_ratio or shared.opts.token_merging_ratio + return self.token_merging_ratio or shared.opts.token_merging_ratio + + +class StableDiffusionProcessingTxt2Img(StableDiffusionProcessing): + + def __init__(self, enable_hr: bool = False, denoising_strength: float = 0.75, firstphase_width: int = 0, firstphase_height: int = 0, hr_scale: float = 2.0, hr_force: bool = False, hr_upscaler: str = None, hr_second_pass_steps: int = 0, hr_resize_x: int = 0, hr_resize_y: int = 0, refiner_steps: int = 5, refiner_start: float = 0, refiner_prompt: str = '', refiner_negative: str = '', **kwargs): + + super().__init__(**kwargs) + if devices.backend == "ipex" and not torch.xpu.has_fp64_dtype() and os.environ.get('DISABLE_IPEX_1024_WA', None) is None: + width_curse = bool(hr_resize_x == 1024 and self.height * (hr_resize_x / self.width) == 1024) + height_curse = bool(hr_resize_y == 1024 and self.width * (hr_resize_y / self.height) == 1024) + if (width_curse != height_curse) or (height_curse and width_curse): + if width_curse: + hr_resize_x = 1080 + if height_curse: + hr_resize_y = 1080 + if self.width * hr_scale == 1024 and self.height * hr_scale == 1024: + hr_scale = 1080 / self.width + if firstphase_width * hr_scale == 1024 and firstphase_height * hr_scale == 1024: + hr_scale = 1080 / firstphase_width + self.enable_hr = enable_hr + self.denoising_strength = denoising_strength + self.hr_scale = hr_scale + self.hr_upscaler = hr_upscaler + self.hr_force = hr_force + self.hr_second_pass_steps = hr_second_pass_steps + self.hr_resize_x = hr_resize_x + self.hr_resize_y = hr_resize_y + self.hr_upscale_to_x = hr_resize_x + self.hr_upscale_to_y = hr_resize_y + if firstphase_width != 0 or firstphase_height != 0: + self.hr_upscale_to_x = self.width + self.hr_upscale_to_y = self.height + self.width = firstphase_width + self.height = firstphase_height + self.truncate_x = 0 + self.truncate_y = 0 + self.applied_old_hires_behavior_to = None + self.refiner_steps = refiner_steps + self.refiner_start = refiner_start + self.refiner_prompt = refiner_prompt + self.refiner_negative = refiner_negative + self.sampler = None + self.scripts = None + self.script_args = [] + + def init(self, all_prompts, all_seeds, all_subseeds): + if shared.backend == shared.Backend.DIFFUSERS: + shared.sd_model = sd_models.set_diffuser_pipe(self.sd_model, sd_models.DiffusersTaskType.TEXT_2_IMAGE) + self.width = self.width or 512 + self.height = self.height or 512 + + def init_hr(self): + if self.hr_resize_x == 0 and self.hr_resize_y == 0: + self.hr_upscale_to_x = int(self.width * self.hr_scale) + self.hr_upscale_to_y = int(self.height * self.hr_scale) + else: + if self.hr_resize_y == 0: + self.hr_upscale_to_x = self.hr_resize_x + self.hr_upscale_to_y = self.hr_resize_x * self.height // self.width + elif self.hr_resize_x == 0: + self.hr_upscale_to_x = self.hr_resize_y * self.width // self.height + self.hr_upscale_to_y = self.hr_resize_y + else: + target_w = self.hr_resize_x + target_h = self.hr_resize_y + src_ratio = self.width / self.height + dst_ratio = self.hr_resize_x / self.hr_resize_y + if src_ratio < dst_ratio: + self.hr_upscale_to_x = self.hr_resize_x + self.hr_upscale_to_y = self.hr_resize_x * self.height // self.width + else: + self.hr_upscale_to_x = self.hr_resize_y * self.width // self.height + self.hr_upscale_to_y = self.hr_resize_y + self.truncate_x = (self.hr_upscale_to_x - target_w) // 8 + self.truncate_y = (self.hr_upscale_to_y - target_h) // 8 + # special case: the user has chosen to do nothing + if (self.hr_upscale_to_x == self.width and self.hr_upscale_to_y == self.height) or self.hr_upscaler is None or self.hr_upscaler == 'None': + self.is_hr_pass = False + return + self.is_hr_pass = True + hypertile_set(self, hr=True) + shared.state.job_count = 2 * self.n_iter + shared.log.debug(f'Init hires: upscaler="{self.hr_upscaler}" sampler="{self.hr_sampler_name}" resize={self.hr_resize_x}x{self.hr_resize_y} upscale={self.hr_upscale_to_x}x{self.hr_upscale_to_y}') + + def sample(self, conditioning, unconditional_conditioning, seeds, subseeds, subseed_strength, prompts): + from modules import processing_original + return processing_original.sample_txt2img(self, conditioning, unconditional_conditioning, seeds, subseeds, subseed_strength, prompts) + + +class StableDiffusionProcessingImg2Img(StableDiffusionProcessing): + + def __init__(self, init_images: list = None, resize_mode: int = 0, resize_name: str = 'None', denoising_strength: float = 0.3, image_cfg_scale: float = None, mask: Any = None, mask_blur: int = 4, inpainting_fill: int = 0, inpaint_full_res: bool = False, inpaint_full_res_padding: int = 0, inpainting_mask_invert: int = 0, initial_noise_multiplier: float = None, scale_by: float = 1, refiner_steps: int = 5, refiner_start: float = 0, refiner_prompt: str = '', refiner_negative: str = '', **kwargs): + super().__init__(**kwargs) + self.init_images = init_images + self.resize_mode: int = resize_mode + self.resize_name: str = resize_name + self.denoising_strength: float = denoising_strength + self.image_cfg_scale: float = image_cfg_scale + self.init_latent = None + self.image_mask = mask + self.latent_mask = None + self.mask_for_overlay = None + self.mask_blur_x = mask_blur # a1111 compatibility item + self.mask_blur_y = mask_blur # a1111 compatibility item + self.mask_blur = mask_blur + self.inpainting_fill = inpainting_fill + self.inpaint_full_res = inpaint_full_res + self.inpaint_full_res_padding = inpaint_full_res_padding + self.inpainting_mask_invert = inpainting_mask_invert + self.initial_noise_multiplier = shared.opts.initial_noise_multiplier if initial_noise_multiplier is None else initial_noise_multiplier + self.mask = None + self.nmask = None + self.image_conditioning = None + self.refiner_steps = refiner_steps + self.refiner_start = refiner_start + self.refiner_prompt = refiner_prompt + self.refiner_negative = refiner_negative + self.enable_hr = None + self.is_batch = False + self.scale_by = scale_by + self.sampler = None + self.scripts = None + self.script_args = [] + + def init(self, all_prompts, all_seeds, all_subseeds): + if shared.backend == shared.Backend.DIFFUSERS and self.image_mask is not None and not self.is_control: + shared.sd_model = sd_models.set_diffuser_pipe(self.sd_model, sd_models.DiffusersTaskType.INPAINTING) + elif shared.backend == shared.Backend.DIFFUSERS and self.image_mask is None and not self.is_control: + shared.sd_model = sd_models.set_diffuser_pipe(self.sd_model, sd_models.DiffusersTaskType.IMAGE_2_IMAGE) + + if self.sampler_name == "PLMS": + self.sampler_name = 'UniPC' + self.sampler = sd_samplers.create_sampler(self.sampler_name, self.sd_model) + if hasattr(self.sampler, "initialize"): + self.sampler.initialize(self) + + if self.image_mask is not None: + self.ops.append('inpaint') + else: + self.ops.append('img2img') + crop_region = None + + if self.image_mask is not None: + if type(self.image_mask) == list: + self.image_mask = self.image_mask[0] + if shared.backend == shared.Backend.ORIGINAL: # original way of processing mask + self.image_mask = processing_helpers.create_binary_mask(self.image_mask) + if self.inpainting_mask_invert: + self.image_mask = ImageOps.invert(self.image_mask) + if self.mask_blur > 0: + np_mask = np.array(self.image_mask) + kernel_size = 2 * int(2.5 * self.mask_blur + 0.5) + 1 + np_mask = cv2.GaussianBlur(np_mask, (kernel_size, 1), self.mask_blur) + np_mask = cv2.GaussianBlur(np_mask, (1, kernel_size), self.mask_blur) + self.image_mask = Image.fromarray(np_mask) + else: + if hasattr(self, 'init_images'): + self.image_mask = masking.run_mask( + input_image=self.init_images, + input_mask=self.image_mask, + return_type='Grayscale', + mask_blur=self.mask_blur, + mask_padding=self.inpaint_full_res_padding, + segment_enable=False, + invert=self.inpainting_mask_invert, + ) + if self.inpaint_full_res: # mask only inpaint + self.mask_for_overlay = self.image_mask + mask = self.image_mask.convert('L') + crop_region = masking.get_crop_region(np.array(mask), self.inpaint_full_res_padding) + crop_region = masking.expand_crop_region(crop_region, self.width, self.height, mask.width, mask.height) + x1, y1, x2, y2 = crop_region + crop_mask = mask.crop(crop_region) + self.image_mask = images.resize_image(2, crop_mask, self.width, self.height) + self.paste_to = (x1, y1, x2-x1, y2-y1) + else: # full image inpaint + self.image_mask = images.resize_image(self.resize_mode, self.image_mask, self.width, self.height) + np_mask = np.array(self.image_mask) + np_mask = np.clip((np_mask.astype(np.float32)) * 2, 0, 255).astype(np.uint8) + self.mask_for_overlay = Image.fromarray(np_mask) + self.overlay_images = [] + + latent_mask = self.latent_mask if self.latent_mask is not None else self.image_mask + + add_color_corrections = shared.opts.img2img_color_correction and self.color_corrections is None + if add_color_corrections: + self.color_corrections = [] + processed_images = [] + if getattr(self, 'init_images', None) is None: + return + if not isinstance(self.init_images, list): + self.init_images = [self.init_images] + for img in self.init_images: + if img is None: + shared.log.warning(f"Skipping empty image: images={self.init_images}") + continue + self.init_img_hash = hashlib.sha256(img.tobytes()).hexdigest()[0:8] # pylint: disable=attribute-defined-outside-init + self.init_img_width = img.width # pylint: disable=attribute-defined-outside-init + self.init_img_height = img.height # pylint: disable=attribute-defined-outside-init + if shared.opts.save_init_img: + images.save_image(img, path=shared.opts.outdir_init_images, basename=None, forced_filename=self.init_img_hash, suffix="-init-image") + image = images.flatten(img, shared.opts.img2img_background_color) + if self.width is None or self.height is None or self.resize_mode == 0: + self.width, self.height = image.width, image.height + if crop_region is None and self.resize_mode != 4 and self.resize_mode > 0: + if image.width != self.width or image.height != self.height: + image = images.resize_image(self.resize_mode, image, self.width, self.height, self.resize_name) + self.width = image.width + self.height = image.height + if self.image_mask is not None: + try: + image_masked = Image.new('RGBa', (image.width, image.height)) + image_to_paste = image.convert("RGBA").convert("RGBa") + image_to_mask = ImageOps.invert(self.mask_for_overlay.convert('L')) if self.mask_for_overlay is not None else None + image_to_mask = image_to_mask.resize((image.width, image.height), Image.Resampling.BILINEAR) if image_to_mask is not None else None + image_masked.paste(image_to_paste, mask=image_to_mask) + self.overlay_images.append(image_masked.convert('RGBA')) + except Exception as e: + shared.log.error(f"Failed to apply mask to image: {e}") + if crop_region is not None: # crop_region is not None if we are doing inpaint full res + image = image.crop(crop_region) + if image.width != self.width or image.height != self.height: + image = images.resize_image(3, image, self.width, self.height, self.resize_name) + if self.image_mask is not None and self.inpainting_fill != 1: + image = masking.fill(image, latent_mask) + if add_color_corrections: + self.color_corrections.append(processing_helpers.setup_color_correction(image)) + processed_images.append(image) + self.init_images = processed_images + self.batch_size = len(self.init_images) + if self.overlay_images is not None: + self.overlay_images = self.overlay_images * self.batch_size + if self.color_corrections is not None and len(self.color_corrections) == 1: + self.color_corrections = self.color_corrections * self.batch_size + if shared.backend == shared.Backend.DIFFUSERS: + return # we've already set self.init_images and self.mask and we dont need any more processing + elif shared.backend == shared.Backend.ORIGINAL: + self.init_images = [np.moveaxis((np.array(image).astype(np.float32) / 255.0), 2, 0) for image in self.init_images] + if len(self.init_images) == 1: + batch_images = np.expand_dims(self.init_images[0], axis=0).repeat(self.batch_size, axis=0) + elif len(self.init_images) <= self.batch_size: + batch_images = np.array(self.init_images) + image = torch.from_numpy(batch_images) + image = 2. * image - 1. + image = image.to(device=shared.device, dtype=devices.dtype_vae) + self.init_latent = self.sd_model.get_first_stage_encoding(self.sd_model.encode_first_stage(image)) + if self.resize_mode == 4: + self.init_latent = torch.nn.functional.interpolate(self.init_latent, size=(self.height // 8, self.width // 8), mode="bilinear") + if self.image_mask is not None: + init_mask = latent_mask + latmask = init_mask.convert('RGB').resize((self.init_latent.shape[3], self.init_latent.shape[2])) + latmask = np.moveaxis(np.array(latmask, dtype=np.float32), 2, 0) / 255 + latmask = latmask[0] + latmask = np.tile(latmask[None], (4, 1, 1)) + latmask = np.around(latmask) + self.mask = torch.asarray(1.0 - latmask).to(device=shared.device, dtype=self.sd_model.dtype) + self.nmask = torch.asarray(latmask).to(device=shared.device, dtype=self.sd_model.dtype) + if self.inpainting_fill == 2: + self.init_latent = self.init_latent * self.mask + processing_helpers.create_random_tensors(self.init_latent.shape[1:], all_seeds[0:self.init_latent.shape[0]]) * self.nmask + elif self.inpainting_fill == 3: + self.init_latent = self.init_latent * self.mask + self.image_conditioning = processing_helpers.img2img_image_conditioning(self, image, self.init_latent, self.image_mask) + + def sample(self, conditioning, unconditional_conditioning, seeds, subseeds, subseed_strength, prompts): + from modules import processing_original + return processing_original.sample_img2img(self, conditioning, unconditional_conditioning, seeds, subseeds, subseed_strength, prompts) + + def get_token_merging_ratio(self, for_hr=False): + return self.token_merging_ratio or ("token_merging_ratio" in self.override_settings and shared.opts.token_merging_ratio) or shared.opts.token_merging_ratio_img2img or shared.opts.token_merging_ratio + +class StableDiffusionProcessingControl(StableDiffusionProcessingImg2Img): + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.strength = None + self.adapter_conditioning_scale = None + self.adapter_conditioning_factor = None + self.guess_mode = None + self.controlnet_conditioning_scale = None + self.control_guidance_start = None + self.control_guidance_end = None + self.reference_attn = None + self.reference_adain = None + self.attention_auto_machine_weight = None + self.gn_auto_machine_weight = None + self.style_fidelity = None + self.ref_image = None + self.image = None + self.query_weight = None + self.adain_weight = None + self.adapter_conditioning_factor = 1.0 + self.attention = 'Attention' + self.fidelity = 0.5 + self.override = None + self.ip_adapter_name = None + self.ip_adapter_scale = 1.0 + self.ip_adapter_image = None + + def sample(self, conditioning, unconditional_conditioning, seeds, subseeds, subseed_strength, prompts): # abstract + pass + + def init_hr(self): + if self.resize_name == 'None' or self.scale_by == 1.0: + return + self.is_hr_pass = True + self.hr_force = True + self.hr_upscaler = self.resize_name + self.hr_upscale_to_x, self.hr_upscale_to_y = int(self.width * self.scale_by), int(self.height * self.scale_by) + self.hr_upscale_to_x, self.hr_upscale_to_y = 8 * math.ceil(self.hr_upscale_to_x / 8), 8 * math.ceil(self.hr_upscale_to_y / 8) + # hypertile_set(self, hr=True) + shared.state.job_count = 2 * self.n_iter + shared.log.debug(f'Control hires: upscaler="{self.hr_upscaler}" upscale={self.scale_by} size={self.hr_upscale_to_x}x{self.hr_upscale_to_y}') diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index b69f325d9..471161208 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -8,13 +8,12 @@ import torch import torchvision.transforms.functional as TF import diffusers from modules import shared, devices, processing, sd_samplers, sd_models, images, errors, masking, prompt_parser_diffusers, sd_hijack_hypertile, processing_correction, processing_vae +from modules.processing_helpers import resize_init_images, resize_hires, fix_prompts, calculate_base_steps, calculate_hires_steps, calculate_refiner_steps from modules.onnx_impl import preprocess_pipeline as preprocess_onnx_pipeline debug = shared.log.trace if os.environ.get('SD_DIFFUSERS_DEBUG', None) is not None else lambda *args, **kwargs: None debug('Trace: DIFFUSERS') -debug_steps = shared.log.trace if os.environ.get('SD_STEPS_DEBUG', None) is not None else lambda *args, **kwargs: None -debug_steps('Trace: STEPS') def process_diffusers(p: processing.StableDiffusionProcessing): @@ -28,45 +27,6 @@ def process_diffusers(p: processing.StableDiffusionProcessing): def is_refiner_enabled(): return p.enable_hr and p.refiner_steps > 0 and p.refiner_start > 0 and p.refiner_start < 1 and shared.sd_refiner is not None - def resize_images(): - if getattr(p, 'image', None) is not None and getattr(p, 'init_images', None) is None: - p.init_images = [p.image] - if getattr(p, 'init_images', None) is not None and len(p.init_images) > 0: - tgt_width, tgt_height = 8 * math.ceil(p.init_images[0].width / 8), 8 * math.ceil(p.init_images[0].height / 8) - if p.init_images[0].size != (tgt_width, tgt_height): - shared.log.debug(f'Resizing init images: original={p.init_images[0].width}x{p.init_images[0].height} target={tgt_width}x{tgt_height}') - p.init_images = [images.resize_image(1, image, tgt_width, tgt_height, upscaler_name=None) for image in p.init_images] - p.height = tgt_height - p.width = tgt_width - sd_hijack_hypertile.hypertile_set(p) - if getattr(p, 'mask', None) is not None and p.mask.size != (tgt_width, tgt_height): - p.mask = images.resize_image(1, p.mask, tgt_width, tgt_height, upscaler_name=None) - if getattr(p, 'init_mask', None) is not None and p.init_mask.size != (tgt_width, tgt_height): - p.init_mask = images.resize_image(1, p.init_mask, tgt_width, tgt_height, upscaler_name=None) - if getattr(p, 'mask_for_overlay', None) is not None and p.mask_for_overlay.size != (tgt_width, tgt_height): - p.mask_for_overlay = images.resize_image(1, p.mask_for_overlay, tgt_width, tgt_height, upscaler_name=None) - return tgt_width, tgt_height - return p.width, p.height - - def hires_resize(latents): # input=latents output=pil - if not torch.is_tensor(latents): - shared.log.warning('Hires: input is not tensor') - first_pass_images = processing_vae.vae_decode(latents=latents, model=shared.sd_model, full_quality=p.full_quality, output_type='pil') - return first_pass_images - latent_upscaler = shared.latent_upscale_modes.get(p.hr_upscaler, None) - shared.log.info(f'Hires: upscaler={p.hr_upscaler} width={p.hr_upscale_to_x} height={p.hr_upscale_to_y} images={latents.shape[0]}') - if latent_upscaler is not None: - latents = torch.nn.functional.interpolate(latents, size=(p.hr_upscale_to_y // 8, p.hr_upscale_to_x // 8), mode=latent_upscaler["mode"], antialias=latent_upscaler["antialias"]) - first_pass_images = processing_vae.vae_decode(latents=latents, model=shared.sd_model, full_quality=p.full_quality, output_type='pil') - resized_images = [] - for img in first_pass_images: - if latent_upscaler is None: - resized_image = images.resize_image(1, img, p.hr_upscale_to_x, p.hr_upscale_to_y, upscaler_name=p.hr_upscaler) - else: - resized_image = img - resized_images.append(resized_image) - return resized_images - def save_intermediate(latents, suffix): for i in range(len(latents)): from modules.processing import create_infotext @@ -116,27 +76,6 @@ def process_diffusers(p: processing.StableDiffusionProcessing): shared.profiler.step() return kwargs - def fix_prompts(prompts, negative_prompts, prompts_2, negative_prompts_2): - if type(prompts) is str: - prompts = [prompts] - if type(negative_prompts) is str: - negative_prompts = [negative_prompts] - while len(negative_prompts) < len(prompts): - negative_prompts.append(negative_prompts[-1]) - while len(prompts) < len(negative_prompts): - prompts.append(prompts[-1]) - if type(prompts_2) is str: - prompts_2 = [prompts_2] - if type(prompts_2) is list: - while len(prompts_2) < len(prompts): - prompts_2.append(prompts_2[-1]) - if type(negative_prompts_2) is str: - negative_prompts_2 = [negative_prompts_2] - if type(negative_prompts_2) is list: - while len(negative_prompts_2) < len(prompts_2): - negative_prompts_2.append(negative_prompts_2[-1]) - return prompts, negative_prompts, prompts_2, negative_prompts_2 - def task_specific_kwargs(model): task_args = {} is_img2img_model = bool('Zero123' in shared.sd_model.__class__.__name__) @@ -175,7 +114,7 @@ def process_diffusers(p: processing.StableDiffusionProcessing): pass else: # fallback p.mask = TF.to_pil_image(torch.ones_like(TF.to_tensor(p.init_images[0]))).convert("L") - width, height = resize_images() + width, height = resize_init_images(p) task_args = { 'image': p.init_images, 'mask_image': p.mask, @@ -353,8 +292,7 @@ def process_diffusers(p: processing.StableDiffusionProcessing): shared.compiled_model_state.width = compile_width shared.compiled_model_state.batch_size = p.batch_size - # Delete UNET after OpenVINO compile - def openvino_post_compile(op="base"): + def openvino_post_compile(op="base"): # delete unet after OpenVINO compile if shared.opts.cuda_compile and shared.opts.cuda_compile_backend == "openvino_fx": if shared.compiled_model_state.first_pass and op == "base": shared.compiled_model_state.first_pass = False @@ -437,46 +375,6 @@ def process_diffusers(p: processing.StableDiffusionProcessing): use_refiner_start = is_txt2img() and is_refiner_enabled() and not p.is_hr_pass and p.refiner_start > 0 and p.refiner_start < 1 use_denoise_start = not is_txt2img() and p.refiner_start > 0 and p.refiner_start < 1 - def calculate_base_steps(): - if not is_txt2img(): - if use_denoise_start and shared.sd_model_type == 'sdxl': - steps = p.steps // (1 - p.refiner_start) - elif p.denoising_strength > 0: - steps = (p.steps // p.denoising_strength) + 1 - else: - steps = p.steps - elif use_refiner_start and shared.sd_model_type == 'sdxl': - steps = (p.steps // p.refiner_start) + 1 - else: - steps = p.steps - debug_steps(f'Steps: type=base input={p.steps} output={steps} task={sd_models.get_diffusers_task(shared.sd_model)} refiner={use_refiner_start} denoise={p.denoising_strength} model={shared.sd_model_type}') - return max(1, int(steps)) - - def calculate_hires_steps(): - if p.hr_second_pass_steps > 0: - steps = (p.hr_second_pass_steps // p.denoising_strength) + 1 - elif p.denoising_strength > 0: - steps = (p.steps // p.denoising_strength) + 1 - else: - steps = 0 - debug_steps(f'Steps: type=hires input={p.hr_second_pass_steps} output={steps} denoise={p.denoising_strength} model={shared.sd_model_type}') - return max(1, int(steps)) - - def calculate_refiner_steps(): - if "StableDiffusionXL" in shared.sd_refiner.__class__.__name__: - if p.refiner_start > 0 and p.refiner_start < 1: - #steps = p.refiner_steps // (1 - p.refiner_start) # SDXL with denoise strenght - steps = (p.refiner_steps // (1 - p.refiner_start) // 2) + 1 - elif p.denoising_strength > 0: - steps = (p.refiner_steps // p.denoising_strength) + 1 - else: - steps = 0 - else: - #steps = p.refiner_steps # SD 1.5 with denoise strenght - steps = (p.refiner_steps * 1.25) + 1 - debug_steps(f'Steps: type=refiner input={p.refiner_steps} output={steps} start={p.refiner_start} denoise={p.denoising_strength}') - return max(1, int(steps)) - shared.sd_model = update_pipeline(shared.sd_model, p) base_args = set_pipeline_args( model=shared.sd_model, @@ -484,7 +382,7 @@ def process_diffusers(p: processing.StableDiffusionProcessing): negative_prompts=p.negative_prompts, prompts_2=[p.refiner_prompt] if len(p.refiner_prompt) > 0 else p.prompts, negative_prompts_2=[p.refiner_negative] if len(p.refiner_negative) > 0 else p.negative_prompts, - num_inference_steps=calculate_base_steps(), + num_inference_steps=calculate_base_steps(p, use_refiner_start=use_refiner_start, use_denoise_start=use_denoise_start), eta=shared.opts.scheduler_eta, guidance_scale=p.cfg_scale, guidance_rescale=p.diffusers_guidance_rescale, @@ -545,7 +443,7 @@ def process_diffusers(p: processing.StableDiffusionProcessing): if shared.opts.save and not p.do_not_save_samples and shared.opts.save_images_before_highres_fix and hasattr(shared.sd_model, 'vae'): save_intermediate(latents=output.images, suffix="-before-hires") shared.state.job = 'upscale' - output.images = hires_resize(latents=output.images) + output.images = resize_hires(p, latents=output.images) if (latent_scale_mode is not None or p.hr_force) and p.denoising_strength > 0: p.ops.append('hires') shared.sd_model = sd_models.set_diffuser_pipe(shared.sd_model, sd_models.DiffusersTaskType.IMAGE_2_IMAGE) @@ -559,7 +457,7 @@ def process_diffusers(p: processing.StableDiffusionProcessing): negative_prompts=[p.refiner_negative] if len(p.refiner_negative) > 0 else p.negative_prompts, prompts_2=[p.refiner_prompt] if len(p.refiner_prompt) > 0 else p.prompts, negative_prompts_2=[p.refiner_negative] if len(p.refiner_negative) > 0 else p.negative_prompts, - num_inference_steps=calculate_hires_steps(), + num_inference_steps=calculate_hires_steps(p), eta=shared.opts.scheduler_eta, guidance_scale=p.image_cfg_scale if p.image_cfg_scale is not None else p.cfg_scale, guidance_rescale=p.diffusers_guidance_rescale, @@ -617,7 +515,7 @@ def process_diffusers(p: processing.StableDiffusionProcessing): model=shared.sd_refiner, prompts=[p.refiner_prompt] if len(p.refiner_prompt) > 0 else p.prompts[i], negative_prompts=[p.refiner_negative] if len(p.refiner_negative) > 0 else p.negative_prompts[i], - num_inference_steps=calculate_refiner_steps(), + num_inference_steps=calculate_refiner_steps(p), eta=shared.opts.scheduler_eta, # strength=p.denoising_strength, noise_level=noise_level, # StableDiffusionUpscalePipeline only diff --git a/modules/processing_helpers.py b/modules/processing_helpers.py new file mode 100644 index 000000000..0da7ab470 --- /dev/null +++ b/modules/processing_helpers.py @@ -0,0 +1,428 @@ +import os +import math +import random +import warnings +from einops import repeat, rearrange +import torch +import numpy as np +import cv2 +from PIL import Image +from skimage import exposure +from blendmodes.blend import blendLayers, BlendType +from modules import shared, devices, images, sd_models, sd_samplers, sd_hijack_hypertile, processing_vae + + +debug = shared.log.trace if os.environ.get('SD_PROCESS_DEBUG', None) is not None else lambda *args, **kwargs: None +debug_steps = shared.log.trace if os.environ.get('SD_STEPS_DEBUG', None) is not None else lambda *args, **kwargs: None +debug_steps('Trace: STEPS') + + +def setup_color_correction(image): + debug("Calibrating color correction") + correction_target = cv2.cvtColor(np.asarray(image.copy()), cv2.COLOR_RGB2LAB) + return correction_target + + +def apply_color_correction(correction, original_image): + shared.log.debug(f"Applying color correction: correction={correction.shape} image={original_image}") + np_image = np.asarray(original_image) + np_recolor = cv2.cvtColor(np_image, cv2.COLOR_RGB2LAB) + np_match = exposure.match_histograms(np_recolor, correction, channel_axis=2) + np_output = cv2.cvtColor(np_match, cv2.COLOR_LAB2RGB) + image = Image.fromarray(np_output.astype("uint8")) + image = blendLayers(image, original_image, BlendType.LUMINOSITY) + return image + + +def apply_overlay(image: Image, paste_loc, index, overlays): + debug(f'Apply overlay: image={image} loc={paste_loc} index={index} overlays={overlays}') + if overlays is None or index >= len(overlays): + return image + overlay = overlays[index] + if paste_loc is not None: + x, y, w, h = paste_loc + if image.width != w or image.height != h or x != 0 or y != 0: + base_image = Image.new('RGBA', (overlay.width, overlay.height)) + image = images.resize_image(2, image, w, h) + base_image.paste(image, (x, y)) + image = base_image + image = image.convert('RGBA') + image.alpha_composite(overlay) + image = image.convert('RGB') + return image + + +def create_binary_mask(image): + if image.mode == 'RGBA' and image.getextrema()[-1] != (255, 255): + image = image.split()[-1].convert("L").point(lambda x: 255 if x > 128 else 0) + else: + image = image.convert('L') + return image + + +def images_tensor_to_samples(image, approximation=None, model=None): # pylint: disable=unused-argument + if model is None: + model = shared.sd_model + model.first_stage_model.to(devices.dtype_vae) + image = image.to(shared.device, dtype=devices.dtype_vae) + image = image * 2 - 1 + if len(image) > 1: + x_latent = torch.stack([ + model.get_first_stage_encoding(model.encode_first_stage(torch.unsqueeze(img, 0)))[0] + for img in image + ]) + else: + x_latent = model.get_first_stage_encoding(model.encode_first_stage(image)) + return x_latent + + +def get_sampler_name(sampler_index: int, img: bool = False) -> str: + sampler_index = sampler_index or 0 + if len(sd_samplers.samplers) > sampler_index: + sampler_name = sd_samplers.samplers[sampler_index].name + else: + sampler_name = "UniPC" + shared.log.warning(f'Sampler not found: index={sampler_index} available={[s.name for s in sd_samplers.samplers]} fallback={sampler_name}') + if img and sampler_name == "PLMS": + sampler_name = "UniPC" + shared.log.warning(f'Sampler not compatible: name=PLMS fallback={sampler_name}') + return sampler_name + + +def slerp(val, low, high): # from https://discuss.pytorch.org/t/help-regarding-slerp-function-for-generative-model-sampling/32475/3 + low_norm = low/torch.norm(low, dim=1, keepdim=True) + high_norm = high/torch.norm(high, dim=1, keepdim=True) + dot = (low_norm*high_norm).sum(1) + + if dot.mean() > 0.9995: + return low * val + high * (1 - val) + + omega = torch.acos(dot) + so = torch.sin(omega) + res = (torch.sin((1.0-val)*omega)/so).unsqueeze(1)*low + (torch.sin(val*omega)/so).unsqueeze(1) * high + return res + + +def create_random_tensors(shape, seeds, subseeds=None, subseed_strength=0.0, seed_resize_from_h=0, seed_resize_from_w=0, p=None): + eta_noise_seed_delta = shared.opts.eta_noise_seed_delta or 0 + xs = [] + # if we have multiple seeds, this means we are working with batch size>1; this then + # enables the generation of additional tensors with noise that the sampler will use during its processing. + # Using those pre-generated tensors instead of simple torch.randn allows a batch with seeds [100, 101] to + # produce the same images as with two batches [100], [101]. + if p is not None and p.sampler is not None and (len(seeds) > 1 and shared.opts.enable_batch_seeds or eta_noise_seed_delta > 0): + sampler_noises = [[] for _ in range(p.sampler.number_of_needed_noises(p))] + else: + sampler_noises = None + for i, seed in enumerate(seeds): + noise_shape = shape if seed_resize_from_h <= 0 or seed_resize_from_w <= 0 else (shape[0], seed_resize_from_h//8, seed_resize_from_w//8) + subnoise = None + if subseeds is not None: + subseed = 0 if i >= len(subseeds) else subseeds[i] + subnoise = devices.randn(subseed, noise_shape) + # randn results depend on device; gpu and cpu get different results for same seed; + # the way I see it, it's better to do this on CPU, so that everyone gets same result; + # but the original script had it like this, so I do not dare change it for now because + # it will break everyone's seeds. + noise = devices.randn(seed, noise_shape) + if subnoise is not None: + noise = slerp(subseed_strength, noise, subnoise) + if noise_shape != shape: + x = devices.randn(seed, shape) + dx = (shape[2] - noise_shape[2]) // 2 + dy = (shape[1] - noise_shape[1]) // 2 + w = noise_shape[2] if dx >= 0 else noise_shape[2] + 2 * dx + h = noise_shape[1] if dy >= 0 else noise_shape[1] + 2 * dy + tx = 0 if dx < 0 else dx + ty = 0 if dy < 0 else dy + dx = max(-dx, 0) + dy = max(-dy, 0) + x[:, ty:ty+h, tx:tx+w] = noise[:, dy:dy+h, dx:dx+w] + noise = x + if sampler_noises is not None: + cnt = p.sampler.number_of_needed_noises(p) + if eta_noise_seed_delta > 0: + torch.manual_seed(seed + eta_noise_seed_delta) + for j in range(cnt): + sampler_noises[j].append(devices.randn_without_seed(tuple(noise_shape))) + xs.append(noise) + if sampler_noises is not None: + p.sampler.sampler_noises = [torch.stack(n).to(shared.device) for n in sampler_noises] + x = torch.stack(xs).to(shared.device) + return x + + +def decode_first_stage(model, x, full_quality=True): + if not shared.opts.keep_incomplete and (shared.state.skipped or shared.state.interrupted): + shared.log.debug(f'Decode VAE: skipped={shared.state.skipped} interrupted={shared.state.interrupted}') + x_sample = torch.zeros((len(x), 3, x.shape[2] * 8, x.shape[3] * 8), dtype=devices.dtype_vae, device=devices.device) + return x_sample + prev_job = shared.state.job + shared.state.job = 'vae' + with devices.autocast(disable = x.dtype==devices.dtype_vae): + try: + if full_quality: + if hasattr(model, 'decode_first_stage'): + x_sample = model.decode_first_stage(x) + elif hasattr(model, 'vae'): + x_sample = model.vae(x) + else: + x_sample = x + shared.log.error('Decode VAE unknown model') + else: + from modules import sd_vae_taesd + x_sample = torch.zeros((len(x), 3, x.shape[2] * 8, x.shape[3] * 8), dtype=devices.dtype_vae, device=devices.device) + for i in range(len(x_sample)): + x_sample[i] = sd_vae_taesd.decode(x[i]) + except Exception as e: + x_sample = x + shared.log.error(f'Decode VAE: {e}') + shared.state.job = prev_job + return x_sample + + +def get_fixed_seed(seed): + if seed is None or seed == '' or seed == -1: + return int(random.randrange(4294967294)) + return seed + + +def fix_seed(p): + p.seed = get_fixed_seed(p.seed) + p.subseed = get_fixed_seed(p.subseed) + + +def old_hires_fix_first_pass_dimensions(width, height): + """old algorithm for auto-calculating first pass size""" + desired_pixel_count = 512 * 512 + actual_pixel_count = width * height + scale = math.sqrt(desired_pixel_count / actual_pixel_count) + width = math.ceil(scale * width / 64) * 64 + height = math.ceil(scale * height / 64) * 64 + return width, height + + +def txt2img_image_conditioning(p, x, width=None, height=None): + width = width or p.width + height = height or p.height + if p.sd_model.model.conditioning_key in {'hybrid', 'concat'}: # Inpainting models + image_conditioning = torch.zeros(x.shape[0], 3, height, width, device=x.device) + image_conditioning = p.sd_model.get_first_stage_encoding(p.sd_model.encode_first_stage(image_conditioning)) + image_conditioning = torch.nn.functional.pad(image_conditioning, (0, 0, 0, 0, 1, 0), value=1.0) # pylint: disable=not-callable + image_conditioning = image_conditioning.to(x.dtype) + return image_conditioning + elif p.sd_model.model.conditioning_key == "crossattn-adm": # UnCLIP models + return x.new_zeros(x.shape[0], 2*p.sd_model.noise_augmentor.time_embed.dim, dtype=x.dtype, device=x.device) + else: + return x.new_zeros(x.shape[0], 5, 1, 1, dtype=x.dtype, device=x.device) + + +def img2img_image_conditioning(p, source_image, latent_image, image_mask=None): + from ldm.models.diffusion.ddpm import LatentDepth2ImageDiffusion + source_image = devices.cond_cast_float(source_image) + + def depth2img_image_conditioning(source_image): + # Use the AddMiDaS helper to Format our source image to suit the MiDaS model + from ldm.data.util import AddMiDaS + transformer = AddMiDaS(model_type="dpt_hybrid") + transformed = transformer({"jpg": rearrange(source_image[0], "c h w -> h w c")}) + midas_in = torch.from_numpy(transformed["midas_in"][None, ...]).to(device=shared.device) + midas_in = repeat(midas_in, "1 ... -> n ...", n=p.batch_size) + conditioning_image = p.sd_model.get_first_stage_encoding(p.sd_model.encode_first_stage(source_image)) + conditioning = torch.nn.functional.interpolate( + p.sd_model.depth_model(midas_in), + size=conditioning_image.shape[2:], + mode="bicubic", + align_corners=False, + ) + (depth_min, depth_max) = torch.aminmax(conditioning) + conditioning = 2. * (conditioning - depth_min) / (depth_max - depth_min) - 1. + return conditioning + + def edit_image_conditioning(source_image): + conditioning_image = p.sd_model.encode_first_stage(source_image).mode() + return conditioning_image + + def unclip_image_conditioning(source_image): + c_adm = p.sd_model.embedder(source_image) + if p.sd_model.noise_augmentor is not None: + noise_level = 0 + c_adm, noise_level_emb = p.sd_model.noise_augmentor(c_adm, noise_level=repeat(torch.tensor([noise_level]).to(c_adm.device), '1 -> b', b=c_adm.shape[0])) + c_adm = torch.cat((c_adm, noise_level_emb), 1) + return c_adm + + def inpainting_image_conditioning(source_image, latent_image, image_mask=None): + # Handle the different mask inputs + if image_mask is not None: + if torch.is_tensor(image_mask): + conditioning_mask = image_mask + else: + conditioning_mask = np.array(image_mask.convert("L")) + conditioning_mask = conditioning_mask.astype(np.float32) / 255.0 + conditioning_mask = torch.from_numpy(conditioning_mask[None, None]) + # Inpainting model uses a discretized mask as input, so we round to either 1.0 or 0.0 + conditioning_mask = torch.round(conditioning_mask) + else: + conditioning_mask = source_image.new_ones(1, 1, *source_image.shape[-2:]) + # Create another latent image, this time with a masked version of the original input. + # Smoothly interpolate between the masked and unmasked latent conditioning image using a parameter. + conditioning_mask = conditioning_mask.to(device=source_image.device, dtype=source_image.dtype) + conditioning_image = torch.lerp( + source_image, + source_image * (1.0 - conditioning_mask), + getattr(p, "inpainting_mask_weight", shared.opts.inpainting_mask_weight) + ) + # Encode the new masked image using first stage of network. + conditioning_image = p.sd_model.get_first_stage_encoding(p.sd_model.encode_first_stage(conditioning_image)) + # Create the concatenated conditioning tensor to be fed to `c_concat` + conditioning_mask = torch.nn.functional.interpolate(conditioning_mask, size=latent_image.shape[-2:]) + conditioning_mask = conditioning_mask.expand(conditioning_image.shape[0], -1, -1, -1) + image_conditioning = torch.cat([conditioning_mask, conditioning_image], dim=1) + image_conditioning = image_conditioning.to(device=shared.device, dtype=source_image.dtype) + return image_conditioning + + def diffusers_image_conditioning(_source_image, latent_image, _image_mask=None): + # shared.log.warning('Diffusers not implemented: img2img_image_conditioning') + return latent_image.new_zeros(latent_image.shape[0], 5, 1, 1) + + # HACK: Using introspection as the Depth2Image model doesn't appear to uniquely + # identify itself with a field common to all models. The conditioning_key is also hybrid. + if shared.backend == shared.Backend.DIFFUSERS: + return diffusers_image_conditioning(source_image, latent_image, image_mask) + if isinstance(p.sd_model, LatentDepth2ImageDiffusion): + return depth2img_image_conditioning(source_image) + if hasattr(p.sd_model, 'cond_stage_key') and p.sd_model.cond_stage_key == "edit": + return edit_image_conditioning(source_image) + if hasattr(p.sampler, 'conditioning_key') and p.sampler.conditioning_key in {'hybrid', 'concat'}: + return inpainting_image_conditioning(source_image, latent_image, image_mask=image_mask) + if hasattr(p.sampler, 'conditioning_key') and p.sampler.conditioning_key == "crossattn-adm": + return unclip_image_conditioning(source_image) + # Dummy zero conditioning if we're not using inpainting or depth model. + return latent_image.new_zeros(latent_image.shape[0], 5, 1, 1) + + +def validate_sample(tensor): + if not isinstance(tensor, np.ndarray) and not isinstance(tensor, torch.Tensor): + return tensor + if tensor.dtype == torch.bfloat16: # numpy does not support bf16 + tensor = tensor.to(torch.float16) + if isinstance(tensor, torch.Tensor) and hasattr(tensor, 'detach'): + sample = tensor.detach().cpu().numpy() + elif isinstance(tensor, np.ndarray): + sample = tensor + else: + shared.log.warning(f'Unknown sample type: {type(tensor)}') + sample = 255.0 * np.moveaxis(sample, 0, 2) if shared.backend == shared.Backend.ORIGINAL else 255.0 * sample + with warnings.catch_warnings(record=True) as w: + cast = sample.astype(np.uint8) + if len(w) > 0: + nans = np.isnan(sample).sum() + shared.log.error(f'Failed to validate samples: sample={sample.shape} invalid={nans}') + cast = np.nan_to_num(sample) + minimum, maximum, mean = np.min(cast), np.max(cast), np.mean(cast) + cast = cast.astype(np.uint8) + shared.log.warning(f'Attempted to correct samples: min={minimum:.2f} max={maximum:.2f} mean={mean:.2f}') + return cast + + +def resize_init_images(p): + if getattr(p, 'image', None) is not None and getattr(p, 'init_images', None) is None: + p.init_images = [p.image] + if getattr(p, 'init_images', None) is not None and len(p.init_images) > 0: + tgt_width, tgt_height = 8 * math.ceil(p.init_images[0].width / 8), 8 * math.ceil(p.init_images[0].height / 8) + if p.init_images[0].size != (tgt_width, tgt_height): + shared.log.debug(f'Resizing init images: original={p.init_images[0].width}x{p.init_images[0].height} target={tgt_width}x{tgt_height}') + p.init_images = [images.resize_image(1, image, tgt_width, tgt_height, upscaler_name=None) for image in p.init_images] + p.height = tgt_height + p.width = tgt_width + sd_hijack_hypertile.hypertile_set(p) + if getattr(p, 'mask', None) is not None and p.mask.size != (tgt_width, tgt_height): + p.mask = images.resize_image(1, p.mask, tgt_width, tgt_height, upscaler_name=None) + if getattr(p, 'init_mask', None) is not None and p.init_mask.size != (tgt_width, tgt_height): + p.init_mask = images.resize_image(1, p.init_mask, tgt_width, tgt_height, upscaler_name=None) + if getattr(p, 'mask_for_overlay', None) is not None and p.mask_for_overlay.size != (tgt_width, tgt_height): + p.mask_for_overlay = images.resize_image(1, p.mask_for_overlay, tgt_width, tgt_height, upscaler_name=None) + return tgt_width, tgt_height + return p.width, p.height + + +def resize_hires(p, latents): # input=latents output=pil + if not torch.is_tensor(latents): + shared.log.warning('Hires: input is not tensor') + first_pass_images = processing_vae.vae_decode(latents=latents, model=shared.sd_model, full_quality=p.full_quality, output_type='pil') + return first_pass_images + latent_upscaler = shared.latent_upscale_modes.get(p.hr_upscaler, None) + shared.log.info(f'Hires: upscaler={p.hr_upscaler} width={p.hr_upscale_to_x} height={p.hr_upscale_to_y} images={latents.shape[0]}') + if latent_upscaler is not None: + latents = torch.nn.functional.interpolate(latents, size=(p.hr_upscale_to_y // 8, p.hr_upscale_to_x // 8), mode=latent_upscaler["mode"], antialias=latent_upscaler["antialias"]) + first_pass_images = processing_vae.vae_decode(latents=latents, model=shared.sd_model, full_quality=p.full_quality, output_type='pil') + resized_images = [] + for img in first_pass_images: + if latent_upscaler is None: + resized_image = images.resize_image(1, img, p.hr_upscale_to_x, p.hr_upscale_to_y, upscaler_name=p.hr_upscaler) + else: + resized_image = img + resized_images.append(resized_image) + return resized_images + +def fix_prompts(prompts, negative_prompts, prompts_2, negative_prompts_2): + if type(prompts) is str: + prompts = [prompts] + if type(negative_prompts) is str: + negative_prompts = [negative_prompts] + while len(negative_prompts) < len(prompts): + negative_prompts.append(negative_prompts[-1]) + while len(prompts) < len(negative_prompts): + prompts.append(prompts[-1]) + if type(prompts_2) is str: + prompts_2 = [prompts_2] + if type(prompts_2) is list: + while len(prompts_2) < len(prompts): + prompts_2.append(prompts_2[-1]) + if type(negative_prompts_2) is str: + negative_prompts_2 = [negative_prompts_2] + if type(negative_prompts_2) is list: + while len(negative_prompts_2) < len(prompts_2): + negative_prompts_2.append(negative_prompts_2[-1]) + return prompts, negative_prompts, prompts_2, negative_prompts_2 + +def calculate_base_steps(p, use_denoise_start, use_refiner_start): + is_txt2img = sd_models.get_diffusers_task(shared.sd_model) == sd_models.DiffusersTaskType.TEXT_2_IMAGE + if not is_txt2img: + if use_denoise_start and shared.sd_model_type == 'sdxl': + steps = p.steps // (1 - p.refiner_start) + elif p.denoising_strength > 0: + steps = (p.steps // p.denoising_strength) + 1 + else: + steps = p.steps + elif use_refiner_start and shared.sd_model_type == 'sdxl': + steps = (p.steps // p.refiner_start) + 1 + else: + steps = p.steps + debug_steps(f'Steps: type=base input={p.steps} output={steps} task={sd_models.get_diffusers_task(shared.sd_model)} refiner={use_refiner_start} denoise={p.denoising_strength} model={shared.sd_model_type}') + return max(1, int(steps)) + +def calculate_hires_steps(p): + if p.hr_second_pass_steps > 0: + steps = (p.hr_second_pass_steps // p.denoising_strength) + 1 + elif p.denoising_strength > 0: + steps = (p.steps // p.denoising_strength) + 1 + else: + steps = 0 + debug_steps(f'Steps: type=hires input={p.hr_second_pass_steps} output={steps} denoise={p.denoising_strength} model={shared.sd_model_type}') + return max(1, int(steps)) + +def calculate_refiner_steps(p): + if "StableDiffusionXL" in shared.sd_refiner.__class__.__name__: + if p.refiner_start > 0 and p.refiner_start < 1: + #steps = p.refiner_steps // (1 - p.refiner_start) # SDXL with denoise strenght + steps = (p.refiner_steps // (1 - p.refiner_start) // 2) + 1 + elif p.denoising_strength > 0: + steps = (p.refiner_steps // p.denoising_strength) + 1 + else: + steps = 0 + else: + #steps = p.refiner_steps # SD 1.5 with denoise strenght + steps = (p.refiner_steps * 1.25) + 1 + debug_steps(f'Steps: type=refiner input={p.refiner_steps} output={steps} start={p.refiner_start} denoise={p.denoising_strength}') + return max(1, int(steps)) diff --git a/modules/processing_info.py b/modules/processing_info.py new file mode 100644 index 000000000..4e1a860da --- /dev/null +++ b/modules/processing_info.py @@ -0,0 +1,143 @@ +import os +from installer import git_commit +from modules import shared, sd_samplers_common, sd_vae, generation_parameters_copypaste +from modules.processing_class import StableDiffusionProcessing + + +if shared.backend == shared.Backend.ORIGINAL: + from modules import sd_hijack +else: + sd_hijack = None + + +def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=None, all_subseeds=None, comments=None, iteration=0, position_in_batch=0, index=None, all_negative_prompts=None): + if p is None: + shared.log.warning('Processing info: no data') + return '' + if not hasattr(shared.sd_model, 'sd_checkpoint_info'): + return '' + if index is None: + index = position_in_batch + iteration * p.batch_size + if all_prompts is None: + all_prompts = p.all_prompts or [p.prompt] + if all_negative_prompts is None: + all_negative_prompts = p.all_negative_prompts or [p.negative_prompt] + if all_seeds is None: + all_seeds = p.all_seeds or [p.seed] + if all_subseeds is None: + all_subseeds = p.all_subseeds or [p.subseed] + while len(all_prompts) <= index: + all_prompts.append(all_prompts[-1]) + while len(all_seeds) <= index: + all_seeds.append(all_seeds[-1]) + while len(all_subseeds) <= index: + all_subseeds.append(all_subseeds[-1]) + while len(all_negative_prompts) <= index: + all_negative_prompts.append(all_negative_prompts[-1]) + comment = ', '.join(comments) if comments is not None and type(comments) is list else None + ops = list(set(p.ops)) + ops.reverse() + args = { + # basic + "Steps": p.steps, + "Seed": all_seeds[index], + "Sampler": p.sampler_name, + "CFG scale": p.cfg_scale, + "Size": f"{p.width}x{p.height}" if hasattr(p, 'width') and hasattr(p, 'height') else None, + "Batch": f'{p.n_iter}x{p.batch_size}' if p.n_iter > 1 or p.batch_size > 1 else None, + "Index": f'{p.iteration + 1}x{index + 1}' if (p.n_iter > 1 or p.batch_size > 1) and index >= 0 else None, + "Parser": shared.opts.prompt_attention, + "Model": None if (not shared.opts.add_model_name_to_info) or (not shared.sd_model.sd_checkpoint_info.model_name) else shared.sd_model.sd_checkpoint_info.model_name.replace(',', '').replace(':', ''), + "Model hash": getattr(p, 'sd_model_hash', None if (not shared.opts.add_model_hash_to_info) or (not shared.sd_model.sd_model_hash) else shared.sd_model.sd_model_hash), + "VAE": (None if not shared.opts.add_model_name_to_info or sd_vae.loaded_vae_file is None else os.path.splitext(os.path.basename(sd_vae.loaded_vae_file))[0]) if p.full_quality else 'TAESD', + "Seed resize from": None if p.seed_resize_from_w == 0 or p.seed_resize_from_h == 0 else f"{p.seed_resize_from_w}x{p.seed_resize_from_h}", + "Clip skip": p.clip_skip if p.clip_skip > 1 else None, + "Prompt2": p.refiner_prompt if len(p.refiner_prompt) > 0 else None, + "Negative2": p.refiner_negative if len(p.refiner_negative) > 0 else None, + "Styles": "; ".join(p.styles) if p.styles is not None and len(p.styles) > 0 else None, + "Tiling": p.tiling if p.tiling else None, + # sdnext + "Backend": 'Diffusers' if shared.backend == shared.Backend.DIFFUSERS else 'Original', + "App": 'SD.Next', + "Version": git_commit, + "Comment": comment, + "Operations": '; '.join(ops).replace('"', '') if len(p.ops) > 0 else 'none', + } + if 'txt2img' in p.ops: + pass + if shared.backend == shared.Backend.ORIGINAL: + args["Variation seed"] = all_subseeds[index] if p.subseed_strength > 0 else None + args["Variation strength"] = p.subseed_strength if p.subseed_strength > 0 else None + if 'hires' in p.ops or 'upscale' in p.ops: + args["Second pass"] = p.enable_hr + args["Hires force"] = p.hr_force + args["Hires steps"] = p.hr_second_pass_steps + args["Hires upscaler"] = p.hr_upscaler + args["Hires upscale"] = p.hr_scale + args["Hires resize"] = f"{p.hr_resize_x}x{p.hr_resize_y}" + args["Hires size"] = f"{p.hr_upscale_to_x}x{p.hr_upscale_to_y}" + args["Denoising strength"] = p.denoising_strength + args["Hires sampler"] = p.hr_sampler_name + args["Image CFG scale"] = p.image_cfg_scale + 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(':', '') + args['Image CFG scale'] = p.image_cfg_scale + args['Refiner steps'] = p.refiner_steps + args['Refiner start'] = p.refiner_start + args["Hires steps"] = p.hr_second_pass_steps + args["Hires sampler"] = p.hr_sampler_name + 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) + args['Resize scale'] = getattr(p, 'scale_by', None) + args["Mask weight"] = getattr(p, "inpainting_mask_weight", shared.opts.inpainting_mask_weight) if p.is_using_inpainting_conditioning else None + args["Denoising strength"] = getattr(p, 'denoising_strength', None) + if args["Size"] is None: + args["Size"] = args["Init image size"] + # lookup by index + if getattr(p, 'resize_mode', None) is not None: + args['Resize mode'] = shared.resize_modes[p.resize_mode] if shared.resize_modes[p.resize_mode] != 'None' else None + if 'face' in p.ops: + args["Face restoration"] = shared.opts.face_restoration_model + if 'color' in p.ops: + args["Color correction"] = True + # embeddings + if sd_hijack is not None and hasattr(sd_hijack.model_hijack, 'embedding_db') and len(sd_hijack.model_hijack.embedding_db.embeddings_used) > 0: # this is for original hijaacked models only, diffusers are handled separately + args["Embeddings"] = ', '.join(sd_hijack.model_hijack.embedding_db.embeddings_used) + # samplers + args["Sampler ENSD"] = shared.opts.eta_noise_seed_delta if shared.opts.eta_noise_seed_delta != 0 and sd_samplers_common.is_sampler_using_eta_noise_seed_delta(p) else None + args["Sampler ENSM"] = p.initial_noise_multiplier if getattr(p, 'initial_noise_multiplier', 1.0) != 1.0 else None + args['Sampler order'] = shared.opts.schedulers_solver_order if shared.opts.schedulers_solver_order != shared.opts.data_labels.get('schedulers_solver_order').default else None + if shared.backend == shared.Backend.DIFFUSERS: + args['Sampler beta schedule'] = shared.opts.schedulers_beta_schedule if shared.opts.schedulers_beta_schedule != shared.opts.data_labels.get('schedulers_beta_schedule').default else None + args['Sampler beta start'] = shared.opts.schedulers_beta_start if shared.opts.schedulers_beta_start != shared.opts.data_labels.get('schedulers_beta_start').default else None + args['Sampler beta end'] = shared.opts.schedulers_beta_end if shared.opts.schedulers_beta_end != shared.opts.data_labels.get('schedulers_beta_end').default else None + args['Sampler DPM solver'] = shared.opts.schedulers_dpm_solver if shared.opts.schedulers_dpm_solver != shared.opts.data_labels.get('schedulers_dpm_solver').default else None + if shared.backend == shared.Backend.ORIGINAL: + args['Sampler brownian'] = shared.opts.schedulers_brownian_noise if shared.opts.schedulers_brownian_noise != shared.opts.data_labels.get('schedulers_brownian_noise').default else None + args['Sampler discard'] = shared.opts.schedulers_discard_penultimate if shared.opts.schedulers_discard_penultimate != shared.opts.data_labels.get('schedulers_discard_penultimate').default else None + args['Sampler dyn threshold'] = shared.opts.schedulers_use_thresholding if shared.opts.schedulers_use_thresholding != shared.opts.data_labels.get('schedulers_use_thresholding').default else None + args['Sampler karras'] = shared.opts.schedulers_use_karras if shared.opts.schedulers_use_karras != shared.opts.data_labels.get('schedulers_use_karras').default else None + args['Sampler low order'] = shared.opts.schedulers_use_loworder if shared.opts.schedulers_use_loworder != shared.opts.data_labels.get('schedulers_use_loworder').default else None + args['Sampler quantization'] = shared.opts.enable_quantization if shared.opts.enable_quantization != shared.opts.data_labels.get('enable_quantization').default else None + args['Sampler sigma'] = shared.opts.schedulers_sigma if shared.opts.schedulers_sigma != shared.opts.data_labels.get('schedulers_sigma').default else None + args['Sampler sigma min'] = shared.opts.s_min if shared.opts.s_min != shared.opts.data_labels.get('s_min').default else None + args['Sampler sigma max'] = shared.opts.s_max if shared.opts.s_max != shared.opts.data_labels.get('s_max').default else None + args['Sampler sigma churn'] = shared.opts.s_churn if shared.opts.s_churn != shared.opts.data_labels.get('s_churn').default else None + args['Sampler sigma uncond'] = shared.opts.s_churn if shared.opts.s_churn != shared.opts.data_labels.get('s_churn').default else None + args['Sampler sigma noise'] = shared.opts.s_noise if shared.opts.s_noise != shared.opts.data_labels.get('s_noise').default else None + args['Sampler sigma tmin'] = shared.opts.s_tmin if shared.opts.s_tmin != shared.opts.data_labels.get('s_tmin').default else None + # tome + token_merging_ratio = p.get_token_merging_ratio() + token_merging_ratio_hr = p.get_token_merging_ratio(for_hr=True) if p.enable_hr else None + args['ToMe'] = token_merging_ratio if token_merging_ratio != 0 else None + args['ToMe hires'] = token_merging_ratio_hr if token_merging_ratio_hr != 0 else None + + args.update(p.extra_generation_params) + params_text = ", ".join([k if k == v else f'{k}: {generation_parameters_copypaste.quote(v)}' for k, v in args.items() if v is not None]) + negative_prompt_text = f"\nNegative prompt: {all_negative_prompts[index]}" if all_negative_prompts[index] else "" + infotext = f"{all_prompts[index]}{negative_prompt_text}\n{params_text}".strip() + return infotext diff --git a/modules/processing_original.py b/modules/processing_original.py new file mode 100644 index 000000000..6e6d2bafe --- /dev/null +++ b/modules/processing_original.py @@ -0,0 +1,163 @@ +import torch +import numpy as np +from PIL import Image +from modules import shared, devices, processing, images, sd_models, sd_vae, sd_samplers, processing_helpers, prompt_parser +from modules.sd_hijack_hypertile import hypertile_set + + +create_binary_mask = processing_helpers.create_binary_mask +apply_overlay = processing_helpers.apply_overlay +apply_color_correction = processing_helpers.apply_color_correction +setup_color_correction = processing_helpers.setup_color_correction +images_tensor_to_samples = processing_helpers.images_tensor_to_samples +txt2img_image_conditioning = processing_helpers.txt2img_image_conditioning +img2img_image_conditioning = processing_helpers.img2img_image_conditioning +get_fixed_seed = processing_helpers.get_fixed_seed +create_random_tensors = processing_helpers.create_random_tensors +decode_first_stage = processing_helpers.decode_first_stage +old_hires_fix_first_pass_dimensions = processing_helpers.old_hires_fix_first_pass_dimensions +validate_sample = processing_helpers.validate_sample + + +def get_conds_with_caching(function, required_prompts, steps, cache): + if cache[0] is not None and (required_prompts, steps) == cache[0]: + return cache[1] + with devices.autocast(): + cache[1] = function(shared.sd_model, required_prompts, steps) + cache[0] = (required_prompts, steps) + return cache[1] + + +def process_original(p: processing.StableDiffusionProcessing): + cached_uc = [None, None] + cached_c = [None, None] + sampler_config = sd_samplers.find_sampler_config(p.sampler_name) + step_multiplier = 2 if sampler_config and sampler_config.options.get("second_order", False) else 1 + uc = get_conds_with_caching(prompt_parser.get_learned_conditioning, p.negative_prompts, p.steps * step_multiplier, cached_uc) + c = get_conds_with_caching(prompt_parser.get_multicond_learned_conditioning, p.prompts, p.steps * step_multiplier, cached_c) + with devices.without_autocast() if devices.unet_needs_upcast else devices.autocast(): + samples_ddim = p.sample(conditioning=c, unconditional_conditioning=uc, seeds=p.seeds, subseeds=p.subseeds, subseed_strength=p.subseed_strength, prompts=p.prompts) + x_samples_ddim = [processing.decode_first_stage(p.sd_model, samples_ddim[i:i+1].to(dtype=devices.dtype_vae), p.full_quality)[0].cpu() for i in range(samples_ddim.size(0))] + try: + for x in x_samples_ddim: + devices.test_for_nans(x, "vae") + except devices.NansException as e: + if not shared.opts.no_half and not shared.opts.no_half_vae and shared.cmd_opts.rollback_vae: + shared.log.warning('Tensor with all NaNs was produced in VAE') + devices.dtype_vae = torch.bfloat16 + vae_file, vae_source = sd_vae.resolve_vae(p.sd_model.sd_model_checkpoint) + sd_vae.load_vae(p.sd_model, vae_file, vae_source) + x_samples_ddim = [processing.decode_first_stage(p.sd_model, samples_ddim[i:i+1].to(dtype=devices.dtype_vae), p.full_quality)[0].cpu() for i in range(samples_ddim.size(0))] + for x in x_samples_ddim: + devices.test_for_nans(x, "vae") + else: + raise e + x_samples_ddim = torch.stack(x_samples_ddim).float() + x_samples_ddim = torch.clamp((x_samples_ddim + 1.0) / 2.0, min=0.0, max=1.0) + del samples_ddim + return x_samples_ddim + + +def sample_txt2img(p: processing.StableDiffusionProcessingTxt2Img, conditioning, unconditional_conditioning, seeds, subseeds, subseed_strength, prompts): + latent_scale_mode = shared.latent_upscale_modes.get(p.hr_upscaler, None) if p.hr_upscaler is not None else shared.latent_upscale_modes.get(shared.latent_upscale_default_mode, "None") + if latent_scale_mode is not None: + p.hr_force = False # no need to force anything + if p.enable_hr and (latent_scale_mode is None or p.hr_force): + if len([x for x in shared.sd_upscalers if x.name == p.hr_upscaler]) == 0: + shared.log.warning(f"Cannot find upscaler for hires: {p.hr_upscaler}") + p.enable_hr = False + + p.ops.append('txt2img') + hypertile_set(p) + p.sampler = sd_samplers.create_sampler(p.sampler_name, p.sd_model) + if hasattr(p.sampler, "initialize"): + p.sampler.initialize(p) + x = create_random_tensors([4, p.height // 8, p.width // 8], seeds=seeds, subseeds=subseeds, subseed_strength=p.subseed_strength, seed_resize_from_h=p.seed_resize_from_h, seed_resize_from_w=p.seed_resize_from_w, p=p) + samples = p.sampler.sample(p, x, conditioning, unconditional_conditioning, image_conditioning=txt2img_image_conditioning(p, x)) + shared.state.nextjob() + if not p.enable_hr or shared.state.interrupted or shared.state.skipped: + return samples + + p.init_hr() + if p.is_hr_pass: + prev_job = shared.state.job + target_width = p.hr_upscale_to_x + target_height = p.hr_upscale_to_y + decoded_samples = None + if shared.opts.save and shared.opts.save_images_before_highres_fix and not p.do_not_save_samples: + decoded_samples = decode_first_stage(p.sd_model, samples.to(dtype=devices.dtype_vae), p.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 = validate_sample(x_sample) + image = Image.fromarray(x_sample) + bak_extra_generation_params, bak_restore_faces = p.extra_generation_params, p.restore_faces + p.extra_generation_params = {} + p.restore_faces = False + info = processing.create_infotext(p, p.all_prompts, p.all_seeds, p.all_subseeds, [], iteration=p.iteration, position_in_batch=i) + p.extra_generation_params, p.restore_faces = bak_extra_generation_params, bak_restore_faces + images.save_image(image, p.outpath_samples, "", seeds[i], prompts[i], shared.opts.samples_format, info=info, suffix="-before-hires") + if latent_scale_mode is None or p.hr_force: # non-latent upscaling + shared.state.job = 'upscale' + if decoded_samples is None: + decoded_samples = decode_first_stage(p.sd_model, samples.to(dtype=devices.dtype_vae), p.full_quality) + 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 = validate_sample(x_sample) + image = Image.fromarray(x_sample) + image = images.resize_image(1, image, target_width, target_height, upscaler_name=p.hr_upscaler) + image = np.array(image).astype(np.float32) / 255.0 + image = np.moveaxis(image, 2, 0) + batch_images.append(image) + resized_samples = torch.from_numpy(np.array(batch_images)) + resized_samples = resized_samples.to(device=shared.device, dtype=devices.dtype_vae) + resized_samples = 2.0 * resized_samples - 1.0 + if shared.opts.sd_vae_sliced_encode and len(decoded_samples) > 1: + samples = torch.stack([p.sd_model.get_first_stage_encoding(p.sd_model.encode_first_stage(torch.unsqueeze(resized_sample, 0)))[0] for resized_sample in resized_samples]) + else: + samples = p.sd_model.get_first_stage_encoding(p.sd_model.encode_first_stage(resized_samples)) + image_conditioning = img2img_image_conditioning(p, resized_samples, samples) + else: + samples = torch.nn.functional.interpolate(samples, size=(target_height // 8, target_width // 8), mode=latent_scale_mode["mode"], antialias=latent_scale_mode["antialias"]) + if getattr(p, "inpainting_mask_weight", shared.opts.inpainting_mask_weight) < 1.0: + image_conditioning = img2img_image_conditioning(p, decode_first_stage(p.sd_model, samples.to(dtype=devices.dtype_vae), p.full_quality), samples) + else: + image_conditioning = txt2img_image_conditioning(p, samples.to(dtype=devices.dtype_vae)) + if p.hr_sampler_name == "PLMS": + p.hr_sampler_name = 'UniPC' + if p.hr_force or latent_scale_mode is not None: + shared.state.job = 'hires' + if p.denoising_strength > 0: + p.ops.append('hires') + devices.torch_gc() # GC now before running the next img2img to prevent running out of memory + p.sampler = sd_samplers.create_sampler(p.hr_sampler_name or p.sampler_name, p.sd_model) + if hasattr(p.sampler, "initialize"): + p.sampler.initialize(p) + samples = samples[:, :, p.truncate_y//2:samples.shape[2]-(p.truncate_y+1)//2, p.truncate_x//2:samples.shape[3]-(p.truncate_x+1)//2] + noise = create_random_tensors(samples.shape[1:], seeds=seeds, subseeds=subseeds, subseed_strength=subseed_strength, p=p) + sd_models.apply_token_merging(p.sd_model, p.get_token_merging_ratio(for_hr=True)) + hypertile_set(p, hr=True) + samples = p.sampler.sample_img2img(p, samples, noise, conditioning, unconditional_conditioning, steps=p.hr_second_pass_steps or p.steps, image_conditioning=image_conditioning) + sd_models.apply_token_merging(p.sd_model, p.get_token_merging_ratio()) + else: + p.ops.append('upscale') + x = None + p.is_hr_pass = False + shared.state.job = prev_job + shared.state.nextjob() + + return samples + + +def sample_img2img(p, conditioning, unconditional_conditioning, seeds, subseeds, subseed_strength, prompts): + hypertile_set(p) + x = create_random_tensors([4, p.height // 8, p.width // 8], seeds=seeds, subseeds=subseeds, subseed_strength=p.subseed_strength, seed_resize_from_h=p.seed_resize_from_h, seed_resize_from_w=p.seed_resize_from_w, p=p) + x *= p.initial_noise_multiplier + samples = p.sampler.sample_img2img(p, p.init_latent, x, conditioning, unconditional_conditioning, image_conditioning=p.image_conditioning) + if p.mask is not None: + samples = samples * p.nmask + p.init_latent * p.mask + del x + devices.torch_gc() + shared.state.nextjob() + + return samples diff --git a/modules/processing_vae.py b/modules/processing_vae.py index 8462e4a8d..6c9c63feb 100644 --- a/modules/processing_vae.py +++ b/modules/processing_vae.py @@ -2,8 +2,7 @@ import os import time import torch import torchvision.transforms.functional as TF -from modules import shared, devices, sd_vae, sd_models -import modules.taesd.sd_vae_taesd as sd_vae_taesd +from modules import shared, devices, sd_models, sd_vae, sd_vae_taesd debug = shared.log.trace if os.environ.get('SD_VAE_DEBUG', None) is not None else lambda *args, **kwargs: None diff --git a/modules/scripts.py b/modules/scripts.py index 4a163aca3..90dc25e08 100644 --- a/modules/scripts.py +++ b/modules/scripts.py @@ -5,13 +5,12 @@ import time from collections import namedtuple import gradio as gr from modules import paths, script_callbacks, extensions, script_loading, scripts_postprocessing, errors, timer -from installer import log, args as cmd_opts AlwaysVisible = object() time_component = {} time_setup = {} -debug = log.trace if os.environ.get('SD_SCRIPT_DEBUG', None) is not None else lambda *args, **kwargs: None +debug = errors.log.trace if os.environ.get('SD_SCRIPT_DEBUG', None) is not None else lambda *args, **kwargs: None class PostprocessImageArgs: @@ -92,6 +91,14 @@ class Script: """ pass # pylint: disable=unnecessary-pass + def process_images(self, p, *args): + """ + This function is called instead of main processing for AlwaysVisible scripts. + You can modify the processing object (p) here, inject hooks, etc. + args contains all values returned by components from ui() + """ + pass # pylint: disable=unnecessary-pass + def before_process_batch(self, p, *args, **kwargs): """ Called before extra networks are parsed from the prompt, so you can add @@ -219,7 +226,7 @@ def list_scripts(scriptdirname, extension): if os.path.isfile(os.path.join(base, "..", ".priority")): with open(os.path.join(base, "..", ".priority"), "r", encoding="utf-8") as f: priority = priority + str(f.read().strip()) - log.debug(f'Script priority override: ${script.name}:{priority}') + errors.log.debug(f'Script priority override: ${script.name}:{priority}') else: priority = priority + script.priority priority_list.append(ScriptFile(script.basedir, script.filename, script.path, priority)) @@ -306,7 +313,7 @@ class ScriptSummary: if total == 0: return scripts = [f'{k}:{v}' for k, v in self.time.items() if v > 0] - log.debug(f'Script: op={self.op} total={total} scripts={scripts}') + errors.log.debug(f'Script: op={self.op} total={total} scripts={scripts}') class ScriptRunner: @@ -352,7 +359,7 @@ class ScriptRunner: self.scripts.append(script) self.selectable_scripts.append(script) except Exception as e: - log.error(f'Script initialize: {path} {e}') + errors.log.error(f'Script initialize: {path} {e}') """ def create_script_ui(self, script): @@ -418,7 +425,7 @@ class ScriptRunner: for control in controls: debug(f'Script control: parent={script.parent} script="{script.name}" label="{control.label}" type={control} id={control.elem_id}') if not isinstance(control, gr.components.IOComponent): - log.error(f'Invalid script control: "{script.filename}" control={control}') + errors.log.error(f'Invalid script control: "{script.filename}" control={control}') continue control.custom_script_source = os.path.basename(script.filename) arg_info = api_models.ScriptArg(label=control.label or "") @@ -522,6 +529,19 @@ class ScriptRunner: s.record(script.title()) s.report() + def process_images(self, p, **kwargs): + s = ScriptSummary('process_images') + processed = None + for script in self.alwayson_scripts: + try: + args = p.per_script_args.get(script.title(), p.script_args[script.args_from:script.args_to]) + processed = script.process_images(p, *args, **kwargs) + except Exception as e: + errors.display(e, f'Running script process images: {script.filename}') + s.record(script.title()) + s.report() + return processed + def before_process_batch(self, p, **kwargs): s = ScriptSummary('before-process-batch') for script in self.alwayson_scripts: diff --git a/modules/sd_samplers_common.py b/modules/sd_samplers_common.py index ec64e8d58..2a46a8a28 100644 --- a/modules/sd_samplers_common.py +++ b/modules/sd_samplers_common.py @@ -3,8 +3,7 @@ from collections import namedtuple import torch import torchvision.transforms as T from PIL import Image -from modules import devices, processing, images, sd_vae_approx, sd_samplers, shared -import modules.taesd.sd_vae_taesd as sd_vae_taesd +from modules import shared, devices, processing, images, sd_vae_approx, sd_vae_taesd, sd_samplers SamplerData = namedtuple('SamplerData', ['name', 'constructor', 'aliases', 'options']) diff --git a/modules/taesd/sd_vae_taesd.py b/modules/sd_vae_taesd.py similarity index 50% rename from modules/taesd/sd_vae_taesd.py rename to modules/sd_vae_taesd.py index 47330d4b8..895dce544 100644 --- a/modules/taesd/sd_vae_taesd.py +++ b/modules/sd_vae_taesd.py @@ -6,11 +6,73 @@ https://github.com/madebyollin/taesd """ import os from PIL import Image +import torch +import torch.nn as nn from modules import devices, paths -from modules.taesd.taesd import TAESD + taesd_models = { 'sd-decoder': None, 'sd-encoder': None, 'sdxl-decoder': None, 'sdxl-encoder': None } + +def conv(n_in, n_out, **kwargs): + return nn.Conv2d(n_in, n_out, 3, padding=1, **kwargs) + +class Clamp(nn.Module): + def forward(self, x): + return torch.tanh(x / 3) * 3 + +class Block(nn.Module): + def __init__(self, n_in, n_out): + super().__init__() + self.conv = nn.Sequential(conv(n_in, n_out), nn.ReLU(), conv(n_out, n_out), nn.ReLU(), conv(n_out, n_out)) + self.skip = nn.Conv2d(n_in, n_out, 1, bias=False) if n_in != n_out else nn.Identity() + self.fuse = nn.ReLU() + def forward(self, x): + return self.fuse(self.conv(x) + self.skip(x)) + +def Encoder(): + return nn.Sequential( + conv(3, 64), Block(64, 64), + conv(64, 64, stride=2, bias=False), Block(64, 64), Block(64, 64), Block(64, 64), + conv(64, 64, stride=2, bias=False), Block(64, 64), Block(64, 64), Block(64, 64), + conv(64, 64, stride=2, bias=False), Block(64, 64), Block(64, 64), Block(64, 64), + conv(64, 4), + ) + +def Decoder(): + return nn.Sequential( + Clamp(), conv(4, 64), nn.ReLU(), + Block(64, 64), Block(64, 64), Block(64, 64), nn.Upsample(scale_factor=2), conv(64, 64, bias=False), + Block(64, 64), Block(64, 64), Block(64, 64), nn.Upsample(scale_factor=2), conv(64, 64, bias=False), + Block(64, 64), Block(64, 64), Block(64, 64), nn.Upsample(scale_factor=2), conv(64, 64, bias=False), + Block(64, 64), conv(64, 3), + ) + +class TAESD(nn.Module): # pylint: disable=abstract-method + latent_magnitude = 3 + latent_shift = 0.5 + + def __init__(self, encoder_path="taesd_encoder.pth", decoder_path="taesd_decoder.pth"): + """Initialize pretrained TAESD on the given device from the given checkpoints.""" + super().__init__() + self.encoder = Encoder() + self.decoder = Decoder() + if encoder_path is not None: + self.encoder.load_state_dict(torch.load(encoder_path, map_location="cpu")) + if decoder_path is not None: + self.decoder.load_state_dict(torch.load(decoder_path, map_location="cpu")) + + @staticmethod + def scale_latents(x): + """raw latents -> [0, 1]""" + return x.div(2 * TAESD.latent_magnitude).add(TAESD.latent_shift).clamp(0, 1) + + @staticmethod + def unscale_latents(x): + """[0, 1] -> raw latents""" + return x.sub(TAESD.latent_shift).mul(2 * TAESD.latent_magnitude) + + def download_model(model_path): model_name = os.path.basename(model_path) model_url = f'https://github.com/madebyollin/taesd/raw/main/{model_name}' @@ -58,21 +120,22 @@ def decode(latents): taesd_models[f'{model_class}-decoder'] = TAESD(decoder_path=model_path, encoder_path=None) shared.log.debug(f'VAE load: type=taesd model={model_path}') vae = taesd_models[f'{model_class}-decoder'] - vae.to(devices.device, devices.dtype_vae) - latents.to(devices.device, devices.dtype_vae) + vae.decoder.to(devices.device, devices.dtype_vae) try: - if len(latents.shape) == 3: - latents = latents.unsqueeze(0) - image = vae.decoder(latents).clamp(0, 1).detach() - image = 2.0 * image - 1.0 # typical normalized range except for preview which runs denormalization - return image[0] - elif len(latents.shape) == 4: - image = vae.decoder(latents).clamp(0, 1).detach() - image = 2.0 * image - 1.0 # typical normalized range except for preview which runs denormalization - return image - else: - shared.log.error(f'TAESD decode unsupported latent type: {latents.shape}') - return latents + with devices.inference_context(): + latents = latents.detach().clone().to(devices.device, devices.dtype_vae) + if len(latents.shape) == 3: + latents = latents.unsqueeze(0) + image = vae.decoder(latents).clamp(0, 1).detach() + image = 2.0 * image - 1.0 # typical normalized range except for preview which runs denormalization + return image[0] + elif len(latents.shape) == 4: + image = vae.decoder(latents).clamp(0, 1).detach() + image = 2.0 * image - 1.0 # typical normalized range except for preview which runs denormalization + return image + else: + shared.log.error(f'TAESD decode unsupported latent type: {latents.shape}') + return latents except Exception as e: shared.log.error(f'VAE decode taesd: {e}') return latents @@ -94,7 +157,7 @@ def encode(image): shared.log.debug(f'VAE load: type=taesd model={model_path}') taesd_models[f'{model_class}-encoder'] = TAESD(encoder_path=model_path, decoder_path=None) vae = taesd_models[f'{model_class}-encoder'] - vae.to(devices.device, devices.dtype_vae) + vae.encoder.to(devices.device, devices.dtype_vae) # image = vae.scale_latents(image) latents = vae.encoder(image) return latents.detach() diff --git a/modules/taesd/taesd.py b/modules/taesd/taesd.py deleted file mode 100644 index 4900a5ab0..000000000 --- a/modules/taesd/taesd.py +++ /dev/null @@ -1,95 +0,0 @@ -#!/usr/bin/env python3 -""" -Tiny AutoEncoder for Stable Diffusion -(DNN for encoding / decoding SD's latent space) -""" -import torch -import torch.nn as nn -from modules import devices - - -def conv(n_in, n_out, **kwargs): - return nn.Conv2d(n_in, n_out, 3, padding=1, **kwargs) - -class Clamp(nn.Module): - def forward(self, x): - return torch.tanh(x / 3) * 3 - -class Block(nn.Module): - def __init__(self, n_in, n_out): - super().__init__() - self.conv = nn.Sequential(conv(n_in, n_out), nn.ReLU(), conv(n_out, n_out), nn.ReLU(), conv(n_out, n_out)) - self.skip = nn.Conv2d(n_in, n_out, 1, bias=False) if n_in != n_out else nn.Identity() - self.fuse = nn.ReLU() - def forward(self, x): - return self.fuse(self.conv(x) + self.skip(x)) - -def Encoder(): - return nn.Sequential( - conv(3, 64), Block(64, 64), - conv(64, 64, stride=2, bias=False), Block(64, 64), Block(64, 64), Block(64, 64), - conv(64, 64, stride=2, bias=False), Block(64, 64), Block(64, 64), Block(64, 64), - conv(64, 64, stride=2, bias=False), Block(64, 64), Block(64, 64), Block(64, 64), - conv(64, 4), - ) - -def Decoder(): - return nn.Sequential( - Clamp(), conv(4, 64), nn.ReLU(), - Block(64, 64), Block(64, 64), Block(64, 64), nn.Upsample(scale_factor=2), conv(64, 64, bias=False), - Block(64, 64), Block(64, 64), Block(64, 64), nn.Upsample(scale_factor=2), conv(64, 64, bias=False), - Block(64, 64), Block(64, 64), Block(64, 64), nn.Upsample(scale_factor=2), conv(64, 64, bias=False), - Block(64, 64), conv(64, 3), - ) - -class TAESD(nn.Module): - latent_magnitude = 3 - latent_shift = 0.5 - - def __init__(self, encoder_path="taesd_encoder.pth", decoder_path="taesd_decoder.pth"): - """Initialize pretrained TAESD on the given device from the given checkpoints.""" - super().__init__() - self.encoder = Encoder() - self.decoder = Decoder() - if encoder_path is not None: - self.encoder.load_state_dict(torch.load(encoder_path, map_location="cpu")) - if decoder_path is not None: - self.decoder.load_state_dict(torch.load(decoder_path, map_location="cpu")) - - @staticmethod - def scale_latents(x): - """raw latents -> [0, 1]""" - return x.div(2 * TAESD.latent_magnitude).add(TAESD.latent_shift).clamp(0, 1) - - @staticmethod - def unscale_latents(x): - """[0, 1] -> raw latents""" - return x.sub(TAESD.latent_shift).mul(2 * TAESD.latent_magnitude) - - -@devices.inference_context() -def main(): - from PIL import Image - import sys - import torchvision.transforms.functional as TF - dev = torch.device("cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu") - print("Using device", dev) - taesd = TAESD().to(dev) - for im_path in sys.argv[1:]: - im = TF.to_tensor(Image.open(im_path).convert("RGB")).unsqueeze(0).to(dev) - - # encode image, quantize, and save to file - im_enc = taesd.scale_latents(taesd.encoder(im)).mul_(255).round_().byte() - enc_path = im_path + ".encoded.png" - TF.to_pil_image(im_enc[0]).save(enc_path) - print(f"Encoded {im_path} to {enc_path}") - - # load the saved file, dequantize, and decode - im_enc = taesd.unscale_latents(TF.to_tensor(Image.open(enc_path)).unsqueeze(0).to(dev)) - im_dec = taesd.decoder(im_enc).clamp(0, 1) - dec_path = im_path + ".decoded.png" - print(f"Decoded {enc_path} to {dec_path}") - TF.to_pil_image(im_dec[0]).save(dec_path) - -if __name__ == "__main__": - main() diff --git a/modules/txt2img.py b/modules/txt2img.py index 8a52003b8..49272cc92 100644 --- a/modules/txt2img.py +++ b/modules/txt2img.py @@ -1,6 +1,5 @@ import os -import modules.scripts -from modules import shared, processing +from modules import shared, processing, scripts from modules.generation_parameters_copypaste import create_override_settings_dict from modules.ui import plaintext_to_html @@ -82,9 +81,9 @@ def txt2img(id_task, hdr_maximize=hdr_maximize, hdr_max_center=hdr_max_center, hdr_max_boundry=hdr_max_boundry, override_settings=override_settings, ) - p.scripts = modules.scripts.scripts_txt2img + p.scripts = scripts.scripts_txt2img p.script_args = args - processed = modules.scripts.scripts_txt2img.run(p, *args) + processed = scripts.scripts_txt2img.run(p, *args) if processed is None: processed = processing.process_images(p) p.close() diff --git a/requirements.txt b/requirements.txt index 3d2de0d0d..cbc27f8a1 100644 --- a/requirements.txt +++ b/requirements.txt @@ -67,7 +67,7 @@ pandas==1.5.3 protobuf==3.20.3 pytorch_lightning==1.9.4 tokenizers==0.15.1 -transformers==4.37.1 +transformers==4.37.2 tomesd==0.1.3 urllib3==1.26.18 Pillow==10.2.0