diff --git a/modules/processing.py b/modules/processing.py index 5ffd0f4..3ddc4f0 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -1,7 +1,7 @@ from PIL import Image, ImageFilter import torch from nodes import common_ksampler, VAEEncode, VAEDecode -from utils import pil_to_tensor, tensor_to_pil, get_crop_region, expand_crop, crop_cond +from utils import pil_to_tensor, tensor_to_pil, get_crop_region, expand_crop, crop_cond, pad_image from modules import shared if (not hasattr(Image, 'Resampling')): # For older versions of Pillow @@ -10,7 +10,7 @@ if (not hasattr(Image, 'Resampling')): # For older versions of Pillow class StableDiffusionProcessing: - def __init__(self, init_img, model, positive, negative, vae, seed, steps, cfg, sampler_name, scheduler, denoise, upscale_by=1): + def __init__(self, init_img, model, positive, negative, vae, seed, steps, cfg, sampler_name, scheduler, denoise, upscale_by, force_uniform_tile_size): # Variables used by the USDU script self.init_images = [init_img] self.image_mask = None @@ -34,6 +34,7 @@ class StableDiffusionProcessing: # Variables used only by this script self.init_size = init_img.width, init_img.height self.upscale_by = upscale_by + self.force_uniform_tile_size = force_uniform_tile_size == "enable" # Other required A1111 variables for the USDU script that is currently unused in this script self.extra_generation_params = {} @@ -63,7 +64,7 @@ def process_images(p: StableDiffusionProcessing) -> Processed: # Locate the white region of the mask outlining the tile and add padding crop_region = get_crop_region(image_mask, p.inpaint_full_res_padding) - crop_region, (p.width, p.height) = expand_crop(crop_region, image_mask.width, image_mask.height) + crop_region, tile_size = expand_crop(crop_region, image_mask.width, image_mask.height) # Blur the mask if p.mask_blur > 0: @@ -72,13 +73,22 @@ def process_images(p: StableDiffusionProcessing) -> Processed: # Crop the images to get the tiles that will be used for generation tiles = [img.crop(crop_region) for img in shared.batch] initial_tile_size = tiles[0].size + w_pad = 0 + h_pad = 0 for i in range(len(tiles)): - if tiles[i].size != (p.width, p.height): - tiles[i] = tiles[i].resize((p.width, p.height), Image.Resampling.LANCZOS) + if tiles[i].size != tile_size: + tiles[i] = tiles[i].resize(tile_size, Image.Resampling.LANCZOS) + + if p.force_uniform_tile_size: + # Pad the tile to center it in an image of size (p.width, p.height) + w_pad = (p.width - tile_size[0]) // 2 + h_pad = (p.height - tile_size[1]) // 2 + tiles[i] = pad_image(tiles[i], left_pad=w_pad, right_pad=w_pad, + top_pad=h_pad, bottom_pad=h_pad, fill=True, blur=True) # Crop conditioning - positive_cropped = crop_cond(p.positive, crop_region, p.init_size, init_image.size, (p.width, p.height)) - negative_cropped = crop_cond(p.negative, crop_region, p.init_size, init_image.size, (p.width, p.height)) + positive_cropped = crop_cond(p.positive, crop_region, p.init_size, init_image.size, tile_size, w_pad, h_pad) + negative_cropped = crop_cond(p.negative, crop_region, p.init_size, init_image.size, tile_size, w_pad, h_pad) # Encode the image vae_encoder = VAEEncode() @@ -99,6 +109,10 @@ def process_images(p: StableDiffusionProcessing) -> Processed: for i, tile_sampled in enumerate(tiles_sampled): init_image = shared.batch[i] + if p.force_uniform_tile_size: + # Crop out the padding from the samples + tile_sampled = tile_sampled.crop((w_pad, h_pad, tile_sampled.width - w_pad, tile_sampled.height - h_pad)) + # Resize back to the original size if tile_sampled.size != initial_tile_size: tile_sampled = tile_sampled.resize(initial_tile_size, Image.Resampling.LANCZOS) diff --git a/nodes.py b/nodes.py index 695315f..376139a 100644 --- a/nodes.py +++ b/nodes.py @@ -52,6 +52,8 @@ def USDU_base_inputs(): ("seam_fix_width", ("INT", {"default": 64, "min": 0, "max": MAX_RESOLUTION, "step": 8})), ("seam_fix_mask_blur", ("INT", {"default": 8, "min": 0, "max": 64, "step": 1})), ("seam_fix_padding", ("INT", {"default": 16, "min": 0, "max": MAX_RESOLUTION, "step": 8})), + # Misc + ("force_uniform_tile_size", (["disable", "enable"], )) ] @@ -95,7 +97,7 @@ class UltimateSDUpscale: steps, cfg, sampler_name, scheduler, denoise, upscale_model, mode_type, tile_width, tile_height, mask_blur, tile_padding, seam_fix_mode, seam_fix_denoise, seam_fix_mask_blur, - seam_fix_width, seam_fix_padding): + seam_fix_width, seam_fix_padding, force_uniform_tile_size): # # Set up A1111 patches # @@ -112,7 +114,7 @@ class UltimateSDUpscale: # Processing sdprocessing = StableDiffusionProcessing( tensor_to_pil(image), model, positive, negative, vae, - seed, steps, cfg, sampler_name, scheduler, denoise, upscale_by + seed, steps, cfg, sampler_name, scheduler, denoise, upscale_by, force_uniform_tile_size ) # @@ -150,14 +152,14 @@ class UltimateSDUpscaleNoUpscale: steps, cfg, sampler_name, scheduler, denoise, mode_type, tile_width, tile_height, mask_blur, tile_padding, seam_fix_mode, seam_fix_denoise, seam_fix_mask_blur, - seam_fix_width, seam_fix_padding): + seam_fix_width, seam_fix_padding, force_uniform_tile_size): shared.sd_upscalers[0] = UpscalerData() shared.actual_upscaler = None shared.batch = [tensor_to_pil(upscaled_image, i) for i in range(len(upscaled_image))] sdprocessing = StableDiffusionProcessing( tensor_to_pil(upscaled_image), model, positive, negative, vae, - seed, steps, cfg, sampler_name, scheduler, denoise + seed, steps, cfg, sampler_name, scheduler, denoise, 1, force_uniform_tile_size ) script = usdu.Script() diff --git a/utils.py b/utils.py index c8cc2a1..d48ca44 100644 --- a/utils.py +++ b/utils.py @@ -1,11 +1,15 @@ import numpy as np -from PIL import Image +from PIL import Image, ImageFilter import torch +import torch.nn.functional as F +from torchvision.transforms import GaussianBlur import math if (not hasattr(Image, 'Resampling')): # For older versions of Pillow Image.Resampling = Image +BLUR_KERNEL_SIZE = 15 + def tensor_to_pil(img_tensor, batch_index=0): # Takes an image in a batch in the form of a tensor of shape [batch_size, channels, height, width] @@ -74,6 +78,67 @@ def fix_crop_region(region, image_size): return x1, y1, x2, y2 +def pad_image(image, left_pad, right_pad, top_pad, bottom_pad, fill=False, blur=False): + ''' + Pads an image with the given number of pixels on each side and fills the padding with data from the edges. + :param image: A PIL image + :param left_pad: The number of pixels to pad on the left side + :param right_pad: The number of pixels to pad on the right side + :param top_pad: The number of pixels to pad on the top side + :param bottom_pad: The number of pixels to pad on the bottom side + :param blur: Whether to blur the padded edges + :return: A PIL image with size (image.width + left_pad + right_pad, image.height + top_pad + bottom_pad) + ''' + left_edge = image.crop((0, 1, 1, image.height - 1)) + right_edge = image.crop((image.width - 1, 1, image.width, image.height - 1)) + top_edge = image.crop((1, 0, image.width - 1, 1)) + bottom_edge = image.crop((1, image.height - 1, image.width - 1, image.height)) + new_width = image.width + left_pad + right_pad + new_height = image.height + top_pad + bottom_pad + padded_image = Image.new('RGB', (new_width, new_height), 0) + padded_image.paste(image, (left_pad, top_pad)) + if fill: + for i in range(left_pad): + edge = left_edge.resize( + (1, new_height - i * (top_pad + bottom_pad) // left_pad), resample=Image.Resampling.NEAREST) + padded_image.paste(edge, (i, i * top_pad // left_pad)) + for i in range(right_pad): + edge = right_edge.resize( + (1, new_height - i * (top_pad + bottom_pad) // right_pad), resample=Image.Resampling.NEAREST) + padded_image.paste(edge, (new_width - 1 - i, i * top_pad // right_pad)) + for i in range(top_pad): + edge = top_edge.resize( + (new_width - i * (left_pad + right_pad) // top_pad, 1), resample=Image.Resampling.NEAREST) + padded_image.paste(edge, (i * left_pad // top_pad, i)) + for i in range(bottom_pad): + edge = bottom_edge.resize( + (new_width - i * (left_pad + right_pad) // bottom_pad, 1), resample=Image.Resampling.NEAREST) + padded_image.paste(edge, (i * left_pad // bottom_pad, new_height - 1 - i)) + if blur and not (left_pad == right_pad == top_pad == bottom_pad == 0): + padded_image = padded_image.filter(ImageFilter.GaussianBlur(BLUR_KERNEL_SIZE)) + padded_image.paste(image, (left_pad, top_pad)) + return padded_image + + +def pad_tensor(tensor, left_pad, right_pad, top_pad, bottom_pad, fill=False, blur=False): + ''' + Pads an image tensor with the given number of pixels on each side and fills the padding with data from the edges. + :param tensor: A tensor of shape [B, H, W, C] + :param left_pad: The number of pixels to pad on the left side + :param right_pad: The number of pixels to pad on the right side + :param top_pad: The number of pixels to pad on the top side + :param bottom_pad: The number of pixels to pad on the bottom side + :param blur: Whether to blur the padded edges + :return: A tensor of shape [B, H + top_pad + bottom_pad, W + left_pad + right_pad, C] + ''' + tensors = [] + for i in range(tensor.shape[0]): + image = tensor_to_pil(tensor, i) + image = pad_image(image, left_pad, right_pad, top_pad, bottom_pad, fill, blur) + tensors.append(pil_to_tensor(image)) + return torch.cat(tensors, dim=0) + + def expand_crop(region, width, height): # Expand the crop region to a multiple of 8 for encoding x1, y1, x2, y2 = region @@ -102,7 +167,6 @@ def expand_crop(region, width, height): height_diff = p_height - (y2 - y1) y2 = min(y2 + height_diff, height) - # Width and height should be the same as p_width and p_height return (x1, y1, x2, y2), (p_width, p_height) @@ -118,7 +182,7 @@ def resize_region(region, init_size, resize_size): return (x1, y1, x2, y2) -def crop_controlnet(cond_dict, region, init_size, canvas_size, tile_size): +def crop_controlnet(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad): if "control" not in cond_dict: return c = cond_dict["control"] @@ -130,6 +194,7 @@ def crop_controlnet(cond_dict, region, init_size, canvas_size, tile_size): resized_crop = resize_region(region, canvas_size, hint.shape[:-3:-1]) hint = crop_tensor(hint.movedim(1, -1), resized_crop).movedim(-1, 1) hint = resize_tensor(hint, tile_size[::-1]) + hint = pad_tensor(hint.movedim(1, -1), w_pad, w_pad, h_pad, h_pad, blur=True).movedim(-1, 1) controlnet.cond_hint_original = hint c = c.previous_controlnet @@ -157,7 +222,7 @@ def region_intersection(region1, region2): return (x1, y1, x2, y2) -def crop_gligen(cond_dict, region, init_size, canvas_size, tile_size): +def crop_gligen(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad): if "gligen" not in cond_dict: return type, model, cond = cond_dict["gligen"] @@ -187,18 +252,26 @@ def crop_gligen(cond_dict, region, init_size, canvas_size, tile_size): x2 -= region[0] y2 -= region[1] + # Add the padding + x1 += w_pad + y1 += h_pad + x2 += w_pad + y2 += h_pad + # Set the new position params h = (y2 - y1) // 8 w = (x2 - x1) // 8 x = x1 // 8 y = y1 // 8 cropped.append((emb, h, w, y, x)) + cond_dict["gligen"] = (type, model, cropped) -def crop_area(cond_dict, region, init_size, canvas_size, tile_size): +def crop_area(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad): if "area" not in cond_dict: return + # Resize the area conditioning to the canvas size and confine it to the tile region h, w, y, x = cond_dict["area"] w, h, x, y = 8 * w, 8 * h, 8 * x, 8 * y @@ -209,52 +282,73 @@ def crop_area(cond_dict, region, init_size, canvas_size, tile_size): del cond_dict["strength"] return x1, y1, x2, y2 = intersection + # Offset origin to the top left of the tile x1 -= region[0] y1 -= region[1] x2 -= region[0] y2 -= region[1] + + # Add the padding + x1 += w_pad + y1 += h_pad + x2 += w_pad + y2 += h_pad + # Set the params for tile w, h = (x2 - x1) // 8, (y2 - y1) // 8 x, y = x1 // 8, y1 // 8 + cond_dict["area"] = (h, w, y, x) -def crop_mask(cond_dict, region, init_size, canvas_size, tile_size): +def crop_mask(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad): if "mask" not in cond_dict: return - mask = cond_dict["mask"] # (1, H, W) - # Convert to PIL image - mask = tensor_to_pil(mask) # W x H - # Resize the mask to the canvas size - mask = mask.resize(canvas_size, Image.Resampling.BICUBIC) - # Crop the mask to the region - mask = mask.crop(region) - # Resize the mask to the tile size - if tile_size != mask.size: - mask = mask.resize(tile_size, Image.Resampling.BICUBIC) - # Remove mask if it is all white - mask_bbox = mask.getbbox() - if mask_bbox is not None: - # Check if mask is completely contains the tile - if region_intersection(region, mask_bbox) == region: - del cond_dict["mask"] - del cond_dict["mask_strength"] - return - # Convert back to tensor - mask = pil_to_tensor(mask) # (1, H, W, 1) - mask = mask.squeeze(-1) # (1, H, W) - cond_dict["mask"] = mask + mask_tensor = cond_dict["mask"] # (B, H, W) + masks = [] + for i in range(mask_tensor.shape[0]): + # Convert to PIL image + mask = tensor_to_pil(mask_tensor, i) # W x H + + # Resize the mask to the canvas size + mask = mask.resize(canvas_size, Image.Resampling.BICUBIC) + + # Crop the mask to the region + mask = mask.crop(region) + + # Add padding + mask = pad_image(mask, w_pad, w_pad, h_pad, h_pad, fill=False) + + # Resize the mask to the tile size + if tile_size != mask.size: + mask = mask.resize(tile_size, Image.Resampling.BICUBIC) + + # # Remove mask if it is all white + # mask_bbox = mask.getbbox() + # if mask_bbox is not None: + # # Check if mask is completely contains the tile + # if region_intersection(region, mask_bbox) == region: + # del cond_dict["mask"] + # del cond_dict["mask_strength"] + # return + + # Convert back to tensor + mask = pil_to_tensor(mask) # (1, H, W, 1) + mask = mask.squeeze(-1) # (1, H, W) + masks.append(mask) + + cond_dict["mask"] = torch.cat(masks, dim=0) # (B, H, W) -def crop_cond(cond, region, init_size, canvas_size, tile_size): +def crop_cond(cond, region, init_size, canvas_size, tile_size, w_pad, h_pad): cropped = [] for emb, x in cond: cond_dict = x.copy() n = [emb, cond_dict] - crop_controlnet(cond_dict, region, init_size, canvas_size, tile_size) - crop_gligen(cond_dict, region, init_size, canvas_size, tile_size) - crop_area(cond_dict, region, init_size, canvas_size, tile_size) - crop_mask(cond_dict, region, init_size, canvas_size, tile_size) + crop_controlnet(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad) + crop_gligen(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad) + crop_area(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad) + crop_mask(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad) cropped.append(n) return cropped