diff --git a/modules/processing.py b/modules/processing.py index a727a93..681d8f3 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -7,13 +7,13 @@ from PIL import Image, ImageFilter class StableDiffusionProcessing: def __init__(self, init_img, model, positive, negative, vae, seed, steps, cfg, sampler_name, scheduler, denoise): - # Variables used by the upscaler script + # Variables used by the USDU script self.init_images = [init_img] self.image_mask = None self.mask_blur = 0 self.inpaint_full_res_padding = 0 - self.width = 0 - self.height = 0 + self.width = init_img.width + self.height = init_img.height # ComfyUI Sampler inputs self.model = model @@ -27,7 +27,10 @@ class StableDiffusionProcessing: self.scheduler = scheduler self.denoise = denoise - # Other required A1111 variables for the upscaler script that is currently unused in this script + # Variables used only by this script + self.init_size = init_img.width, init_img.height + + # Other required A1111 variables for the USDU script that is currently unused in this script self.extra_generation_params = {} @@ -68,8 +71,8 @@ def process_images(p: StableDiffusionProcessing) -> Processed: tile = tile.resize((p.width, p.height), Image.Resampling.LANCZOS) # Crop conditioning - positive_cropped = crop_cond(p.positive, crop_region, (p.width, p.height), init_image.size) - negative_cropped = crop_cond(p.negative, crop_region, (p.width, p.height), init_image.size) + 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)) # Encode the image vae_encoder = VAEEncode() diff --git a/utils.py b/utils.py index 7e8371a..bb496d3 100644 --- a/utils.py +++ b/utils.py @@ -75,8 +75,8 @@ def expand_crop(region, width, height): return (x1, y1, x2, y2), (p_width, p_height) -def resize_crop(region, init_size, resize_size): - # Resize a crop region so that it fits an image that was resized to the given width and height +def resize_region(region, init_size, resize_size): + # Resize a crop so that it fits an image that was resized to the given width and height x1, y1, x2, y2 = region init_width, init_height = init_size resize_width, resize_height = resize_size @@ -87,16 +87,53 @@ def resize_crop(region, init_size, resize_size): return (x1, y1, x2, y2) -def crop_cond(cond, region, p_size, image_size): +def crop_controlnet(controlnet, region, canvas_size, tile_size): + im = controlnet_hint_to_pil(controlnet.cond_hint_original) + resized_crop = resize_region(region, canvas_size, im.size) + im = im.crop(resized_crop) + im = im.resize(tile_size, Image.Resampling.NEAREST) + controlnet.cond_hint = pil_to_controlnet_hint(im).to(controlnet.device) + + +def crop_gligen(gligen, region, init_size, canvas_size): + type, model, cond = gligen + for i, c in enumerate(cond): + emb, h, w, y, x = c + if type == "position": + x1 = x * 8 + y1 = y * 8 + x2 = x1 + w * 8 + y2 = y1 + h * 8 + x1, y1, x2, y2 = resize_region((x1, y1, x2, y2), init_size, canvas_size) + + # Calculate the intersection of the gligen box and the region + x1_, y1_, x2_, y2_ = region + x1 = max(x1, x1_) + y1 = max(y1, y1_) + x2 = min(x2, x2_) + y2 = min(y2, y2_) + + # Set the new position params + h = (y2 - y1) // 8 + w = (x2 - x1) // 8 + x = x1 // 8 + y = y1 // 8 + cond[i] = (emb, h, w, y, x) + + else: + from warnings import warn + warn(f"Cropping of gligen method of type \"{type}\" is not implemented yet") + + +def crop_cond(cond, region, init_size, canvas_size, tile_size): cropped = [] for emb, x in cond: n = [emb, x.copy()] if "control" in n[1]: cnet = n[1]["control"] - im = controlnet_hint_to_pil(cnet.cond_hint_original) - resized_crop = resize_crop(region, image_size, im.size) - im = im.crop(resized_crop) - im = im.resize(p_size, Image.Resampling.NEAREST) - cnet.cond_hint = pil_to_controlnet_hint(im).to(cnet.device) + crop_controlnet(cnet, region, canvas_size, tile_size) + if "gligen" in n[1]: + gligen = n[1]["gligen"] + crop_gligen(gligen, region, init_size, canvas_size) cropped.append(n) return cropped