Added option for uniformly sized tiles
This commit is contained in:
+21
-7
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user