Added option for uniformly sized tiles

This commit is contained in:
ssit
2023-07-11 23:18:12 -04:00
parent c89eccf0de
commit 86ffbc1a0c
3 changed files with 154 additions and 44 deletions
+21 -7
View File
@@ -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)
+6 -4
View File
@@ -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()
+127 -33
View File
@@ -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