diff --git a/.gitignore b/.gitignore index 3bbe7b6..9699f99 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,71 @@ +# Byte-compiled / optimized / DLL files __pycache__/ -*.pyc -*.pyo +*.py[cod] +*$py.class + +# C extensions +*.so + +# Distribution / packaging +.Python +build/ +develop-eggs/ +dist/ +downloads/ +eggs/ +.eggs/ +lib/ +lib64/ +parts/ +sdist/ +var/ +wheels/ +share/python-wheels/ +*.egg-info/ +.installed.cfg +*.egg +MANIFEST + +# Virtual environments +.env +.venv +env/ +venv/ +ENV/ +env.bak/ +venv.bak/ + +# Unit test / coverage / cache reports +htmlcov/ +.tox/ +.nox/ +.coverage +.coverage.* +.cache +nosetests.xml +coverage.xml +*.cover +*.py,cover +.hypothesis/ +.pytest_cache/ +.mypy_cache/ +.ruff_cache/ + +# IDEs, editors, OS files +.vscode/ +.idea/ +*.swp +*.swo +*~ +.DS_Store +Thumbs.db +.directory + +# Temporary test outputs / logs +*.tmp +*.log +scratch/ + +# ComfyUI runtime directories (when running within ComfyUI) +output/ +temp/ diff --git a/inpaint_cropandstitch.py b/inpaint_cropandstitch.py index 6e23a84..df5fd24 100644 --- a/inpaint_cropandstitch.py +++ b/inpaint_cropandstitch.py @@ -1,14 +1,32 @@ -import comfy.utils -import comfy.model_management import math -import nodes +from abc import ABC, abstractmethod import numpy as np import torch import torch.nn.functional as TF import torchvision.transforms.functional as F from PIL import Image from scipy.ndimage import gaussian_filter, grey_dilation, binary_closing, binary_fill_holes -from abc import ABC, abstractmethod + +try: + import comfy.utils + import comfy.model_management +except ImportError: + class _MockComfyModelManagement: + @staticmethod + def get_torch_device(): + return torch.device("cuda" if torch.cuda.is_available() else "cpu") + class _MockComfy: + model_management = _MockComfyModelManagement + utils = None + comfy = _MockComfy() + +try: + import nodes +except ImportError: + class _MockNodes: + MAX_RESOLUTION = 16384 + nodes = _MockNodes() + class ProcessorLogic(ABC): @abstractmethod @@ -90,19 +108,143 @@ class ProcessorLogic(ABC): def crop_magic_im(self, image, mask, x, y, w, h, target_w, target_h, padding, downscale_algorithm, upscale_algorithm, resize_output=True): pass - @abstractmethod def stitch_magic_im(self, canvas_image, inpainted_image, mask, ctc_x, ctc_y, ctc_w, ctc_h, cto_x, cto_y, cto_w, cto_h, downscale_algorithm, upscale_algorithm): - pass + canvas_image = canvas_image.clone() + inpainted_image = inpainted_image.clone() + mask = mask.clone() + + # Ensure inpainted_image is 4D [B, H, W, C] + if inpainted_image.ndim == 3: + inpainted_image = inpainted_image.unsqueeze(0) + if mask.ndim == 2: + mask = mask.unsqueeze(0) + + ctc_w = max(1, int(ctc_w)) + ctc_h = max(1, int(ctc_h)) + ctc_x = int(ctc_x) + ctc_y = int(ctc_y) + + # Resize inpainted image and mask to match the context size + B, h, w, _ = inpainted_image.shape + if ctc_w > w or ctc_h > h: # Upscaling + resized_image = self.rescale_i(inpainted_image, ctc_w, ctc_h, upscale_algorithm) + resized_mask = self.rescale_m(mask, ctc_w, ctc_h, upscale_algorithm) + else: # Downscaling + resized_image = self.rescale_i(inpainted_image, ctc_w, ctc_h, downscale_algorithm) + resized_mask = self.rescale_m(mask, ctc_w, ctc_h, downscale_algorithm) + + # Clamp mask to [0, 1] and expand to match image channels + resized_mask = resized_mask.clamp(0, 1).unsqueeze(-1) # shape: [B, H, W, 1] + + # Ensure canvas_crop is within canvas bounds + canvas_h, canvas_w = canvas_image.shape[1], canvas_image.shape[2] + crop_y1 = max(0, min(canvas_h, ctc_y)) + crop_y2 = max(0, min(canvas_h, ctc_y + ctc_h)) + crop_x1 = max(0, min(canvas_w, ctc_x)) + crop_x2 = max(0, min(canvas_w, ctc_x + ctc_w)) + + if crop_y2 <= crop_y1 or crop_x2 <= crop_x1: + # Nothing to stitch / out of bounds + output_image = canvas_image[:, cto_y:cto_y + cto_h, cto_x:cto_x + cto_w] + return output_image + + # Extract the canvas region we're about to overwrite + canvas_crop = canvas_image[:, crop_y1:crop_y2, crop_x1:crop_x2] + + # If canvas_crop size does not match resized dimensions (e.g. edge clipping), + # slice the corresponding part of resized_image and resized_mask + actual_h = crop_y2 - crop_y1 + actual_w = crop_x2 - crop_x1 + if resized_image.shape[1] != actual_h or resized_image.shape[2] != actual_w: + offset_y = crop_y1 - ctc_y + offset_x = crop_x1 - ctc_x + resized_image = resized_image[:, offset_y:offset_y + actual_h, offset_x:offset_x + actual_w] + resized_mask = resized_mask[:, offset_y:offset_y + actual_h, offset_x:offset_x + actual_w] + + # --- Channel reconciliation (defensive handling of RGB vs RGBA vs Grayscale) --- + c_canvas = canvas_crop.shape[-1] + c_inpaint = resized_image.shape[-1] + + if c_canvas == 4 and c_inpaint == 3: + # Canvas has alpha (RGBA), inpaint is RGB. Preserve the canvas's alpha channel! + alpha = canvas_crop[:, :, :, 3:4].to(device=resized_image.device, dtype=resized_image.dtype) + resized_image = torch.cat([resized_image, alpha], dim=-1) + elif c_canvas == 3 and c_inpaint == 4: + # Canvas is RGB, inpaint has alpha (RGBA). Drop alpha channel to blend into RGB canvas. + resized_image = resized_image[:, :, :, :3] + elif c_canvas == 1 and c_inpaint == 3: + # Grayscale canvas, RGB inpaint: convert inpaint to grayscale + resized_image = (0.2989 * resized_image[:, :, :, 0:1] + 0.5870 * resized_image[:, :, :, 1:2] + 0.1140 * resized_image[:, :, :, 2:3]) + elif c_canvas == 3 and c_inpaint == 1: + # RGB canvas, Grayscale inpaint: repeat to 3 channels + resized_image = resized_image.repeat(1, 1, 1, 3) + elif c_canvas == 4 and c_inpaint == 1: + # RGBA canvas, Grayscale inpaint: repeat to 3 channels and preserve canvas alpha + alpha = canvas_crop[:, :, :, 3:4].to(device=resized_image.device, dtype=resized_image.dtype) + resized_image = torch.cat([resized_image.repeat(1, 1, 1, 3), alpha], dim=-1) + elif c_canvas != c_inpaint: + # Fallback for unexpected channel counts + if c_inpaint > c_canvas: + resized_image = resized_image[:, :, :, :c_canvas] + else: + repeats = (c_canvas + c_inpaint - 1) // c_inpaint + resized_image = resized_image.repeat(1, 1, 1, repeats)[:, :, :, :c_canvas] + + # Ensure device and dtype match + resized_image = resized_image.to(device=canvas_crop.device, dtype=canvas_crop.dtype) + resized_mask = resized_mask.to(device=canvas_crop.device, dtype=canvas_crop.dtype) + + # Blend: new = mask * inpainted + (1 - mask) * canvas + blended = resized_mask * resized_image + (1.0 - resized_mask) * canvas_crop + + # Paste the blended region back onto the canvas + canvas_image[:, crop_y1:crop_y2, crop_x1:crop_x2] = blended.to(dtype=canvas_image.dtype, device=canvas_image.device) + + # Final crop to get back the original image area + out_y1 = max(0, min(canvas_h, int(cto_y))) + out_y2 = max(0, min(canvas_h, int(cto_y + cto_h))) + out_x1 = max(0, min(canvas_w, int(cto_x))) + out_x2 = max(0, min(canvas_w, int(cto_x + cto_w))) + output_image = canvas_image[:, out_y1:out_y2, out_x1:out_x2] + + return output_image + + +def _get_pil_resampling(algorithm: str): + if not isinstance(algorithm, str): + return Image.Resampling.BILINEAR + algo = algorithm.lower().strip() + mapping = { + "nearest": Image.Resampling.NEAREST, + "nearest-exact": Image.Resampling.NEAREST, + "bilinear": Image.Resampling.BILINEAR, + "bicubic": Image.Resampling.BICUBIC, + "lanczos": Image.Resampling.LANCZOS, + "box": Image.Resampling.BOX, + "area": Image.Resampling.BOX, + "hamming": Image.Resampling.HAMMING, + } + if algo in mapping: + return mapping[algo] + try: + return getattr(Image.Resampling, algorithm.upper()) + except (AttributeError, ValueError): + try: + return getattr(Image, algorithm.upper()) + except (AttributeError, ValueError): + return Image.Resampling.BILINEAR class CPUProcessorLogic(ProcessorLogic): def rescale_i(self, samples, width, height, algorithm: str): # samples shape: [B, H, W, C] + width = max(1, int(width)) + height = max(1, int(height)) samples = samples.movedim(-1, 1) # [B, C, H, W] - algorithm_enum = getattr(Image, algorithm.upper()) # i.e. Image.BICUBIC + algorithm_enum = _get_pil_resampling(algorithm) results = [] for i in range(samples.shape[0]): - samples_pil: Image.Image = F.to_pil_image(samples[i].cpu()).resize((width, height), algorithm_enum) + samples_pil: Image.Image = F.to_pil_image(samples[i].float().cpu()).resize((width, height), algorithm_enum) results.append(F.to_tensor(samples_pil)) samples = torch.stack(results, dim=0) samples = samples.movedim(1, -1) @@ -110,26 +252,37 @@ class CPUProcessorLogic(ProcessorLogic): def rescale_m(self, samples, width, height, algorithm: str): # samples shape: [B, H, W] - algorithm_enum = getattr(Image, algorithm.upper()) # i.e. Image.BICUBIC + width = max(1, int(width)) + height = max(1, int(height)) + if samples.ndim == 2: + samples = samples.unsqueeze(0) + algorithm_enum = _get_pil_resampling(algorithm) results = [] for i in range(samples.shape[0]): - samples_pil: Image.Image = F.to_pil_image(samples[i].cpu()).resize((width, height), algorithm_enum) + samples_pil: Image.Image = F.to_pil_image(samples[i].float().cpu()).resize((width, height), algorithm_enum) results.append(F.to_tensor(samples_pil).squeeze(0)) samples = torch.stack(results, dim=0) return samples def fillholes_iterative_hipass_fill_m(self, samples): + is_2d = False + if samples.ndim == 2: + is_2d = True + samples = samples.unsqueeze(0) thresholds = [1, 0.99, 0.97, 0.95, 0.93, 0.9, 0.8, 0.7, 0.6, 0.5, 0.4, 0.3, 0.2, 0.1] results = [] for i in range(samples.shape[0]): - mask_np = samples[i].cpu().numpy() + mask_np = samples[i].float().cpu().numpy() for threshold in thresholds: thresholded_mask = mask_np >= threshold closed_mask = binary_closing(thresholded_mask, structure=np.ones((3, 3)), border_value=1) filled_mask = binary_fill_holes(closed_mask) mask_np = np.maximum(mask_np, np.where(filled_mask != 0, threshold, 0)) results.append(torch.from_numpy(mask_np.astype(np.float32))) - return torch.stack(results, dim=0) + res = torch.stack(results, dim=0).to(samples.device) + if is_2d: + res = res.squeeze(0) + return res def hipassfilter_m(self, samples, threshold): filtered_mask = samples.clone() @@ -137,38 +290,66 @@ class CPUProcessorLogic(ProcessorLogic): return filtered_mask def expand_m(self, mask, pixels): + if pixels <= 0: + return mask + is_2d = False + if mask.ndim == 2: + is_2d = True + mask = mask.unsqueeze(0) sigma = pixels / 4 kernel_size = math.ceil(sigma * 1.5 + 1) kernel = np.ones((kernel_size, kernel_size), dtype=np.uint8) results = [] for i in range(mask.shape[0]): - mask_np = mask[i].cpu().numpy() + mask_np = mask[i].float().cpu().numpy() dilated_mask = grey_dilation(mask_np, footprint=kernel, mode='reflect') results.append(torch.from_numpy(dilated_mask.astype(np.float32)).clamp(0.0, 1.0)) - return torch.stack(results, dim=0) + res = torch.stack(results, dim=0).to(mask.device) + if is_2d: + res = res.squeeze(0) + return res def invert_m(self, samples): + if samples.dtype == torch.bool: + return (~samples).float() inverted_mask = samples.clone() + if not inverted_mask.is_floating_point(): + inverted_mask = inverted_mask.float() inverted_mask = 1.0 - inverted_mask return inverted_mask def blur_m(self, samples, pixels): if pixels <= 0: return samples + is_2d = False + if samples.ndim == 2: + is_2d = True + samples = samples.unsqueeze(0) sigma = pixels / 4 results = [] for i in range(samples.shape[0]): - mask_np = samples[i].cpu().numpy() + mask_np = samples[i].float().cpu().numpy() blurred_mask = gaussian_filter(mask_np, sigma=sigma, mode='reflect') results.append(torch.from_numpy(blurred_mask).float().clamp(0.0, 1.0)) - return torch.stack(results, dim=0) + res = torch.stack(results, dim=0).to(samples.device) + if is_2d: + res = res.squeeze(0) + return res def debug_context_location_in_image(self, image, x, y, w, h): debug_image = image.clone() - debug_image[:, y:y+h, x:x+w, :] = 1.0 - debug_image[:, y:y+h, x:x+w, :] + B, img_h, img_w, C = image.shape + x1 = max(0, min(img_w, int(x))) + y1 = max(0, min(img_h, int(y))) + x2 = max(0, min(img_w, int(x + w))) + y2 = max(0, min(img_h, int(y + h))) + if x2 > x1 and y2 > y1: + debug_image[:, y1:y2, x1:x2, :] = 1.0 - debug_image[:, y1:y2, x1:x2, :] return debug_image def pad_to_multiple(self, value, multiple): + if multiple <= 0: + return int(value) return int(math.ceil(value / multiple) * multiple) def preresize_imm(self, image, mask, optional_context_mask, downscale_algorithm, upscale_algorithm, preresize_mode, preresize_min_width, preresize_min_height, preresize_max_width, preresize_max_height): @@ -260,9 +441,12 @@ class CPUProcessorLogic(ProcessorLogic): assert new_H >= 0, f"Error: Trying to crop too much, height ({new_H}) must be >= 0" assert new_W >= 0, f"Error: Trying to crop too much, width ({new_W}) must be >= 0" - expanded_image = torch.zeros(B, new_H, new_W, C, device=image.device) - expanded_mask = torch.ones(B, new_H, new_W, device=mask.device) - expanded_optional_context_mask = torch.zeros(B, new_H, new_W, device=optional_context_mask.device) + if optional_context_mask is None: + optional_context_mask = torch.zeros(B, H, W, device=image.device, dtype=mask.dtype) + + expanded_image = torch.zeros(B, new_H, new_W, C, device=image.device, dtype=image.dtype) + expanded_mask = torch.ones(B, new_H, new_W, device=mask.device, dtype=mask.dtype) + expanded_optional_context_mask = torch.zeros(B, new_H, new_W, device=optional_context_mask.device, dtype=optional_context_mask.dtype) up_padding = int(H * (extend_up_factor - 1.0)) down_padding = new_H - H - up_padding @@ -387,8 +571,13 @@ class CPUProcessorLogic(ProcessorLogic): mask = mask.clone() # Check for invalid inputs - if target_w <= 0 or target_h <= 0 or w == 0 or h == 0: - return image, 0, 0, image.shape[2], image.shape[1], image, mask, 0, 0, image.shape[2], image.shape[1] + if target_w <= 0 or target_h <= 0 or w <= 0 or h <= 0: + crop_im = image + crop_m = mask + if resize_output and target_w > 0 and target_h > 0: + crop_im = self.rescale_i(crop_im, target_w, target_h, downscale_algorithm) + crop_m = self.rescale_m(crop_m, target_w, target_h, downscale_algorithm) + return image, 0, 0, image.shape[2], image.shape[1], crop_im, crop_m, 0, 0, image.shape[2], image.shape[1] # Step 1: Pad target dimensions to be multiples of padding if padding != 0: @@ -507,8 +696,8 @@ class CPUProcessorLogic(ProcessorLogic): expanded_image_h += down_padding # Step 5: Create the new image and mask - expanded_image = torch.zeros((image.shape[0], expanded_image_h, expanded_image_w, image.shape[3]), device=image.device) - expanded_mask = torch.ones((mask.shape[0], expanded_image_h, expanded_image_w), device=mask.device) + expanded_image = torch.zeros((image.shape[0], expanded_image_h, expanded_image_w, image.shape[3]), device=image.device, dtype=image.dtype) + expanded_mask = torch.ones((mask.shape[0], expanded_image_h, expanded_image_w), device=mask.device, dtype=mask.dtype) # Reorder the tensors to match the required dimension format for padding image = image.permute(0, 3, 1, 2) # [B, H, W, C] -> [B, C, H, W] @@ -566,47 +755,20 @@ class CPUProcessorLogic(ProcessorLogic): return canvas_image, cto_x, cto_y, cto_w, cto_h, cropped_image, cropped_mask, ctc_x, ctc_y, ctc_w, ctc_h - def stitch_magic_im(self, canvas_image, inpainted_image, mask, ctc_x, ctc_y, ctc_w, ctc_h, cto_x, cto_y, cto_w, cto_h, downscale_algorithm, upscale_algorithm): - canvas_image = canvas_image.clone() - inpainted_image = inpainted_image.clone() - mask = mask.clone() - - # Resize inpainted image and mask to match the context size - B, h, w, _ = inpainted_image.shape - if ctc_w > w or ctc_h > h: # Upscaling - resized_image = self.rescale_i(inpainted_image, ctc_w, ctc_h, upscale_algorithm) - resized_mask = self.rescale_m(mask, ctc_w, ctc_h, upscale_algorithm) - else: # Downscaling - resized_image = self.rescale_i(inpainted_image, ctc_w, ctc_h, downscale_algorithm) - resized_mask = self.rescale_m(mask, ctc_w, ctc_h, downscale_algorithm) - - # Clamp mask to [0, 1] and expand to match image channels - resized_mask = resized_mask.clamp(0, 1).unsqueeze(-1) # shape: [B, H, W, 1] - - # Extract the canvas region we're about to overwrite - canvas_crop = canvas_image[:, ctc_y:ctc_y + ctc_h, ctc_x:ctc_x + ctc_w] - - # Blend: new = mask * inpainted + (1 - mask) * canvas - blended = resized_mask * resized_image + (1.0 - resized_mask) * canvas_crop - - # Paste the blended region back onto the canvas - canvas_image[:, ctc_y:ctc_y + ctc_h, ctc_x:ctc_x + ctc_w] = blended - - # Final crop to get back the original image area - output_image = canvas_image[:, cto_y:cto_y + cto_h, cto_x:cto_x + cto_w] - - return output_image + # stitch_magic_im is inherited from ProcessorLogic class GPUProcessorLogic(ProcessorLogic): def rescale_i(self, samples, width, height, algorithm: str): # samples shape: [B, H, W, C] + width = max(1, int(width)) + height = max(1, int(height)) mode = algorithm.lower() # CPU works better, fallback to CPU for rescaling original_device = samples.device samples = samples.movedim(-1, 1) # [B, C, H, W] - algorithm_enum = getattr(Image, algorithm.upper()) + algorithm_enum = _get_pil_resampling(algorithm) results = [] for i in range(samples.shape[0]): samples_pil: Image.Image = F.to_pil_image(samples[i].float().cpu()).resize((width, height), algorithm_enum) @@ -614,49 +776,45 @@ class GPUProcessorLogic(ProcessorLogic): samples = torch.stack(results, dim=0).to(original_device) samples = samples.movedim(1, -1) return samples - - #samples = samples.movedim(-1, 1) # [B, C, H, W] - #samples = TF.interpolate(samples, size=(height, width), mode=mode, align_corners=False if mode not in ['nearest', 'area'] else None) - #samples = samples.movedim(1, -1) - #return samples def rescale_m(self, samples, width, height, algorithm: str): # samples shape: [B, H, W] + width = max(1, int(width)) + height = max(1, int(height)) + if samples.ndim == 2: + samples = samples.unsqueeze(0) mode = algorithm.lower() # CPU works better, fallback to CPU for rescaling original_device = samples.device - algorithm_enum = getattr(Image, algorithm.upper()) + algorithm_enum = _get_pil_resampling(algorithm) results = [] for i in range(samples.shape[0]): samples_pil: Image.Image = F.to_pil_image(samples[i].float().cpu()).resize((width, height), algorithm_enum) results.append(F.to_tensor(samples_pil).squeeze(0)) samples = torch.stack(results, dim=0).to(original_device) return samples - - #samples = samples.unsqueeze(1) # [B, H, W] -> [B, 1, H, W] - #samples = TF.interpolate(samples, size=(height, width), mode=mode, align_corners=False if mode not in ['nearest', 'area'] else None) - #samples = samples.squeeze(1) - #return samples def fillholes_iterative_hipass_fill_m(self, samples): - # We want this to always run in CPU for simplicity of implementation. - # Just convert whatever inputs from GPU to CPU at the beginning of the function, - # then convert them back to GPU at the end of the function. - # The implementation is verbatim from CPUProcessorLogic. - + is_2d = False + if samples.ndim == 2: + is_2d = True + samples = samples.unsqueeze(0) thresholds = [1, 0.99, 0.97, 0.95, 0.93, 0.9, 0.8, 0.7, 0.6, 0.5, 0.4, 0.3, 0.2, 0.1] results = [] original_device = samples.device for i in range(samples.shape[0]): - mask_np = samples[i].cpu().numpy() + mask_np = samples[i].float().cpu().numpy() for threshold in thresholds: thresholded_mask = mask_np >= threshold closed_mask = binary_closing(thresholded_mask, structure=np.ones((3, 3)), border_value=1) filled_mask = binary_fill_holes(closed_mask) mask_np = np.maximum(mask_np, np.where(filled_mask != 0, threshold, 0)) results.append(torch.from_numpy(mask_np.astype(np.float32))) - return torch.stack(results, dim=0).to(original_device) + res = torch.stack(results, dim=0).to(original_device) + if is_2d: + res = res.squeeze(0) + return res def hipassfilter_m(self, samples, threshold): filtered_mask = samples.clone() @@ -664,6 +822,15 @@ class GPUProcessorLogic(ProcessorLogic): return filtered_mask def expand_m(self, mask, pixels): + if pixels <= 0: + return mask + is_2d = False + if mask.ndim == 2: + is_2d = True + mask = mask.unsqueeze(0) + orig_dtype = mask.dtype + if not mask.is_floating_point(): + mask = mask.float() # Dilation can be approximated with max pooling sigma = pixels / 4 kernel_size = math.ceil(sigma * 1.5 + 1) @@ -675,22 +842,40 @@ class GPUProcessorLogic(ProcessorLogic): # mask is [B, H, W] -> [B, 1, H, W] mask_in = mask.unsqueeze(1) - # Reflect padding to avoid border transparency - mask_padded = TF.pad(mask_in, (padding, padding, padding, padding), mode='reflect') + # Reflect padding to avoid border transparency (fallback to replicate if padding >= dimension) + pad_mode = 'reflect' if (padding < mask.shape[1] and padding < mask.shape[2]) else 'replicate' + mask_padded = TF.pad(mask_in, (padding, padding, padding, padding), mode=pad_mode) # MaxPool2d is equivalent to dilation with a square kernel of 1s dilated = TF.max_pool2d(mask_padded, kernel_size=kernel_size, stride=1, padding=0) - return dilated.squeeze(1) + res = dilated.squeeze(1) + if orig_dtype == torch.bool: + res = res > 0.5 + elif not orig_dtype.is_floating_point: + res = res.to(orig_dtype) + if is_2d: + res = res.squeeze(0) + return res def invert_m(self, samples): + if samples.dtype == torch.bool: + return (~samples).float() inverted_mask = samples.clone() + if not inverted_mask.is_floating_point(): + inverted_mask = inverted_mask.float() inverted_mask = 1.0 - inverted_mask return inverted_mask def blur_m(self, samples, pixels): if pixels <= 0: return samples + is_2d = False + if samples.ndim == 2: + is_2d = True + samples = samples.unsqueeze(0) + if not samples.is_floating_point(): + samples = samples.float() sigma = pixels / 4 # Gaussian blur implementation on GPU (Separable 2-pass 1D convolution for memory optimization) kernel_size = 2 * int(4.0 * sigma + 0.5) + 1 @@ -707,20 +892,33 @@ class GPUProcessorLogic(ProcessorLogic): pad = kernel_size // 2 # Reflect padding and separable 1D convolutions (horizontal then vertical) - padded_h = TF.pad(mask_in, (pad, pad, 0, 0), mode='reflect') + pad_mode_h = 'reflect' if pad < samples.shape[2] else 'replicate' + padded_h = TF.pad(mask_in, (pad, pad, 0, 0), mode=pad_mode_h) blurred_h = TF.conv2d(padded_h, kernel_h, padding=0) - padded_v = TF.pad(blurred_h, (0, 0, pad, pad), mode='reflect') + pad_mode_v = 'reflect' if pad < samples.shape[1] else 'replicate' + padded_v = TF.pad(blurred_h, (0, 0, pad, pad), mode=pad_mode_v) blurred = TF.conv2d(padded_v, kernel_v, padding=0) - return blurred.squeeze(1).clamp(0.0, 1.0) + res = blurred.squeeze(1).clamp(0.0, 1.0) + if is_2d: + res = res.squeeze(0) + return res def debug_context_location_in_image(self, image, x, y, w, h): debug_image = image.clone() - debug_image[:, y:y+h, x:x+w, :] = 1.0 - debug_image[:, y:y+h, x:x+w, :] + B, img_h, img_w, C = image.shape + x1 = max(0, min(img_w, int(x))) + y1 = max(0, min(img_h, int(y))) + x2 = max(0, min(img_w, int(x + w))) + y2 = max(0, min(img_h, int(y + h))) + if x2 > x1 and y2 > y1: + debug_image[:, y1:y2, x1:x2, :] = 1.0 - debug_image[:, y1:y2, x1:x2, :] return debug_image def pad_to_multiple(self, value, multiple): + if multiple <= 0: + return int(value) return int(math.ceil(value / multiple) * multiple) def preresize_imm(self, image, mask, optional_context_mask, downscale_algorithm, upscale_algorithm, preresize_mode, preresize_min_width, preresize_min_height, preresize_max_width, preresize_max_height): @@ -812,9 +1010,12 @@ class GPUProcessorLogic(ProcessorLogic): assert new_H >= 0, f"Error: Trying to crop too much, height ({new_H}) must be >= 0" assert new_W >= 0, f"Error: Trying to crop too much, width ({new_W}) must be >= 0" - expanded_image = torch.zeros(B, new_H, new_W, C, device=image.device) - expanded_mask = torch.ones(B, new_H, new_W, device=mask.device) - expanded_optional_context_mask = torch.zeros(B, new_H, new_W, device=optional_context_mask.device) + if optional_context_mask is None: + optional_context_mask = torch.zeros(B, H, W, device=image.device, dtype=mask.dtype) + + expanded_image = torch.zeros(B, new_H, new_W, C, device=image.device, dtype=image.dtype) + expanded_mask = torch.ones(B, new_H, new_W, device=mask.device, dtype=mask.dtype) + expanded_optional_context_mask = torch.zeros(B, new_H, new_W, device=optional_context_mask.device, dtype=optional_context_mask.dtype) up_padding = int(H * (extend_up_factor - 1.0)) down_padding = new_H - H - up_padding @@ -943,8 +1144,13 @@ class GPUProcessorLogic(ProcessorLogic): mask = mask.clone() # Check for invalid inputs - if target_w <= 0 or target_h <= 0 or w == 0 or h == 0: - return image, 0, 0, image.shape[2], image.shape[1], image, mask, 0, 0, image.shape[2], image.shape[1] + if target_w <= 0 or target_h <= 0 or w <= 0 or h <= 0: + crop_im = image + crop_m = mask + if resize_output and target_w > 0 and target_h > 0: + crop_im = self.rescale_i(crop_im, target_w, target_h, downscale_algorithm) + crop_m = self.rescale_m(crop_m, target_w, target_h, downscale_algorithm) + return image, 0, 0, image.shape[2], image.shape[1], crop_im, crop_m, 0, 0, image.shape[2], image.shape[1] # Step 1: Pad target dimensions to be multiples of padding if padding != 0: @@ -1063,8 +1269,8 @@ class GPUProcessorLogic(ProcessorLogic): expanded_image_h += down_padding # Step 5: Create the new image and mask - expanded_image = torch.zeros((image.shape[0], expanded_image_h, expanded_image_w, image.shape[3]), device=image.device) - expanded_mask = torch.ones((mask.shape[0], expanded_image_h, expanded_image_w), device=mask.device) + expanded_image = torch.zeros((image.shape[0], expanded_image_h, expanded_image_w, image.shape[3]), device=image.device, dtype=image.dtype) + expanded_mask = torch.ones((mask.shape[0], expanded_image_h, expanded_image_w), device=mask.device, dtype=mask.dtype) # Reorder the tensors to match the required dimension format for padding image = image.permute(0, 3, 1, 2) # [B, H, W, C] -> [B, C, H, W] @@ -1122,36 +1328,7 @@ class GPUProcessorLogic(ProcessorLogic): return canvas_image, cto_x, cto_y, cto_w, cto_h, cropped_image, cropped_mask, ctc_x, ctc_y, ctc_w, ctc_h - def stitch_magic_im(self, canvas_image, inpainted_image, mask, ctc_x, ctc_y, ctc_w, ctc_h, cto_x, cto_y, cto_w, cto_h, downscale_algorithm, upscale_algorithm): - canvas_image = canvas_image.clone() - inpainted_image = inpainted_image.clone() - mask = mask.clone() - - # Resize inpainted image and mask to match the context size - B, h, w, _ = inpainted_image.shape - if ctc_w > w or ctc_h > h: # Upscaling - resized_image = self.rescale_i(inpainted_image, ctc_w, ctc_h, upscale_algorithm) - resized_mask = self.rescale_m(mask, ctc_w, ctc_h, upscale_algorithm) - else: # Downscaling - resized_image = self.rescale_i(inpainted_image, ctc_w, ctc_h, downscale_algorithm) - resized_mask = self.rescale_m(mask, ctc_w, ctc_h, downscale_algorithm) - - # Clamp mask to [0, 1] and expand to match image channels - resized_mask = resized_mask.clamp(0, 1).unsqueeze(-1) # shape: [B, H, W, 1] - - # Extract the canvas region we're about to overwrite - canvas_crop = canvas_image[:, ctc_y:ctc_y + ctc_h, ctc_x:ctc_x + ctc_w] - - # Blend: new = mask * inpainted + (1 - mask) * canvas - blended = resized_mask * resized_image + (1.0 - resized_mask) * canvas_crop - - # Paste the blended region back onto the canvas - canvas_image[:, ctc_y:ctc_y + ctc_h, ctc_x:ctc_x + ctc_w] = blended - - # Final crop to get back the original image area - output_image = canvas_image[:, cto_y:cto_y + cto_h, cto_x:cto_x + cto_w] - - return output_image + # stitch_magic_im is inherited from ProcessorLogic class InpaintCropImproved: @classmethod @@ -1210,16 +1387,10 @@ class InpaintCropImproved: CATEGORY = "inpaint" DESCRIPTION = "Crops an image around a mask for inpainting, the optional context mask defines an extra area to keep for the context." - # Remove the following # to turn on debug mode (extra outputs, print statements) - #''' - DEBUG_MODE = False - RETURN_TYPES = ("STITCHER", "IMAGE", "MASK") - RETURN_NAMES = ("stitcher", "cropped_image", "cropped_mask") + VERBOSE = True - ''' - - DEBUG_MODE = True - RETURN_TYPES = ("STITCHER", "IMAGE", "MASK", + DEBUG_RETURN_TYPES = ( + "STITCHER", "IMAGE", "MASK", # DEBUG "IMAGE", "MASK", @@ -1245,7 +1416,8 @@ class InpaintCropImproved: "IMAGE", "MASK", ) - RETURN_NAMES = ("stitcher", "cropped_image", "cropped_mask", + DEBUG_RETURN_NAMES = ( + "stitcher", "cropped_image", "cropped_mask", # DEBUG "DEBUG_preresize_image", "DEBUG_preresize_mask", @@ -1271,18 +1443,70 @@ class InpaintCropImproved: "DEBUG_cropped_in_canvas_location", "DEBUG_cropped_mask_blend", ) + + # Remove the following # to turn on debug mode (extra outputs, print statements) + #''' + DEBUG_MODE = False + RETURN_TYPES = ("STITCHER", "IMAGE", "MASK") + RETURN_NAMES = ("stitcher", "cropped_image", "cropped_mask") + + ''' + + DEBUG_MODE = True + RETURN_TYPES = DEBUG_RETURN_TYPES + RETURN_NAMES = DEBUG_RETURN_NAMES #''' def inpaint_crop(self, image, downscale_algorithm, upscale_algorithm, preresize, preresize_mode, preresize_min_width, preresize_min_height, preresize_max_width, preresize_max_height, extend_for_outpainting, extend_up_factor, extend_down_factor, extend_left_factor, extend_right_factor, mask_hipass_filter, mask_fill_holes, mask_expand_pixels, mask_invert, mask_blend_pixels, context_from_mask_extend_factor, output_resize_to_target_size, output_target_width, output_target_height, output_padding, device_mode, mask=None, optional_context_mask=None): image = image.clone() + if not image.is_floating_point(): + image = image.float() / 255.0 if (image.numel() > 0 and image.max() > 1.0) else image.float() + if image.ndim == 3: + image = image.unsqueeze(0) # (H, W, C) -> (1, H, W, C) + + # Defensive layer handling for input image channels (1->RGB, 2->RGBA) + if image.shape[-1] == 1: + image = image.repeat(1, 1, 1, 3) + elif image.shape[-1] == 2: + image = torch.cat([image[..., 0:1].repeat(1, 1, 1, 3), image[..., 1:2]], dim=-1) + + if mask is not None and mask.numel() == 0: + mask = None + if mask is not None: mask = mask.clone() + if not mask.is_floating_point(): + mask = mask.float() if mask.ndim == 2: mask = mask.unsqueeze(0) # (H, W) -> (1, H, W) + elif mask.ndim == 4: + if mask.shape[1] == 1: + mask = mask.squeeze(1) + elif mask.shape[-1] == 1: + mask = mask.squeeze(-1) + # Defensive range normalization (e.g. 0-255 masks) + if mask.numel() > 0 and mask.max() > 1.0: + mask = mask / 255.0 + mask = mask.clamp(0.0, 1.0) + + if optional_context_mask is not None and optional_context_mask.numel() == 0: + optional_context_mask = None + if optional_context_mask is not None: optional_context_mask = optional_context_mask.clone() + if not optional_context_mask.is_floating_point(): + optional_context_mask = optional_context_mask.float() if optional_context_mask.ndim == 2: optional_context_mask = optional_context_mask.unsqueeze(0) # (H, W) -> (1, H, W) + elif optional_context_mask.ndim == 4: + if optional_context_mask.shape[1] == 1: + optional_context_mask = optional_context_mask.squeeze(1) + elif optional_context_mask.shape[-1] == 1: + optional_context_mask = optional_context_mask.squeeze(-1) + # Defensive range normalization + if optional_context_mask.numel() > 0 and optional_context_mask.max() > 1.0: + optional_context_mask = optional_context_mask / 255.0 + optional_context_mask = optional_context_mask.clamp(0.0, 1.0) if device_mode == "gpu (much faster)": device = comfy.model_management.get_torch_device() @@ -1293,14 +1517,18 @@ class InpaintCropImproved: else: processor = CPUProcessorLogic() - output_padding = int(output_padding) - - # Check that some parameters make sense - if preresize and preresize_mode == "ensure minimum and maximum resolution": - assert preresize_max_width >= preresize_min_width, "Preresize maximum width must be greater than or equal to minimum width" - assert preresize_max_height >= preresize_min_height, "Preresize maximum height must be greater than or equal to minimum height" + output_padding = max(0, int(output_padding)) + output_target_width = max(1, int(output_target_width)) + output_target_height = max(1, int(output_target_height)) - if self.DEBUG_MODE: + # Check that resolution parameters make sense (swap if min > max) + if preresize and preresize_mode == "ensure minimum and maximum resolution": + if preresize_min_width > preresize_max_width: + preresize_min_width, preresize_max_width = preresize_max_width, preresize_min_width + if preresize_min_height > preresize_max_height: + preresize_min_height, preresize_max_height = preresize_max_height, preresize_min_height + + if self.DEBUG_MODE and getattr(self, "VERBOSE", True): print('Inpaint Crop Batch input') print(image.shape, type(image), image.dtype) if mask is not None: @@ -1308,8 +1536,13 @@ class InpaintCropImproved: if optional_context_mask is not None: print(optional_context_mask.shape, type(optional_context_mask), optional_context_mask.dtype) - if image.shape[0] > 1: - assert output_resize_to_target_size, "output_resize_to_target_size must be enabled when input is a batch of images, given all images in the batch output have to be the same size" + # Batch inputs require uniform output target sizes + if image.shape[0] > 1 and not output_resize_to_target_size: + output_resize_to_target_size = True + if output_target_width <= 1: + output_target_width = image.shape[2] + if output_target_height <= 1: + output_target_height = image.shape[1] # When a LoadImage node passes a mask without user editing, it may be the wrong shape. # Detect and fix that to avoid shape mismatch errors. @@ -1317,11 +1550,15 @@ class InpaintCropImproved: if mask.shape[1] != image.shape[1] or mask.shape[2] != image.shape[2]: if torch.count_nonzero(mask) == 0: mask = torch.zeros((mask.shape[0], image.shape[1], image.shape[2]), device=image.device, dtype=image.dtype) + else: + mask = processor.rescale_m(mask, image.shape[2], image.shape[1], "bilinear") if optional_context_mask is not None and (image.shape[0] == 1 or optional_context_mask.shape[0] == 1 or optional_context_mask.shape[0] == image.shape[0]): if optional_context_mask.shape[1] != image.shape[1] or optional_context_mask.shape[2] != image.shape[2]: if torch.count_nonzero(optional_context_mask) == 0: optional_context_mask = torch.zeros((optional_context_mask.shape[0], image.shape[1], image.shape[2]), device=image.device, dtype=image.dtype) + else: + optional_context_mask = processor.rescale_m(optional_context_mask, image.shape[2], image.shape[1], "bilinear") # If no mask is provided, create one with the shape of the image if mask is None: @@ -1346,7 +1583,7 @@ class InpaintCropImproved: assert optional_context_mask.dim() == 3, f"Expected 3D BHW optional_context_mask tensor, got {optional_context_mask.shape}" optional_context_mask = optional_context_mask.expand(image.shape[0], -1, -1).clone() - if self.DEBUG_MODE: + if self.DEBUG_MODE and getattr(self, "VERBOSE", True): print('Inpaint Crop Batch ready') print(image.shape, type(image), image.dtype) print(mask.shape, type(mask), mask.dtype) @@ -1539,7 +1776,8 @@ class InpaintCropImproved: if self.DEBUG_MODE: # Everything is already on CPU, stack will be memory-safe final_debug_outputs = [] - for name in self.RETURN_NAMES: + return_names = getattr(self, "DEBUG_RETURN_NAMES", self.RETURN_NAMES) if len(self.RETURN_NAMES) <= 3 else self.RETURN_NAMES + for name in return_names: if name.startswith("DEBUG_"): values = debug_outputs[name] if not values: @@ -1592,7 +1830,19 @@ class InpaintStitchImproved: def inpaint_stitch(self, stitcher, inpainted_image): inpainted_image = inpainted_image.clone() + if not inpainted_image.is_floating_point(): + inpainted_image = inpainted_image.float() / 255.0 if (inpainted_image.numel() > 0 and inpainted_image.max() > 1.0) else inpainted_image.float() + if inpainted_image.ndim == 3: + if inpainted_image.shape[-1] in [1, 3, 4]: + inpainted_image = inpainted_image.unsqueeze(0) + else: + inpainted_image = inpainted_image.unsqueeze(-1) results = [] + + required_keys = ['cropped_to_canvas_x', 'cropped_to_canvas_y', 'cropped_to_canvas_w', 'cropped_to_canvas_h', 'canvas_image', 'cropped_mask_for_blend', 'canvas_to_orig_x', 'canvas_to_orig_y', 'canvas_to_orig_w', 'canvas_to_orig_h'] + for k in required_keys: + if k not in stitcher: + raise ValueError(f"InpaintStitchImproved: Provided stitcher is missing required key '{k}'. Ensure it was generated by InpaintCropImproved.") device_mode = stitcher.get('device_mode', 'cpu (compatible)') @@ -1610,10 +1860,7 @@ class InpaintStitchImproved: stitcher[key] = [t.to(device) if torch.is_tensor(t) else t for t in stitcher[key]] batch_size = inpainted_image.shape[0] - assert len(stitcher['cropped_to_canvas_x']) == batch_size or len(stitcher['cropped_to_canvas_x']) == 1, "Stitch batch size doesn't match image batch size" - override = False - if len(stitcher['cropped_to_canvas_x']) != batch_size and len(stitcher['cropped_to_canvas_x']) == 1: - override = True + stitcher_len = max(1, len(stitcher['cropped_to_canvas_x'])) for i in range(batch_size): one_image = inpainted_image[i:i+1] @@ -1622,10 +1869,8 @@ class InpaintStitchImproved: for key in ['downscale_algorithm', 'upscale_algorithm', 'blend_pixels']: one_stitcher[key] = stitcher[key] for key in ['canvas_to_orig_x', 'canvas_to_orig_y', 'canvas_to_orig_w', 'canvas_to_orig_h', 'canvas_image', 'cropped_to_canvas_x', 'cropped_to_canvas_y', 'cropped_to_canvas_w', 'cropped_to_canvas_h', 'cropped_mask_for_blend']: - if override: - one_stitcher[key] = stitcher[key][0] - else: - one_stitcher[key] = stitcher[key][i] + idx = 0 if stitcher_len == 1 else (i % stitcher_len) + one_stitcher[key] = stitcher[key][idx] one_image, = self.inpaint_stitch_single_image(one_stitcher, one_image, processor) results.append(one_image.squeeze(0)) diff --git a/pyproject.toml b/pyproject.toml index 6004392..b686883 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui-inpaint-cropandstitch" description = "The '✂️ Inpaint Crop' and '✂️ Inpaint Stitch' nodes enable inpainting only on masked area very easily: crop the image around the masked area with the Crop node, then use any standard workflow for sampling, then connect the sampled image to the Stitch node, which will put it back in place in the original image. These nodes enable faster sampling of smaller areas and take care of downsampling and upsampling to fit specific model and resource needs." -version = "3.0.14" +version = "3.0.15" license = { file = "LICENSE" } [project.urls] diff --git a/run_tests.sh b/run_tests.sh new file mode 100755 index 0000000..702bdbf --- /dev/null +++ b/run_tests.sh @@ -0,0 +1,30 @@ +#!/usr/bin/env bash +set -e + +# ComfyUI-Inpaint-CropAndStitch Test Runner +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +cd "$SCRIPT_DIR" + +echo "============================================================" +echo " Running ComfyUI-Inpaint-CropAndStitch Unit & Integration Tests " +echo "============================================================" + +# Ensure python3 is available +if ! command -v python3 &>/dev/null; then + echo "Error: python3 is not installed or not in PATH." + exit 1 +fi + +# Run test discovery +if python3 -m unittest discover -s tests -p "test_*.py" -v "$@"; then + echo "============================================================" + echo " [SUCCESS] All tests passed successfully!" + echo "============================================================" + exit 0 +else + EXIT_CODE=$? + echo "============================================================" + echo " [FAILURE] Some tests failed. Exit code: $EXIT_CODE" + echo "============================================================" + exit $EXIT_CODE +fi diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..d4839a6 --- /dev/null +++ b/tests/__init__.py @@ -0,0 +1 @@ +# Tests package diff --git a/tests/test_context_areas.py b/tests/test_context_areas.py new file mode 100644 index 0000000..19330fa --- /dev/null +++ b/tests/test_context_areas.py @@ -0,0 +1,127 @@ +import unittest +import torch +from inpaint_cropandstitch import CPUProcessorLogic, GPUProcessorLogic + + +class TestContextAreas(unittest.TestCase): + def setUp(self): + self.processors = [ + ("cpu", CPUProcessorLogic()), + ("gpu", GPUProcessorLogic()), + ] + + def test_findcontextarea_empty_mask(self): + for name, proc in self.processors: + with self.subTest(processor=name): + # All zeros mask should return -1 indicator + mask = torch.zeros(1, 40, 40, dtype=torch.float32) + _, bx, by, bw, bh = proc.batched_findcontextarea_m(mask) + self.assertEqual(bx[0].item(), -1) + self.assertEqual(by[0].item(), -1) + self.assertEqual(bw[0].item(), -1) + self.assertEqual(bh[0].item(), -1) + + def test_findcontextarea_single_pixel_center(self): + for name, proc in self.processors: + with self.subTest(processor=name): + mask = torch.zeros(1, 50, 50, dtype=torch.float32) + mask[0, 25, 25] = 1.0 + _, bx, by, bw, bh = proc.batched_findcontextarea_m(mask) + self.assertEqual(bx[0].item(), 25) + self.assertEqual(by[0].item(), 25) + self.assertEqual(bw[0].item(), 1) + self.assertEqual(bh[0].item(), 1) + + def test_findcontextarea_borders(self): + for name, proc in self.processors: + with self.subTest(processor=name): + # Pixel at top-left (0, 0) + mask1 = torch.zeros(1, 50, 50, dtype=torch.float32) + mask1[0, 0, 0] = 1.0 + _, bx1, by1, bw1, bh1 = proc.batched_findcontextarea_m(mask1) + self.assertEqual(bx1[0].item(), 0) + self.assertEqual(by1[0].item(), 0) + + # Pixel at bottom-right (49, 49) + mask2 = torch.zeros(1, 50, 50, dtype=torch.float32) + mask2[0, 49, 49] = 1.0 + _, bx2, by2, bw2, bh2 = proc.batched_findcontextarea_m(mask2) + self.assertEqual(bx2[0].item(), 49) + self.assertEqual(by2[0].item(), 49) + + def test_findcontextarea_multi_blob(self): + for name, proc in self.processors: + with self.subTest(processor=name): + mask = torch.zeros(1, 100, 100, dtype=torch.float32) + # Two blobs + mask[0, 10:20, 10:20] = 1.0 + mask[0, 60:80, 50:70] = 1.0 + _, bx, by, bw, bh = proc.batched_findcontextarea_m(mask) + self.assertEqual(bx[0].item(), 10) + self.assertEqual(by[0].item(), 10) + self.assertEqual(bw[0].item(), 60) # 69 - 10 + 1 + self.assertEqual(bh[0].item(), 70) # 79 - 10 + 1 + + def test_findcontextarea_batch_dimension(self): + for name, proc in self.processors: + with self.subTest(processor=name): + # Batch of 2 different masks + mask = torch.zeros(2, 60, 60, dtype=torch.float32) + mask[0, 10:20, 10:20] = 1.0 + mask[1, 30:50, 30:50] = 1.0 + _, bx, by, bw, bh = proc.batched_findcontextarea_m(mask) + self.assertEqual(len(bx), 2) + self.assertEqual(bx[0].item(), 10) + self.assertEqual(bw[0].item(), 10) + self.assertEqual(bx[1].item(), 30) + self.assertEqual(bw[1].item(), 20) + + def test_growcontextarea(self): + for name, proc in self.processors: + with self.subTest(processor=name): + mask = torch.zeros(1, 100, 100, dtype=torch.float32) + mask[0, 40:60, 40:60] = 1.0 + _, x, y, w, h = proc.batched_findcontextarea_m(mask) + + # Grow with factor 1.0 (no change) + _, gx1, gy1, gw1, gh1 = proc.batched_growcontextarea_m(mask, x, y, w, h, extend_factor=1.0) + self.assertEqual(gw1[0].item(), 20) + self.assertEqual(gh1[0].item(), 20) + + # Grow with factor 2.0 (expand around center) + _, gx2, gy2, gw2, gh2 = proc.batched_growcontextarea_m(mask, x, y, w, h, extend_factor=2.0) + self.assertGreater(gw2[0].item(), 20) + self.assertGreater(gh2[0].item(), 20) + self.assertLess(gx2[0].item(), 40) + self.assertLess(gy2[0].item(), 40) + + # Empty mask (w == -1) -> should fill entire image + empty_w = torch.tensor([-1]) + empty_h = torch.tensor([-1]) + empty_x = torch.tensor([-1]) + empty_y = torch.tensor([-1]) + _, egx, egy, egw, egh = proc.batched_growcontextarea_m(mask, empty_x, empty_y, empty_w, empty_h, extend_factor=1.5) + self.assertEqual(egx[0].item(), 0) + self.assertEqual(egy[0].item(), 0) + self.assertEqual(egw[0].item(), 100) + self.assertEqual(egh[0].item(), 100) + + def test_combinecontextmask(self): + for name, proc in self.processors: + with self.subTest(processor=name): + mask = torch.zeros(1, 100, 100, dtype=torch.float32) + mask[0, 40:60, 40:60] = 1.0 + _, x, y, w, h = proc.batched_findcontextarea_m(mask) + + opt_mask = torch.zeros(1, 100, 100, dtype=torch.float32) + opt_mask[0, 10:20, 10:20] = 1.0 + + _, cx, cy, cw, ch = proc.batched_combinecontextmask_m(mask, x, y, w, h, opt_mask) + self.assertEqual(cx[0].item(), 10) + self.assertEqual(cy[0].item(), 10) + self.assertGreaterEqual(cx[0].item() + cw[0].item(), 60) + self.assertGreaterEqual(cy[0].item() + ch[0].item(), 60) + + +if __name__ == '__main__': + unittest.main() diff --git a/tests/test_crop_and_stitch.py b/tests/test_crop_and_stitch.py new file mode 100644 index 0000000..a63bc0b --- /dev/null +++ b/tests/test_crop_and_stitch.py @@ -0,0 +1,232 @@ +import unittest +import torch +from inpaint_cropandstitch import CPUProcessorLogic, GPUProcessorLogic + + +class TestCropAndStitch(unittest.TestCase): + def setUp(self): + self.processors = [ + ("cpu", CPUProcessorLogic()), + ("gpu", GPUProcessorLogic()), + ] + + def test_crop_magic_im_normal(self): + for name, proc in self.processors: + with self.subTest(processor=name): + img = torch.rand(1, 100, 100, 3, dtype=torch.float32) + mask = torch.zeros(1, 100, 100, dtype=torch.float32) + canvas, cto_x, cto_y, cto_w, cto_h, crop_im, crop_m, ctc_x, ctc_y, ctc_w, ctc_h = proc.crop_magic_im( + img, mask, x=20, y=20, w=40, h=40, target_w=64, target_h=64, padding=8, + downscale_algorithm="bilinear", upscale_algorithm="bicubic", resize_output=True + ) + self.assertEqual(crop_im.shape, (1, 64, 64, 3)) + self.assertEqual(crop_m.shape, (1, 64, 64)) + self.assertEqual(crop_im.dtype, torch.float32) + + def test_crop_magic_im_non_positive_dimensions(self): + for name, proc in self.processors: + with self.subTest(processor=name): + img = torch.rand(1, 50, 50, 3, dtype=torch.float32) + mask = torch.zeros(1, 50, 50, dtype=torch.float32) + # w=0, h=0 + res = proc.crop_magic_im( + img, mask, x=0, y=0, w=0, h=0, target_w=64, target_h=64, padding=8, + downscale_algorithm="bilinear", upscale_algorithm="bicubic", resize_output=True + ) + self.assertEqual(res[5].shape, (1, 64, 64, 3)) + + # target_w <= 0 + res_bad_target = proc.crop_magic_im( + img, mask, x=0, y=0, w=20, h=20, target_w=-10, target_h=0, padding=8, + downscale_algorithm="bilinear", upscale_algorithm="bicubic", resize_output=True + ) + self.assertEqual(res_bad_target[5].shape, img.shape) + + def test_crop_magic_im_channels(self): + for name, proc in self.processors: + for c in [1, 3, 4]: + with self.subTest(processor=name, channels=c): + img = torch.rand(1, 60, 60, c, dtype=torch.float32) + mask = torch.zeros(1, 60, 60, dtype=torch.float32) + _, _, _, _, _, crop_im, crop_m, _, _, _, _ = proc.crop_magic_im( + img, mask, x=10, y=10, w=30, h=30, target_w=32, target_h=32, padding=0, + downscale_algorithm="bilinear", upscale_algorithm="bicubic", resize_output=True + ) + self.assertEqual(crop_im.shape[-1], c) + + def test_stitch_magic_im_reproduce_rgba_canvas_rgb_inpaint(self): + """ + Direct test for the user-reported bug: + RuntimeError: The size of tensor a (3) must match the size of tensor b (4) at non-singleton dimension 3 + blended = resized_mask * resized_image + (1.0 - resized_mask) * canvas_crop + """ + for name, proc in self.processors: + with self.subTest(processor=name): + # Canvas is RGBA (4 channels) + canvas_image = torch.ones(1, 100, 100, 4, dtype=torch.float32) + # Set specific alpha in canvas to verify preservation + canvas_image[:, :, :, 3] = 0.75 + + # Inpainted crop is RGB (3 channels) as returned by standard VAE decode + inpainted_image = torch.zeros(1, 40, 40, 3, dtype=torch.float32) + # Mask + mask = torch.ones(1, 40, 40, dtype=torch.float32) + + ctc_x, ctc_y, ctc_w, ctc_h = 20, 20, 40, 40 + cto_x, cto_y, cto_w, cto_h = 0, 0, 100, 100 + + output = proc.stitch_magic_im( + canvas_image, inpainted_image, mask, + ctc_x, ctc_y, ctc_w, ctc_h, + cto_x, cto_y, cto_w, cto_h, + downscale_algorithm="bilinear", upscale_algorithm="bicubic" + ) + + # Output should have 4 channels and match original canvas size + self.assertEqual(output.shape, (1, 100, 100, 4)) + # Stitched region RGB should be 0.0 (from inpaint) + self.assertAlmostEqual(output[0, 30, 30, 0].item(), 0.0) + # Stitched region Alpha should preserve original canvas alpha (0.75) + self.assertAlmostEqual(output[0, 30, 30, 3].item(), 0.75) + + def test_stitch_magic_im_rgb_canvas_rgba_inpaint(self): + # Canvas is RGB (3 channels), Inpaint is RGBA (4 channels) + for name, proc in self.processors: + with self.subTest(processor=name): + canvas_image = torch.ones(1, 100, 100, 3, dtype=torch.float32) + inpainted_image = torch.zeros(1, 40, 40, 4, dtype=torch.float32) + mask = torch.ones(1, 40, 40, dtype=torch.float32) + + output = proc.stitch_magic_im( + canvas_image, inpainted_image, mask, + 20, 20, 40, 40, 0, 0, 100, 100, + downscale_algorithm="bilinear", upscale_algorithm="bicubic" + ) + self.assertEqual(output.shape, (1, 100, 100, 3)) + self.assertAlmostEqual(output[0, 30, 30, 0].item(), 0.0) + + def test_stitch_magic_im_channel_matching_variants(self): + # Grayscale canvas (1) + RGB inpaint (3) + for name, proc in self.processors: + with self.subTest(processor=name, mode="gray_rgb"): + canvas_image = torch.ones(1, 80, 80, 1, dtype=torch.float32) + inpainted_image = torch.zeros(1, 30, 30, 3, dtype=torch.float32) + mask = torch.ones(1, 30, 30, dtype=torch.float32) + output = proc.stitch_magic_im( + canvas_image, inpainted_image, mask, + 10, 10, 30, 30, 0, 0, 80, 80, + "bilinear", "bicubic" + ) + self.assertEqual(output.shape, (1, 80, 80, 1)) + + # RGB canvas (3) + Grayscale inpaint (1) + with self.subTest(processor=name, mode="rgb_gray"): + canvas_image = torch.ones(1, 80, 80, 3, dtype=torch.float32) + inpainted_image = torch.zeros(1, 30, 30, 1, dtype=torch.float32) + mask = torch.ones(1, 30, 30, dtype=torch.float32) + output = proc.stitch_magic_im( + canvas_image, inpainted_image, mask, + 10, 10, 30, 30, 0, 0, 80, 80, + "bilinear", "bicubic" + ) + self.assertEqual(output.shape, (1, 80, 80, 3)) + + def test_stitch_magic_im_boundary_clipping_and_out_of_bounds(self): + for name, proc in self.processors: + with self.subTest(processor=name): + canvas_image = torch.ones(1, 100, 100, 3, dtype=torch.float32) + inpainted_image = torch.zeros(1, 50, 50, 3, dtype=torch.float32) + mask = torch.ones(1, 50, 50, dtype=torch.float32) + + # Coordinate extending past image right and bottom: x=80, w=50 (exceeds 100) + output = proc.stitch_magic_im( + canvas_image, inpainted_image, mask, + ctc_x=80, ctc_y=80, ctc_w=50, ctc_h=50, + cto_x=0, cto_y=0, cto_w=100, cto_h=100, + downscale_algorithm="bilinear", upscale_algorithm="bicubic" + ) + self.assertEqual(output.shape, (1, 100, 100, 3)) + + # Negative coordinates: x=-10, y=-10 + output_neg = proc.stitch_magic_im( + canvas_image, inpainted_image, mask, + ctc_x=-10, ctc_y=-10, ctc_w=50, ctc_h=50, + cto_x=0, cto_y=0, cto_w=100, cto_h=100, + downscale_algorithm="bilinear", upscale_algorithm="bicubic" + ) + self.assertEqual(output_neg.shape, (1, 100, 100, 3)) + + # Completely outside canvas: x=500, y=500 + output_disjoint = proc.stitch_magic_im( + canvas_image, inpainted_image, mask, + ctc_x=500, ctc_y=500, ctc_w=50, ctc_h=50, + cto_x=0, cto_y=0, cto_w=100, cto_h=100, + downscale_algorithm="bilinear", upscale_algorithm="bicubic" + ) + self.assertEqual(output_disjoint.shape, (1, 100, 100, 3)) + + def test_stitch_magic_im_dtypes_and_shapes(self): + for name, proc in self.processors: + with self.subTest(processor=name): + # canvas float32, inpaint float16 + canvas = torch.ones(1, 60, 60, 3, dtype=torch.float32) + inpaint = torch.zeros(1, 30, 30, 3, dtype=torch.float16) + mask = torch.ones(1, 30, 30, dtype=torch.float32) + + out = proc.stitch_magic_im( + canvas, inpaint, mask, + 15, 15, 30, 30, 0, 0, 60, 60, + "bilinear", "bicubic" + ) + self.assertEqual(out.dtype, torch.float32) + + # 3D inpainted image [H, W, C] without batch dimension + inpaint_3d = torch.zeros(30, 30, 3, dtype=torch.float32) + mask_2d = torch.ones(30, 30, dtype=torch.float32) + out_3d = proc.stitch_magic_im( + canvas, inpaint_3d, mask_2d, + 15, 15, 30, 30, 0, 0, 60, 60, + "bilinear", "bicubic" + ) + def test_crop_magic_im_pixel_accuracy(self): + for name, proc in self.processors: + with self.subTest(processor=name): + img = torch.zeros(1, 50, 50, 3, dtype=torch.float32) + # Specific marker pixel + img[0, 25, 25, :] = torch.tensor([0.123, 0.456, 0.789]) + mask = torch.zeros(1, 50, 50, dtype=torch.float32) + + # Crop 10x10 around (20, 20) without resize + canvas, cto_x, cto_y, cto_w, cto_h, crop_im, crop_m, ctc_x, ctc_y, ctc_w, ctc_h = proc.crop_magic_im( + img, mask, x=20, y=20, w=10, h=10, target_w=10, target_h=10, padding=0, + downscale_algorithm="bilinear", upscale_algorithm="bicubic", resize_output=False + ) + # Marker should be at local offset (5, 5) + self.assertAlmostEqual(crop_im[0, 5, 5, 0].item(), 0.123, places=3) + self.assertAlmostEqual(crop_im[0, 5, 5, 1].item(), 0.456, places=3) + self.assertAlmostEqual(crop_im[0, 5, 5, 2].item(), 0.789, places=3) + + def test_stitch_magic_im_exact_blending(self): + for name, proc in self.processors: + with self.subTest(processor=name): + # Canvas has value 0.2 + canvas = torch.full((1, 40, 40, 3), 0.2, dtype=torch.float32) + # Inpaint has value 0.8 + inpaint = torch.full((1, 20, 20, 3), 0.8, dtype=torch.float32) + # Mask has value 0.5 (exact 50/50 blend) + mask = torch.full((1, 20, 20), 0.5, dtype=torch.float32) + + out = proc.stitch_magic_im( + canvas, inpaint, mask, + ctc_x=10, ctc_y=10, ctc_w=20, ctc_h=20, + cto_x=0, cto_y=0, cto_w=40, cto_h=40, + downscale_algorithm="bilinear", upscale_algorithm="bicubic" + ) + # 0.5 * 0.8 + 0.5 * 0.2 = 0.500 (allowing for 8-bit PIL quantization) + self.assertAlmostEqual(out[0, 15, 15, 0].item(), 0.5, places=2) + # Outside stitched box, canvas remains 0.2 + self.assertAlmostEqual(out[0, 0, 0, 0].item(), 0.2, places=2) + + +if __name__ == '__main__': + unittest.main() diff --git a/tests/test_nodes_pipeline.py b/tests/test_nodes_pipeline.py new file mode 100644 index 0000000..b743749 --- /dev/null +++ b/tests/test_nodes_pipeline.py @@ -0,0 +1,366 @@ +import unittest +import torch +from inpaint_cropandstitch import InpaintCropImproved, InpaintStitchImproved + + +class TestNodesPipeline(unittest.TestCase): + def setUp(self): + self.crop_node = InpaintCropImproved() + self.stitch_node = InpaintStitchImproved() + + def _default_crop_args(self, image, mask=None, optional_context_mask=None, **kwargs): + args = { + "image": image, + "mask": mask, + "optional_context_mask": optional_context_mask, + "downscale_algorithm": "bilinear", + "upscale_algorithm": "bicubic", + "preresize": False, + "preresize_mode": "no preresize", + "preresize_min_width": 256, + "preresize_min_height": 256, + "preresize_max_width": 1024, + "preresize_max_height": 1024, + "extend_for_outpainting": False, + "extend_up_factor": 1.0, + "extend_down_factor": 1.0, + "extend_left_factor": 1.0, + "extend_right_factor": 1.0, + "mask_hipass_filter": 0.0, + "mask_fill_holes": False, + "mask_expand_pixels": 0, + "mask_invert": False, + "mask_blend_pixels": 0, + "context_from_mask_extend_factor": 1.2, + "output_resize_to_target_size": False, + "output_target_width": 256, + "output_target_height": 256, + "output_padding": 8, + "device_mode": "cpu (compatible)", + } + args.update(kwargs) + return args + + def test_node_metadata(self): + crop_inputs = InpaintCropImproved.INPUT_TYPES() + self.assertIn("required", crop_inputs) + self.assertIn("image", crop_inputs["required"]) + self.assertEqual(InpaintCropImproved.FUNCTION, "inpaint_crop") + + stitch_inputs = InpaintStitchImproved.INPUT_TYPES() + self.assertIn("required", stitch_inputs) + self.assertIn("stitcher", stitch_inputs["required"]) + self.assertIn("inpainted_image", stitch_inputs["required"]) + self.assertEqual(InpaintStitchImproved.FUNCTION, "inpaint_stitch") + + def test_roundtrip_rgb_pipeline(self): + # Standard RGB image [1, 100, 100, 3] + orig_img = torch.rand(1, 100, 100, 3, dtype=torch.float32) + mask = torch.zeros(1, 100, 100, dtype=torch.float32) + mask[0, 30:50, 30:50] = 1.0 + + args = self._default_crop_args(orig_img, mask=mask) + crop_results = self.crop_node.inpaint_crop(**args) + + stitcher, cropped_image, cropped_mask = crop_results[:3] + + self.assertEqual(cropped_image.ndim, 4) + self.assertEqual(cropped_mask.ndim, 3) + + # Simulate inpainting: invert color in crop + inpainted_crop = 1.0 - cropped_image + stitch_results = self.stitch_node.inpaint_stitch(stitcher, inpainted_crop) + output_image = stitch_results[0] + + self.assertEqual(output_image.shape, orig_img.shape) + # Inside inpaint area, output should reflect inverted crop + self.assertNotEqual(output_image[0, 40, 40, 0].item(), orig_img[0, 40, 40, 0].item()) + # Outside inpaint area, output should match original image + self.assertAlmostEqual(output_image[0, 5, 5, 0].item(), orig_img[0, 5, 5, 0].item(), places=5) + + def test_roundtrip_rgba_canvas_rgb_inpaint_pipeline(self): + """ + End-to-end integration test of the reported bug: + RGBA canvas input into InpaintCropImproved, standard RGB output from VAE into InpaintStitchImproved. + """ + # Canvas has 4 channels (RGBA) + rgba_img = torch.rand(1, 80, 80, 4, dtype=torch.float32) + rgba_img[:, :, :, 3] = 0.8 # Specific alpha channel + mask = torch.zeros(1, 80, 80, dtype=torch.float32) + mask[0, 20:40, 20:40] = 1.0 + + args = self._default_crop_args(rgba_img, mask=mask) + crop_results = self.crop_node.inpaint_crop(**args) + + stitcher, cropped_image, _ = crop_results[:3] + + # Simulate VAE decoding only 3 channels (RGB) + rgb_inpainted_crop = cropped_image[..., :3].clone() * 0.5 + + # Stitch back together + stitch_results = self.stitch_node.inpaint_stitch(stitcher, rgb_inpainted_crop) + final_image = stitch_results[0] + + # Final image must be 4 channels (RGBA preserved) + self.assertEqual(final_image.shape, (1, 80, 80, 4)) + # Alpha channel must be preserved from original canvas + self.assertAlmostEqual(final_image[0, 30, 30, 3].item(), 0.8, places=5) + + def test_roundtrip_rgb_canvas_rgba_inpaint_pipeline(self): + # Canvas has 3 channels (RGB), Inpainted crop has 4 channels (RGBA) + rgb_img = torch.rand(1, 80, 80, 3, dtype=torch.float32) + mask = torch.zeros(1, 80, 80, dtype=torch.float32) + mask[0, 20:40, 20:40] = 1.0 + + args = self._default_crop_args(rgb_img, mask=mask) + crop_results = self.crop_node.inpaint_crop(**args) + + stitcher, cropped_image, _ = crop_results[:3] + + # Inpainted crop has an extra alpha channel + rgba_inpainted_crop = torch.cat([cropped_image, torch.ones_like(cropped_image[..., :1])], dim=-1) + + stitch_results = self.stitch_node.inpaint_stitch(stitcher, rgba_inpainted_crop) + final_image = stitch_results[0] + + # Final image must be 3 channels + self.assertEqual(final_image.shape, (1, 80, 80, 3)) + + def test_mask_dimensions_normalization(self): + orig_img = torch.rand(1, 60, 60, 3, dtype=torch.float32) + + # 2D mask [H, W] + mask_2d = torch.zeros(60, 60, dtype=torch.float32) + mask_2d[20:30, 20:30] = 1.0 + res_2d = self.crop_node.inpaint_crop(**self._default_crop_args(orig_img, mask=mask_2d)) + self.assertEqual(res_2d[1].ndim, 4) + + # 4D mask [B, 1, H, W] + mask_4d_b1hw = mask_2d.unsqueeze(0).unsqueeze(1) + res_4d_1 = self.crop_node.inpaint_crop(**self._default_crop_args(orig_img, mask=mask_4d_b1hw)) + self.assertEqual(res_4d_1[1].ndim, 4) + + # 4D mask [B, H, W, 1] + mask_4d_bhw1 = mask_2d.unsqueeze(0).unsqueeze(-1) + res_4d_2 = self.crop_node.inpaint_crop(**self._default_crop_args(orig_img, mask=mask_4d_bhw1)) + self.assertEqual(res_4d_2[1].ndim, 4) + + def test_unbatched_image_input(self): + # 3D image [H, W, C] + img_3d = torch.rand(60, 60, 3, dtype=torch.float32) + mask_2d = torch.zeros(60, 60, dtype=torch.float32) + mask_2d[20:30, 20:30] = 1.0 + + res = self.crop_node.inpaint_crop(**self._default_crop_args(img_3d, mask=mask_2d)) + self.assertEqual(res[1].ndim, 4) + + def test_outpainting_extension_pipeline(self): + orig_img = torch.rand(1, 60, 60, 3, dtype=torch.float32) + mask = torch.zeros(1, 60, 60, dtype=torch.float32) + + args = self._default_crop_args( + orig_img, mask=mask, + extend_for_outpainting=True, + extend_up_factor=1.5, + extend_down_factor=1.2, + extend_left_factor=1.0, + extend_right_factor=1.0 + ) + crop_results = self.crop_node.inpaint_crop(**args) + stitcher, cropped_image, _ = crop_results[:3] + + stitch_results = self.stitch_node.inpaint_stitch(stitcher, cropped_image) + # Stitched image has the extended (outpainted) image size: 60 * 1.7 = 102 height + self.assertEqual(stitch_results[0].shape, (1, 102, 60, 3)) + + def test_preresize_modes_pipeline(self): + orig_img = torch.rand(1, 50, 50, 3, dtype=torch.float32) + mask = torch.zeros(1, 50, 50, dtype=torch.float32) + mask[0, 10:20, 10:20] = 1.0 + + # ensure minimum resolution + args_min = self._default_crop_args( + orig_img, mask=mask, + preresize=True, + preresize_mode="ensure minimum resolution", + preresize_min_width=100, + preresize_min_height=100 + ) + res_min = self.crop_node.inpaint_crop(**args_min) + stitcher = res_min[0] + # canvas_image is stored as a list of images [img_1, img_2, ...] + self.assertGreaterEqual(stitcher['canvas_image'][0].shape[2], 100) + + def test_device_mode_gpu(self): + orig_img = torch.rand(1, 40, 40, 3, dtype=torch.float32) + mask = torch.zeros(1, 40, 40, dtype=torch.float32) + mask[0, 10:20, 10:20] = 1.0 + + args = self._default_crop_args(orig_img, mask=mask, device_mode="gpu (much faster)") + crop_results = self.crop_node.inpaint_crop(**args) + stitcher, cropped_image, _ = crop_results[:3] + stitch_results = self.stitch_node.inpaint_stitch(stitcher, cropped_image) + self.assertEqual(stitch_results[0].shape, orig_img.shape) + + def test_invalid_stitcher_validation(self): + # Missing required keys in stitcher dict + corrupted_stitcher = {"canvas_image": [torch.zeros(1, 10, 10, 3)]} + with self.assertRaises(ValueError): + self.stitch_node.inpaint_stitch(corrupted_stitcher, torch.zeros(1, 10, 10, 3)) + + def test_mask_none_default(self): + # No mask provided: node creates default mask + img = torch.rand(1, 40, 40, 3, dtype=torch.float32) + res = self.crop_node.inpaint_crop(**self._default_crop_args(img, mask=None)) + stitcher, crop_im, crop_m = res[:3] + self.assertEqual(crop_im.shape[-1], 3) + self.assertIsNotNone(stitcher) + + def test_mask_mismatched_spatial_resolution(self): + img = torch.rand(1, 64, 64, 3, dtype=torch.float32) + # Empty mismatched mask (32x32) + empty_mask = torch.zeros(1, 32, 32, dtype=torch.float32) + res_empty = self.crop_node.inpaint_crop(**self._default_crop_args(img, mask=empty_mask)) + self.assertIsNotNone(res_empty[0]) + + # Non-empty mismatched mask (32x32 with content) + content_mask = torch.zeros(1, 32, 32, dtype=torch.float32) + content_mask[0, 10:20, 10:20] = 1.0 + res_content = self.crop_node.inpaint_crop(**self._default_crop_args(img, mask=content_mask)) + self.assertIsNotNone(res_content[0]) + + def test_output_resize_to_target_size(self): + img = torch.rand(1, 100, 100, 3, dtype=torch.float32) + mask = torch.zeros(1, 100, 100, dtype=torch.float32) + mask[0, 20:40, 20:40] = 1.0 + + args = self._default_crop_args( + img, mask=mask, + output_resize_to_target_size=True, + output_target_width=128, + output_target_height=128 + ) + res = self.crop_node.inpaint_crop(**args) + stitcher, crop_im, crop_m = res[:3] + self.assertEqual(crop_im.shape[1], 128) + self.assertEqual(crop_im.shape[2], 128) + + def test_batch_processing(self): + img_batch = torch.rand(2, 64, 64, 3, dtype=torch.float32) + mask_batch = torch.zeros(2, 64, 64, dtype=torch.float32) + mask_batch[0, 10:30, 10:30] = 1.0 + mask_batch[1, 20:40, 20:40] = 1.0 + + args = self._default_crop_args( + img_batch, mask=mask_batch, + output_resize_to_target_size=True, + output_target_width=64, + output_target_height=64 + ) + res = self.crop_node.inpaint_crop(**args) + stitcher, crop_im, crop_m = res[:3] + self.assertEqual(crop_im.shape[0], 2) + + # Stitch + def test_inpaint_crop_grayscale_1channel(self): + # Grayscale 1-channel image input + img_gray = torch.rand(1, 40, 40, 1, dtype=torch.float32) + mask = torch.zeros(1, 40, 40, dtype=torch.float32) + mask[0, 10:20, 10:20] = 1.0 + res = self.crop_node.inpaint_crop(**self._default_crop_args(img_gray, mask=mask)) + stitcher, crop_im, crop_m = res[:3] + # Auto-converted to 3 channels RGB + self.assertEqual(crop_im.shape[-1], 3) + + def test_inpaint_crop_2channel_grayscale_alpha(self): + # 2-channel (grayscale + alpha) + img_ga = torch.rand(1, 40, 40, 2, dtype=torch.float32) + mask = torch.zeros(1, 40, 40, dtype=torch.float32) + mask[0, 10:20, 10:20] = 1.0 + res = self.crop_node.inpaint_crop(**self._default_crop_args(img_ga, mask=mask)) + stitcher, crop_im, crop_m = res[:3] + # Auto-converted to 4 channels RGBA + self.assertEqual(crop_im.shape[-1], 4) + + def test_mask_range_255_normalization(self): + # Mask with 0-255 values + img = torch.rand(1, 50, 50, 3, dtype=torch.float32) + mask_255 = torch.zeros(1, 50, 50, dtype=torch.float32) + mask_255[0, 15:35, 15:35] = 255.0 + res = self.crop_node.inpaint_crop(**self._default_crop_args(img, mask=mask_255)) + stitcher, crop_im, crop_m = res[:3] + # Mask is clamped and normalized into [0.0, 1.0] + self.assertLessEqual(crop_m.max().item(), 1.0) + + def test_preresize_min_greater_than_max_swap(self): + # Inverted min/max resolutions (min=800, max=400) + img = torch.rand(1, 60, 60, 3, dtype=torch.float32) + mask = torch.zeros(1, 60, 60, dtype=torch.float32) + args = self._default_crop_args( + img, mask=mask, + preresize=True, + preresize_mode="ensure minimum and maximum resolution", + preresize_min_width=800, + preresize_max_width=400, + preresize_min_height=800, + preresize_max_height=400 + ) + # Should gracefully swap min and max instead of crashing with AssertionError + res = self.crop_node.inpaint_crop(**args) + self.assertIsNotNone(res[0]) + + def test_batch_without_target_resize_auto_handled(self): + # Batch of 2 images with output_resize_to_target_size=False + img_batch = torch.rand(2, 60, 60, 3, dtype=torch.float32) + mask_batch = torch.zeros(2, 60, 60, dtype=torch.float32) + mask_batch[0, 10:20, 10:20] = 1.0 + mask_batch[1, 20:30, 20:30] = 1.0 + args = self._default_crop_args(img_batch, mask=mask_batch, output_resize_to_target_size=False) + # Should auto-enable target resize and process batch cleanly + res = self.crop_node.inpaint_crop(**args) + self.assertEqual(res[1].shape[0], 2) + + def test_flexible_batch_ratio_inpaint_stitch(self): + # 1 stitcher with 4 inpainted candidate images + img_1 = torch.rand(1, 60, 60, 3, dtype=torch.float32) + mask_1 = torch.zeros(1, 60, 60, dtype=torch.float32) + mask_1[0, 15:35, 15:35] = 1.0 + crop_res = self.crop_node.inpaint_crop(**self._default_crop_args(img_1, mask=mask_1)) + stitcher, crop_im, _ = crop_res[:3] + + # 4 variations generated from sampler + variations_4 = crop_im.repeat(4, 1, 1, 1) + stitch_res = self.stitch_node.inpaint_stitch(stitcher, variations_4) + self.assertEqual(stitch_res[0].shape, (4, 60, 60, 3)) + + def test_bool_and_uint8_inputs(self): + # uint8 image (0-255) and boolean mask + img_u8 = torch.randint(0, 256, (1, 60, 60, 3), dtype=torch.uint8) + mask_bool = torch.zeros(1, 60, 60, dtype=torch.bool) + mask_bool[0, 15:35, 15:35] = True + + for device_mode in ["cpu (compatible)", "gpu (much faster)"]: + with self.subTest(device_mode=device_mode): + args = self._default_crop_args(img_u8, mask=mask_bool, device_mode=device_mode) + res = self.crop_node.inpaint_crop(**args) + stitcher, crop_im, crop_m = res[:3] + self.assertEqual(crop_im.dtype, torch.float32) + self.assertEqual(crop_m.dtype, torch.float32) + + # uint8 inpainted image + inpaint_u8 = (crop_im * 255).to(torch.uint8) + stitch_res = self.stitch_node.inpaint_stitch(stitcher, inpaint_u8) + self.assertEqual(stitch_res[0].shape, (1, 60, 60, 3)) + self.assertEqual(stitch_res[0].dtype, torch.float32) + + def test_empty_tensor_masks(self): + img = torch.rand(1, 40, 40, 3, dtype=torch.float32) + args = self._default_crop_args(img, mask=torch.empty(0), optional_context_mask=torch.empty(0)) + res = self.crop_node.inpaint_crop(**args) + stitcher, crop_im, crop_m = res[:3] + self.assertIsNotNone(stitcher) + self.assertEqual(crop_im.shape[-1], 3) + + +if __name__ == '__main__': + unittest.main() diff --git a/tests/test_processor_cpu.py b/tests/test_processor_cpu.py new file mode 100644 index 0000000..3a62b15 --- /dev/null +++ b/tests/test_processor_cpu.py @@ -0,0 +1,234 @@ +import unittest +import torch +import numpy as np +from inpaint_cropandstitch import CPUProcessorLogic + + +class TestCPUProcessorLogic(unittest.TestCase): + def setUp(self): + self.processor = CPUProcessorLogic() + + def test_rescale_i_algorithms(self): + img = torch.rand(1, 32, 32, 3, dtype=torch.float32) + algorithms = ["bicubic", "bilinear", "nearest", "nearest-exact", "lanczos", "box", "area", "hamming"] + for algo in algorithms: + with self.subTest(algorithm=algo): + res = self.processor.rescale_i(img, 64, 48, algo) + self.assertEqual(res.shape, (1, 48, 64, 3)) + self.assertEqual(res.dtype, torch.float32) + + def test_rescale_i_channels_and_batches(self): + for c in [1, 3, 4]: + with self.subTest(channels=c): + img = torch.rand(2, 20, 20, c) + res = self.processor.rescale_i(img, 40, 30, "bilinear") + self.assertEqual(res.shape, (2, 30, 40, c)) + + def test_rescale_i_dtypes(self): + for dtype in [torch.float32, torch.float16, torch.bfloat16]: + with self.subTest(dtype=dtype): + img = torch.rand(1, 16, 16, 3, dtype=dtype) + res = self.processor.rescale_i(img, 32, 32, "bicubic") + self.assertEqual(res.shape, (1, 32, 32, 3)) + self.assertIn(res.dtype, [torch.float32, dtype]) + + def test_rescale_i_boundary_sizes(self): + img = torch.rand(1, 10, 10, 3) + res_1x1 = self.processor.rescale_i(img, 1, 1, "bilinear") + self.assertEqual(res_1x1.shape, (1, 1, 1, 3)) + + res_zero = self.processor.rescale_i(img, 0, -5, "bilinear") + self.assertEqual(res_zero.shape, (1, 1, 1, 3)) + + def test_rescale_m_algorithms_and_shapes(self): + mask_3d = torch.rand(1, 24, 24, dtype=torch.float32) + mask_2d = torch.rand(24, 24, dtype=torch.float32) + algorithms = ["bicubic", "bilinear", "nearest", "nearest-exact", "lanczos", "box", "area", "hamming"] + for algo in algorithms: + with self.subTest(algorithm=algo): + res_3d = self.processor.rescale_m(mask_3d, 48, 36, algo) + self.assertEqual(res_3d.shape, (1, 36, 48)) + res_2d = self.processor.rescale_m(mask_2d, 48, 36, algo) + self.assertEqual(res_2d.shape, (1, 36, 48)) + + def test_rescale_m_dtypes(self): + for dtype in [torch.float32, torch.float16, torch.bfloat16]: + with self.subTest(dtype=dtype): + mask = torch.rand(1, 16, 16, dtype=dtype) + res = self.processor.rescale_m(mask, 32, 32, "bilinear") + self.assertEqual(res.shape, (1, 32, 32)) + + def test_fillholes_functional_hollow_ring(self): + # Frame enclosing empty hole in center + mask = torch.zeros(1, 40, 40, dtype=torch.float32) + mask[:, 5:35, 5:10] = 1.0 + mask[:, 5:35, 30:35] = 1.0 + mask[:, 5:10, 5:35] = 1.0 + mask[:, 30:35, 5:35] = 1.0 + + # Center is originally zero + self.assertEqual(mask[0, 20, 20].item(), 0.0) + filled = self.processor.fillholes_iterative_hipass_fill_m(mask) + # Inside center must now be filled to 1.0 + self.assertAlmostEqual(filled[0, 20, 20].item(), 1.0) + # Outside corner remains 0.0 + self.assertEqual(filled[0, 0, 0].item(), 0.0) + + def test_fillholes_multilevel_gradient(self): + # Soft threshold ring (0.8) with 0.0 hole + mask = torch.zeros(1, 30, 30, dtype=torch.float32) + mask[:, 5:25, 5:25] = 0.8 + mask[:, 10:20, 10:20] = 0.0 + filled = self.processor.fillholes_iterative_hipass_fill_m(mask) + self.assertAlmostEqual(filled[0, 15, 15].item(), 0.8, places=2) + self.assertEqual(filled[0, 0, 0].item(), 0.0) + + def test_fillholes_2d_mask(self): + mask_2d = torch.zeros(40, 40, dtype=torch.float32) + mask_2d[5:35, 5:10] = 1.0 + mask_2d[5:35, 30:35] = 1.0 + mask_2d[5:10, 5:35] = 1.0 + mask_2d[30:35, 5:35] = 1.0 + filled_2d = self.processor.fillholes_iterative_hipass_fill_m(mask_2d) + self.assertEqual(filled_2d.ndim, 2) + self.assertAlmostEqual(filled_2d[20, 20].item(), 1.0) + self.assertEqual(filled_2d[0, 0].item(), 0.0) + + def test_fillholes_solid_and_empty(self): + # Solid mask should remain solid + solid = torch.ones(1, 20, 20, dtype=torch.float32) + self.assertTrue(torch.equal(self.processor.fillholes_iterative_hipass_fill_m(solid), solid)) + + # Empty mask should remain empty + empty = torch.zeros(1, 20, 20, dtype=torch.float32) + self.assertTrue(torch.equal(self.processor.fillholes_iterative_hipass_fill_m(empty), empty)) + + def test_hipassfilter_m(self): + mask = torch.tensor([[[0.1, 0.4], [0.6, 0.9]]], dtype=torch.float32) + filtered = self.processor.hipassfilter_m(mask, 0.5) + self.assertEqual(filtered[0, 0, 0].item(), 0.0) + self.assertEqual(filtered[0, 0, 1].item(), 0.0) + self.assertAlmostEqual(filtered[0, 1, 0].item(), 0.6) + self.assertAlmostEqual(filtered[0, 1, 1].item(), 0.9) + + filtered_neg = self.processor.hipassfilter_m(mask, -1.0) + self.assertTrue(torch.equal(filtered_neg, mask)) + + filtered_high = self.processor.hipassfilter_m(mask, 1.5) + self.assertEqual(torch.count_nonzero(filtered_high).item(), 0) + + def test_expand_m_exact_radii(self): + mask = torch.zeros(1, 31, 31, dtype=torch.float32) + mask[0, 15, 15] = 1.0 + exp = self.processor.expand_m(mask, 4) + self.assertEqual(exp[0, 15, 15].item(), 1.0) + self.assertEqual(exp[0, 15, 16].item(), 1.0) + self.assertEqual(exp[0, 15, 14].item(), 1.0) + self.assertEqual(exp[0, 0, 0].item(), 0.0) + + def test_expand_m_boundary_cases(self): + mask = torch.zeros(1, 20, 20, dtype=torch.float32) + mask[0, 10, 10] = 1.0 + self.assertTrue(torch.equal(self.processor.expand_m(mask, 0), mask)) + self.assertTrue(torch.equal(self.processor.expand_m(mask, -5), mask)) + + # 2D mask + mask_2d = torch.zeros(20, 20, dtype=torch.float32) + mask_2d[10, 10] = 1.0 + exp_2d = self.processor.expand_m(mask_2d, 4) + self.assertEqual(exp_2d.ndim, 2) + + # Huge expansion + exp_large = self.processor.expand_m(mask, 100) + self.assertEqual(exp_large.shape, mask.shape) + + def test_invert_m(self): + mask = torch.tensor([[[0.0, 0.25], [0.75, 1.0]]], dtype=torch.float32) + inv = self.processor.invert_m(mask) + self.assertAlmostEqual(inv[0, 0, 0].item(), 1.0) + self.assertAlmostEqual(inv[0, 0, 1].item(), 0.75) + self.assertAlmostEqual(inv[0, 1, 0].item(), 0.25) + self.assertAlmostEqual(inv[0, 1, 1].item(), 0.0) + + # Bool mask + mask_bool = torch.tensor([[[False, True], [True, False]]], dtype=torch.bool) + inv_bool = self.processor.invert_m(mask_bool) + self.assertAlmostEqual(inv_bool[0, 0, 0].item(), 1.0) + self.assertAlmostEqual(inv_bool[0, 0, 1].item(), 0.0) + + def test_blur_m_bool_and_int_dtypes(self): + mask_bool = torch.zeros(1, 21, 21, dtype=torch.bool) + mask_bool[0, 10, 10] = True + blurred_bool = self.processor.blur_m(mask_bool, 3) + self.assertGreater(blurred_bool[0, 10, 10].item(), 0.0) + + mask_u8 = torch.zeros(1, 21, 21, dtype=torch.uint8) + mask_u8[0, 10, 10] = 255 + blurred_u8 = self.processor.blur_m(mask_u8, 3) + self.assertGreater(blurred_u8[0, 10, 10].item(), 0.0) + + def test_blur_m_gaussian_decay(self): + mask = torch.zeros(1, 31, 31, dtype=torch.float32) + mask[0, 15, 15] = 1.0 + blurred = self.processor.blur_m(mask, 4) + + peak = blurred[0, 15, 15].item() + d1 = blurred[0, 15, 16].item() + d2 = blurred[0, 15, 17].item() + self.assertGreater(peak, d1) + self.assertGreater(d1, d2) + self.assertAlmostEqual(blurred[0, 0, 0].item(), 0.0, places=3) + + def test_pad_to_multiple(self): + for val, mult, expected in [ + (0, 8, 0), (7, 8, 8), (8, 8, 8), (9, 8, 16), + (60, 8, 64), (64, 8, 64), (65, 8, 72), + (64, 16, 64), (65, 16, 80), (100, 32, 128), + (50, 0, 50), (50, -8, 50) + ]: + with self.subTest(val=val, mult=mult): + self.assertEqual(self.processor.pad_to_multiple(val, mult), expected) + + def test_debug_context_location_in_image_inversion(self): + img = torch.full((1, 30, 30, 3), 0.25, dtype=torch.float32) + deb = self.processor.debug_context_location_in_image(img, 10, 10, 10, 10) + # Inside box: 1.0 - 0.25 = 0.75 + self.assertAlmostEqual(deb[0, 15, 15, 0].item(), 0.75) + # Outside box: remains 0.25 + self.assertAlmostEqual(deb[0, 0, 0, 0].item(), 0.25) + + # Clamping out-of-bounds coordinates + deb_out = self.processor.debug_context_location_in_image(img, -10, -10, 20, 20) + self.assertEqual(deb_out.shape, img.shape) + deb_far = self.processor.debug_context_location_in_image(img, 100, 100, 20, 20) + self.assertEqual(deb_far.shape, img.shape) + + def test_extend_imm_edge_preservation(self): + img = torch.zeros(1, 10, 10, 3, dtype=torch.float32) + img[:, 0, :, 0] = 0.42 + img[:, -1, :, 1] = 0.77 + mask = torch.zeros(1, 10, 10, dtype=torch.float32) + e_img, e_mask, _ = self.processor.extend_imm(img, mask, None, 2.0, 2.0, 1.0, 1.0) + self.assertAlmostEqual(e_img[0, 0, 5, 0].item(), 0.42) + self.assertAlmostEqual(e_img[0, -1, 5, 1].item(), 0.77) + self.assertAlmostEqual(e_mask[0, 0, 5].item(), 1.0) + self.assertAlmostEqual(e_mask[0, -1, 5].item(), 1.0) + + def test_preresize_imm_modes(self): + img = torch.rand(1, 100, 100, 3, dtype=torch.float32) + mask = torch.zeros(1, 100, 100, dtype=torch.float32) + opt_mask = torch.zeros(1, 100, 100, dtype=torch.float32) + + # Ensure min resolution + p_img, p_mask, p_opt = self.processor.preresize_imm(img, mask, opt_mask, "bilinear", "bicubic", "ensure minimum resolution", 200, 150, 400, 400) + self.assertGreaterEqual(p_img.shape[2], 200) + self.assertGreaterEqual(p_img.shape[1], 150) + + # Ensure max resolution + p_img2, p_mask2, p_opt2 = self.processor.preresize_imm(img, mask, opt_mask, "bilinear", "bicubic", "ensure maximum resolution", 10, 10, 50, 80) + self.assertLessEqual(p_img2.shape[2], 50) + self.assertLessEqual(p_img2.shape[1], 80) + + +if __name__ == '__main__': + unittest.main() diff --git a/tests/test_processor_gpu.py b/tests/test_processor_gpu.py new file mode 100644 index 0000000..c938098 --- /dev/null +++ b/tests/test_processor_gpu.py @@ -0,0 +1,217 @@ +import unittest +import torch +from inpaint_cropandstitch import GPUProcessorLogic + + +class TestGPUProcessorLogic(unittest.TestCase): + @classmethod + def setUpClass(cls): + cls.devices = ["cpu"] + if torch.cuda.is_available(): + cls.devices.append("cuda") + + def setUp(self): + self.processor = GPUProcessorLogic() + + def test_rescale_i_algorithms(self): + for dev in self.devices: + img = torch.rand(1, 32, 32, 3, dtype=torch.float32, device=dev) + algorithms = ["bicubic", "bilinear", "nearest", "nearest-exact", "lanczos", "box", "area", "hamming"] + for algo in algorithms: + with self.subTest(device=dev, algorithm=algo): + res = self.processor.rescale_i(img, 64, 48, algo) + self.assertEqual(res.shape, (1, 48, 64, 3)) + self.assertEqual(res.device.type, dev) + + def test_rescale_i_channels_and_dtypes(self): + for dev in self.devices: + for c in [1, 3, 4]: + with self.subTest(device=dev, channels=c): + img = torch.rand(2, 20, 20, c, device=dev) + res = self.processor.rescale_i(img, 30, 30, "bilinear") + self.assertEqual(res.shape, (2, 30, 30, c)) + + for dtype in [torch.float32, torch.float16, torch.bfloat16]: + with self.subTest(device=dev, dtype=dtype): + img = torch.rand(1, 16, 16, 3, dtype=dtype, device=dev) + res = self.processor.rescale_i(img, 32, 32, "bicubic") + self.assertEqual(res.shape, (1, 32, 32, 3)) + + def test_rescale_m_algorithms_and_shapes(self): + for dev in self.devices: + mask_3d = torch.rand(1, 24, 24, dtype=torch.float32, device=dev) + mask_2d = torch.rand(24, 24, dtype=torch.float32, device=dev) + for algo in ["bicubic", "bilinear", "nearest-exact", "box"]: + with self.subTest(device=dev, algorithm=algo): + res_3d = self.processor.rescale_m(mask_3d, 40, 40, algo) + self.assertEqual(res_3d.shape, (1, 40, 40)) + self.assertEqual(res_3d.device.type, dev) + res_2d = self.processor.rescale_m(mask_2d, 40, 40, algo) + self.assertEqual(res_2d.shape, (1, 40, 40)) + + def test_fillholes_functional_hollow_ring(self): + for dev in self.devices: + mask = torch.zeros(1, 40, 40, dtype=torch.float32, device=dev) + mask[:, 5:35, 5:10] = 1.0 + mask[:, 5:35, 30:35] = 1.0 + mask[:, 5:10, 5:35] = 1.0 + mask[:, 30:35, 5:35] = 1.0 + + self.assertEqual(mask[0, 20, 20].item(), 0.0) + filled = self.processor.fillholes_iterative_hipass_fill_m(mask) + self.assertEqual(filled.device.type, dev) + self.assertAlmostEqual(filled[0, 20, 20].item(), 1.0) + self.assertEqual(filled[0, 0, 0].item(), 0.0) + + def test_fillholes_2d_mask(self): + for dev in self.devices: + mask_2d = torch.zeros(40, 40, dtype=torch.float32, device=dev) + mask_2d[5:35, 5:10] = 1.0 + mask_2d[5:35, 30:35] = 1.0 + mask_2d[5:10, 5:35] = 1.0 + mask_2d[30:35, 5:35] = 1.0 + filled_2d = self.processor.fillholes_iterative_hipass_fill_m(mask_2d) + self.assertEqual(filled_2d.ndim, 2) + self.assertAlmostEqual(filled_2d[20, 20].item(), 1.0) + self.assertEqual(filled_2d[0, 0].item(), 0.0) + + def test_fillholes_bfloat16(self): + for dev in self.devices: + mask = torch.zeros(1, 20, 20, dtype=torch.bfloat16, device=dev) + mask[:, 5:15, 5:15] = 1.0 + mask[:, 8:12, 8:12] = 0.0 + filled = self.processor.fillholes_iterative_hipass_fill_m(mask) + self.assertEqual(filled.shape, mask.shape) + + def test_hipassfilter_m(self): + for dev in self.devices: + mask = torch.tensor([[[0.1, 0.4], [0.6, 0.9]]], dtype=torch.float32, device=dev) + filtered = self.processor.hipassfilter_m(mask, 0.5) + self.assertEqual(filtered[0, 0, 0].item(), 0.0) + self.assertAlmostEqual(filtered[0, 1, 1].item(), 0.9) + + def test_expand_m_exact_radii(self): + for dev in self.devices: + mask = torch.zeros(1, 31, 31, dtype=torch.float32, device=dev) + mask[0, 15, 15] = 1.0 + exp = self.processor.expand_m(mask, 4) + self.assertEqual(exp[0, 15, 15].item(), 1.0) + self.assertEqual(exp[0, 15, 16].item(), 1.0) + self.assertEqual(exp[0, 15, 14].item(), 1.0) + self.assertEqual(exp[0, 0, 0].item(), 0.0) + + def test_expand_m_boundary_cases(self): + for dev in self.devices: + mask = torch.zeros(1, 20, 20, dtype=torch.float32, device=dev) + mask[0, 10, 10] = 1.0 + self.assertTrue(torch.equal(self.processor.expand_m(mask, 0), mask)) + self.assertTrue(torch.equal(self.processor.expand_m(mask, -3), mask)) + + mask_2d = torch.zeros(20, 20, dtype=torch.float32, device=dev) + mask_2d[10, 10] = 1.0 + exp_2d = self.processor.expand_m(mask_2d, 4) + self.assertEqual(exp_2d.ndim, 2) + + exp_huge = self.processor.expand_m(mask, 80) + self.assertEqual(exp_huge.shape, mask.shape) + + def test_expand_m_bool_and_int_dtypes(self): + for dev in self.devices: + mask_bool = torch.zeros(1, 21, 21, dtype=torch.bool, device=dev) + mask_bool[0, 10, 10] = True + exp_bool = self.processor.expand_m(mask_bool, 3) + self.assertEqual(exp_bool[0, 10, 10].item(), True) + self.assertEqual(exp_bool[0, 10, 11].item(), True) + self.assertEqual(exp_bool[0, 0, 0].item(), False) + + mask_u8 = torch.zeros(1, 21, 21, dtype=torch.uint8, device=dev) + mask_u8[0, 10, 10] = 255 + exp_u8 = self.processor.expand_m(mask_u8, 3) + self.assertEqual(exp_u8[0, 10, 10].item(), 255) + + def test_invert_m(self): + for dev in self.devices: + mask = torch.tensor([[[0.0, 0.25], [0.75, 1.0]]], dtype=torch.float32, device=dev) + inv = self.processor.invert_m(mask) + self.assertAlmostEqual(inv[0, 0, 0].item(), 1.0) + self.assertAlmostEqual(inv[0, 1, 1].item(), 0.0) + + # Bool mask + mask_bool = torch.tensor([[[False, True], [True, False]]], dtype=torch.bool, device=dev) + inv_bool = self.processor.invert_m(mask_bool) + self.assertAlmostEqual(inv_bool[0, 0, 0].item(), 1.0) + self.assertAlmostEqual(inv_bool[0, 0, 1].item(), 0.0) + + def test_blur_m_gaussian_decay(self): + for dev in self.devices: + mask = torch.zeros(1, 31, 31, dtype=torch.float32, device=dev) + mask[0, 15, 15] = 1.0 + blurred = self.processor.blur_m(mask, 4) + + peak = blurred[0, 15, 15].item() + d1 = blurred[0, 15, 16].item() + d2 = blurred[0, 15, 17].item() + self.assertGreater(peak, d1) + self.assertGreater(d1, d2) + self.assertAlmostEqual(blurred[0, 0, 0].item(), 0.0, places=3) + + def test_blur_m_bool_and_int_dtypes(self): + for dev in self.devices: + mask_bool = torch.zeros(1, 21, 21, dtype=torch.bool, device=dev) + mask_bool[0, 10, 10] = True + blurred_bool = self.processor.blur_m(mask_bool, 3) + self.assertGreater(blurred_bool[0, 10, 10].item(), 0.0) + self.assertAlmostEqual(blurred_bool[0, 0, 0].item(), 0.0, places=3) + + mask_u8 = torch.zeros(1, 21, 21, dtype=torch.uint8, device=dev) + mask_u8[0, 10, 10] = 1 + blurred_u8 = self.processor.blur_m(mask_u8, 3) + self.assertGreater(blurred_u8[0, 10, 10].item(), 0.0) + + def test_pad_to_multiple(self): + self.assertEqual(self.processor.pad_to_multiple(60, 8), 64) + self.assertEqual(self.processor.pad_to_multiple(50, 0), 50) + self.assertEqual(self.processor.pad_to_multiple(50, -4), 50) + + def test_debug_context_location_in_image(self): + for dev in self.devices: + img = torch.full((1, 30, 30, 3), 0.25, dtype=torch.float32, device=dev) + deb = self.processor.debug_context_location_in_image(img, 10, 10, 10, 10) + self.assertAlmostEqual(deb[0, 15, 15, 0].item(), 0.75) + self.assertAlmostEqual(deb[0, 0, 0, 0].item(), 0.25) + + def test_extend_imm_edge_preservation(self): + for dev in self.devices: + img = torch.zeros(1, 10, 10, 3, dtype=torch.float32, device=dev) + img[:, 0, :, 0] = 0.55 + img[:, -1, :, 2] = 0.88 + mask = torch.zeros(1, 10, 10, dtype=torch.float32, device=dev) + e_img, e_mask, _ = self.processor.extend_imm(img, mask, None, 2.0, 2.0, 1.0, 1.0) + self.assertAlmostEqual(e_img[0, 0, 5, 0].item(), 0.55) + self.assertAlmostEqual(e_img[0, -1, 5, 2].item(), 0.88) + self.assertAlmostEqual(e_mask[0, 0, 5].item(), 1.0) + + def test_preresize_imm_modes(self): + for dev in self.devices: + img = torch.rand(1, 50, 60, 3, dtype=torch.float32, device=dev) + mask = torch.rand(1, 50, 60, dtype=torch.float32, device=dev) + + # ensure minimum + r_img, r_mask, _ = self.processor.preresize_imm( + img, mask, mask, "bilinear", "bicubic", "ensure minimum resolution", + 100, 100, 200, 200 + ) + self.assertGreaterEqual(r_img.shape[2], 100) + self.assertGreaterEqual(r_img.shape[1], 100) + + # ensure maximum + r_img2, r_mask2, _ = self.processor.preresize_imm( + img, mask, mask, "bilinear", "bicubic", "ensure maximum resolution", + 20, 20, 30, 30 + ) + self.assertLessEqual(r_img2.shape[2], 30) + self.assertLessEqual(r_img2.shape[1], 30) + + +if __name__ == '__main__': + unittest.main() diff --git a/tests/test_workflows.py b/tests/test_workflows.py new file mode 100644 index 0000000..1ed943b --- /dev/null +++ b/tests/test_workflows.py @@ -0,0 +1,184 @@ +import os +import unittest +import torch +from tests.workflow_runner import ( + WorkflowRunner, + validate_workflow_run, + validate_crop_outputs, + validate_stitch_outputs, + MockLoadImage, + MockMaskToImage, + MockImageInvert, + MockImpactMakeImageBatch, + MockImpactMakeMaskBatch, + MockImageCompositeMasked, +) + + +class TestMockNodes(unittest.TestCase): + """Unit tests for the self-contained mock nodes used in workflow execution.""" + + @classmethod + def setUpClass(cls): + cls.repo_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + cls.testimgs_dir = os.path.join(cls.repo_root, "testimgs") + + def test_mock_load_image_rgb(self): + loader = MockLoadImage(self.testimgs_dir) + img, mask = loader.load_image("example.png") + self.assertEqual(img.ndim, 4) + self.assertEqual(img.shape[0], 1) + self.assertEqual(img.shape[-1], 3) + self.assertEqual(mask.ndim, 3) + self.assertEqual(mask.shape, (1, 64, 64)) # No alpha -> 64x64 zero mask + self.assertEqual(mask.sum().item(), 0.0) + + def test_mock_load_image_clipspace_rgba(self): + loader = MockLoadImage(self.testimgs_dir) + img, mask = loader.load_image("clipspace/clipspace-mask-105444.59999999404.png [input]") + self.assertEqual(img.ndim, 4) + self.assertEqual(img.shape[0], 1) + self.assertEqual(img.shape[-1], 3) + self.assertEqual(mask.ndim, 3) + self.assertEqual(mask.shape[1:], img.shape[1:3]) + self.assertTrue((mask >= 0.0).all() and (mask <= 1.0).all()) + + def test_mock_load_image_not_found(self): + loader = MockLoadImage(self.testimgs_dir) + with self.assertRaises(FileNotFoundError): + loader.load_image("non_existent_image_12345.png") + + def test_mock_mask_to_image(self): + node = MockMaskToImage() + mask_3d = torch.rand(2, 32, 48) + (img_4d,) = node.mask_to_image(mask_3d) + self.assertEqual(img_4d.shape, (2, 32, 48, 3)) + # Channels should be identical grayscale replicated + self.assertTrue(torch.equal(img_4d[..., 0], img_4d[..., 1])) + self.assertTrue(torch.equal(img_4d[..., 1], img_4d[..., 2])) + + mask_2d = torch.rand(32, 48) + (img_from_2d,) = node.mask_to_image(mask_2d) + self.assertEqual(img_from_2d.shape, (1, 32, 48, 3)) + + def test_mock_image_invert(self): + node = MockImageInvert() + img = torch.tensor([[[[0.0, 0.25], [0.75, 1.0]]]]) + (inv,) = node.invert(img) + self.assertTrue(torch.allclose(inv, 1.0 - img)) + + def test_mock_impact_image_batch(self): + node = MockImpactMakeImageBatch() + img1 = torch.rand(1, 16, 16, 3) + img2 = torch.rand(2, 16, 16, 3) + (batched,) = node.make_batch(image1=img1, image2=img2, image3=None) + self.assertEqual(batched.shape, (3, 16, 16, 3)) + + def test_mock_impact_mask_batch(self): + node = MockImpactMakeMaskBatch() + m1 = torch.rand(1, 16, 16) + m2 = torch.rand(16, 16) # 2D mask + (batched,) = node.make_batch(mask1=m1, mask2=m2, mask3=None) + self.assertEqual(batched.shape, (2, 16, 16)) + + def test_mock_image_composite_masked(self): + node = MockImageCompositeMasked() + dest = torch.zeros(1, 40, 40, 3) + src = torch.ones(1, 20, 20, 3) + mask = torch.ones(1, 20, 20) + (res,) = node.composite(dest, src, x=10, y=10, mask=mask) + self.assertEqual(res.shape, (1, 40, 40, 3)) + # Inner region (10:30, 10:30) should be 1.0, outer should be 0.0 + self.assertAlmostEqual(res[0, 15, 15, 0].item(), 1.0) + self.assertAlmostEqual(res[0, 0, 0, 0].item(), 0.0) + + +class TestWorkflowExecution(unittest.TestCase): + """Executes testscpu.json and testsgpu.json end-to-end and validates outputs.""" + + @classmethod + def setUpClass(cls): + cls.repo_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + cls.cpu_path = os.path.join(cls.repo_root, "testscpu.json") + cls.gpu_path = os.path.join(cls.repo_root, "testsgpu.json") + cls.testimgs_dir = os.path.join(cls.repo_root, "testimgs") + cls.runner = WorkflowRunner(testimgs_dir=cls.testimgs_dir, verbose=False) + # Run each workflow once and cache for all assertion methods + cls.cpu_result = cls.runner.run_file(cls.cpu_path) + cls.gpu_result = cls.runner.run_file(cls.gpu_path) + + def test_testscpu_workflow_execution_and_outputs(self): + """Execute all 793 nodes in testscpu.json and validate outputs.""" + result = self.cpu_result + self.assertEqual(len(result.nodes), 793) + self.assertEqual(len(result.outputs), 793) + self.assertEqual(len(result.crop_node_ids), 105) + self.assertEqual(len(result.stitch_node_ids), 34) + self.assertEqual(len(result.preview_node_ids), 349) + self.assertEqual(len(result.load_node_ids), 78) + + # Full validation across all crop, stitch, and preview nodes + validate_workflow_run(result) + + # Specific sample checks: + # Check node 15 (InpaintCropImproved) + crop_15_outs = result.outputs[15] + self.assertEqual(len(crop_15_outs), 26) + stitcher_15 = crop_15_outs[0] + self.assertEqual(stitcher_15["device_mode"], "cpu (compatible)") + self.assertIn("canvas_image", stitcher_15) + cropped_img_15 = crop_15_outs[1] + self.assertEqual(cropped_img_15.ndim, 4) + self.assertEqual(cropped_img_15.shape[-1], 3) + + # Check node 478 (InpaintStitchImproved) + stitch_478_out = result.outputs[478][0] + self.assertEqual(stitch_478_out.ndim, 4) + self.assertEqual(stitch_478_out.shape[-1], 3) + self.assertTrue((stitch_478_out >= 0.0).all() and (stitch_478_out <= 1.0).all()) + + def test_testsgpu_workflow_execution_and_outputs(self): + """Execute all 793 nodes in testsgpu.json and validate outputs.""" + result = self.gpu_result + self.assertEqual(len(result.nodes), 793) + self.assertEqual(len(result.outputs), 793) + self.assertEqual(len(result.crop_node_ids), 105) + self.assertEqual(len(result.stitch_node_ids), 34) + self.assertEqual(len(result.preview_node_ids), 349) + + # Full validation across all crop, stitch, and preview nodes + validate_workflow_run(result) + + # Verify device mode was passed through stitcher + sample_crop_id = result.crop_node_ids[0] + stitcher = result.outputs[sample_crop_id][0] + self.assertEqual(stitcher["device_mode"], "gpu (much faster)") + + def test_cpu_gpu_workflow_consistency(self): + """Compare CPU vs GPU workflows: output shapes and coordinate consistency.""" + cpu_res = self.cpu_result + gpu_res = self.gpu_result + + # Verify all PreviewImage nodes have identical output shapes + for nid in cpu_res.preview_node_ids: + cpu_tensor = cpu_res.outputs[nid][0] + gpu_tensor = gpu_res.outputs[nid][0] + self.assertEqual( + cpu_tensor.shape, + gpu_tensor.shape, + f"Preview node {nid} shape mismatch between CPU and GPU", + ) + + # Verify all InpaintStitchImproved nodes have identical shapes + for nid in cpu_res.stitch_node_ids: + cpu_stitched = cpu_res.outputs[nid][0] + gpu_stitched = gpu_res.outputs[nid][0] + self.assertEqual( + cpu_stitched.shape, + gpu_stitched.shape, + f"Stitch node {nid} shape mismatch between CPU and GPU", + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/workflow_runner.py b/tests/workflow_runner.py new file mode 100644 index 0000000..5e6e774 --- /dev/null +++ b/tests/workflow_runner.py @@ -0,0 +1,465 @@ +import os +import json +import math +import torch +import numpy as np +from PIL import Image, ImageOps +import inpaint_cropandstitch +from inpaint_cropandstitch import InpaintCropImproved, InpaintStitchImproved + + +def repeat_to_batch_size(tensor, batch_size, dim=0): + """Repeat tensor along dimension to match batch_size (matching comfy.utils).""" + if tensor.shape[dim] > batch_size: + return tensor.narrow(dim, 0, batch_size) + elif tensor.shape[dim] < batch_size: + repeats = dim * [1] + [math.ceil(batch_size / tensor.shape[dim])] + [1] * (len(tensor.shape) - 1 - dim) + return tensor.repeat(repeats).narrow(dim, 0, batch_size) + return tensor + + +def image_alpha_fix(destination, source): + """Align alpha channel dimension between destination and source (matching node_helpers).""" + if destination.shape[-1] < source.shape[-1]: + source = source[..., :destination.shape[-1]] + elif destination.shape[-1] > source.shape[-1]: + source = torch.nn.functional.pad(source, (0, 1)) + source[..., -1] = 1.0 + return destination, source + + +def composite_images(destination, source, x, y, mask=None, multiplier=1, resize_source=False): + """ + Self-contained implementation of ComfyUI's composite function for images. + destination and source in [B, C, H, W] format. + """ + source = source.to(destination.device) + if resize_source: + source = torch.nn.functional.interpolate( + source, size=(destination.shape[-2], destination.shape[-1]), mode="bilinear" + ) + + source = repeat_to_batch_size(source, destination.shape[0]) + + x = max(-source.shape[-1] * multiplier, min(x, destination.shape[-1] * multiplier)) + y = max(-source.shape[-2] * multiplier, min(y, destination.shape[-2] * multiplier)) + + left, top = (x // multiplier, y // multiplier) + right, bottom = (left + source.shape[-1], top + source.shape[-2]) + + if mask is None: + mask = torch.ones_like(source) + else: + mask = mask.to(destination.device, copy=True) + mask = torch.nn.functional.interpolate( + mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), + size=(source.shape[-2], source.shape[-1]), + mode="bilinear", + ) + mask = repeat_to_batch_size(mask, source.shape[0]) + + visible_width = destination.shape[-1] - left + min(0, x) + visible_height = destination.shape[-2] - top + min(0, y) + + mask = mask[:, :, :visible_height, :visible_width] + if mask.ndim < source.ndim: + mask = mask.unsqueeze(1) + + inverse_mask = torch.ones_like(mask) - mask + + source_portion = mask * source[..., :visible_height, :visible_width] + destination_portion = inverse_mask * destination[..., top:bottom, left:right] + + destination[..., top:bottom, left:right] = source_portion + destination_portion + return destination + + +class MockLoadImage: + """Self-contained mock for ComfyUI's LoadImage node.""" + + def __init__(self, base_dir=None): + if base_dir is None: + # Default to repo root / testimgs + base_dir = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "testimgs") + self.base_dir = base_dir + + def load_image(self, image_name): + clean_name = image_name.replace(" [input]", "") + candidates = [ + os.path.join(self.base_dir, clean_name), + os.path.join(self.base_dir, os.path.basename(clean_name)), + clean_name, + ] + found_path = None + for c in candidates: + if os.path.exists(c): + found_path = c + break + + if found_path is None: + raise FileNotFoundError(f"MockLoadImage: Image file not found: {image_name}. Tried {candidates}") + + with Image.open(found_path) as img: + img = ImageOps.exif_transpose(img) + rgb = img.convert("RGB") + image_np = np.array(rgb).astype(np.float32) / 255.0 + image_tensor = torch.from_numpy(image_np).unsqueeze(0) # [1, H, W, 3] + + if "A" in img.getbands(): + mask_np = np.array(img.getchannel("A")).astype(np.float32) / 255.0 + mask_tensor = 1.0 - torch.from_numpy(mask_np) + else: + mask_tensor = torch.zeros((64, 64), dtype=torch.float32) + mask_tensor = mask_tensor.unsqueeze(0) # [1, H, W] + + return (image_tensor, mask_tensor) + + +class MockMaskToImage: + """Self-contained mock for ComfyUI's MaskToImage node.""" + + def mask_to_image(self, mask): + if mask.ndim == 2: + mask = mask.unsqueeze(0) + result = mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])).movedim(1, -1).expand(-1, -1, -1, 3) + return (result,) + + +class MockImageInvert: + """Self-contained mock for ComfyUI's ImageInvert node.""" + + def invert(self, image): + return (1.0 - image,) + + +class MockImpactMakeImageBatch: + """Self-contained mock for ImpactMakeImageBatch node.""" + + def make_batch(self, **kwargs): + imgs = [v for k, v in sorted(kwargs.items()) if k.startswith("image") and v is not None] + if not imgs: + raise ValueError("ImpactMakeImageBatch: No images provided to batch.") + # Ensure 4D + imgs_4d = [img.unsqueeze(0) if img.ndim == 3 else img for img in imgs] + return (torch.cat(imgs_4d, dim=0),) + + +class MockImpactMakeMaskBatch: + """Self-contained mock for ImpactMakeMaskBatch node.""" + + def make_batch(self, **kwargs): + ms = [v for k, v in sorted(kwargs.items()) if k.startswith("mask") and v is not None] + if not ms: + raise ValueError("ImpactMakeMaskBatch: No masks provided to batch.") + # Ensure 3D + ms_3d = [m.unsqueeze(0) if m.ndim == 2 else m for m in ms] + return (torch.cat(ms_3d, dim=0),) + + +class MockImageCompositeMasked: + """Self-contained mock for ImageCompositeMasked node.""" + + def composite(self, destination, source, x=0, y=0, resize_source=False, mask=None): + destination, source = image_alpha_fix(destination, source) + dest_ch = destination.clone().movedim(-1, 1) + src_ch = source.movedim(-1, 1) + output = composite_images(dest_ch, src_ch, x, y, mask, multiplier=1, resize_source=resize_source).movedim(1, -1) + return (output,) + + +class WorkflowRunResult: + """Holds all results and statistics of a workflow execution run.""" + + def __init__(self, wf_name, nodes, links, outputs): + self.wf_name = wf_name + self.nodes = nodes + self.links = links + self.outputs = outputs # node_id -> tuple of output values + + @property + def crop_node_ids(self): + return [nid for nid, n in self.nodes.items() if n.get("type") == "InpaintCropImproved"] + + @property + def stitch_node_ids(self): + return [nid for nid, n in self.nodes.items() if n.get("type") == "InpaintStitchImproved"] + + @property + def preview_node_ids(self): + return [nid for nid, n in self.nodes.items() if n.get("type") == "PreviewImage"] + + @property + def load_node_ids(self): + return [nid for nid, n in self.nodes.items() if n.get("type") == "LoadImage"] + + +class WorkflowRunner: + """ + Parses and executes ComfyUI workflows without requiring an external ComfyUI server. + """ + + CROP_WIDGET_NAMES = [ + "downscale_algorithm", + "upscale_algorithm", + "preresize", + "preresize_mode", + "preresize_min_width", + "preresize_min_height", + "preresize_max_width", + "preresize_max_height", + "mask_fill_holes", + "mask_expand_pixels", + "mask_invert", + "mask_blend_pixels", + "mask_hipass_filter", + "extend_for_outpainting", + "extend_up_factor", + "extend_down_factor", + "extend_left_factor", + "extend_right_factor", + "context_from_mask_extend_factor", + "output_resize_to_target_size", + "output_target_width", + "output_target_height", + "output_padding", + "device_mode", + ] + + def __init__(self, testimgs_dir=None, verbose=False): + self.verbose = verbose + self.load_image_node = MockLoadImage(testimgs_dir) + self.mask_to_image_node = MockMaskToImage() + self.image_invert_node = MockImageInvert() + self.image_batch_node = MockImpactMakeImageBatch() + self.mask_batch_node = MockImpactMakeMaskBatch() + self.composite_node = MockImageCompositeMasked() + + # Initialize Crop & Stitch nodes with DEBUG_MODE enabled + self.crop_node = InpaintCropImproved() + self.crop_node.DEBUG_MODE = True + self.crop_node.VERBOSE = verbose + self.crop_node.RETURN_NAMES = InpaintCropImproved.DEBUG_RETURN_NAMES + + self.stitch_node = InpaintStitchImproved() + + def run_file(self, json_path): + with open(json_path, "r", encoding="utf-8") as f: + wf_data = json.load(f) + return self.run_dict(wf_data, name=os.path.basename(json_path)) + + def run_dict(self, wf_data, name="workflow"): + nodes = {n["id"]: n for n in wf_data.get("nodes", [])} + links = {l[0]: l for l in wf_data.get("links", [])} + memo = {} + + def get_node_output(node_id): + if node_id in memo: + return memo[node_id] + + n = nodes[node_id] + ntype = n.get("type") + wv = n.get("widgets_values") or [] + inputs = n.get("inputs") or [] + + # Resolve linked inputs + resolved_inputs = {} + for inp in inputs: + iname = inp.get("name") + lid = inp.get("link") + if lid is not None: + link_info = links[lid] + from_node_id = link_info[1] + from_slot_idx = link_info[2] + from_outs = get_node_output(from_node_id) + resolved_inputs[iname] = from_outs[from_slot_idx] + + # Execute node based on type + if ntype == "LoadImage": + filename = wv[0] if wv else "example.png" + res = self.load_image_node.load_image(filename) + + elif ntype == "ImageInvert": + res = self.image_invert_node.invert(resolved_inputs["image"]) + + elif ntype == "MaskToImage": + res = self.mask_to_image_node.mask_to_image(resolved_inputs["mask"]) + + elif ntype == "ImpactMakeImageBatch": + res = self.image_batch_node.make_batch(**resolved_inputs) + + elif ntype == "ImpactMakeMaskBatch": + res = self.mask_batch_node.make_batch(**resolved_inputs) + + elif ntype == "ImageCompositeMasked": + dest = resolved_inputs["destination"] + src = resolved_inputs["source"] + mask = resolved_inputs.get("mask") + x = wv[0] if len(wv) > 0 else 0 + y = wv[1] if len(wv) > 1 else 0 + resize_source = wv[2] if len(wv) > 2 else False + res = self.composite_node.composite(dest, src, x=x, y=y, resize_source=resize_source, mask=mask) + + elif ntype == "InpaintCropImproved": + kwargs = {} + for wname, wval in zip(self.CROP_WIDGET_NAMES, wv): + kwargs[wname] = wval + kwargs["image"] = resolved_inputs["image"] + kwargs["mask"] = resolved_inputs.get("mask", None) + kwargs["optional_context_mask"] = resolved_inputs.get("optional_context_mask", None) + res = self.crop_node.inpaint_crop(**kwargs) + + elif ntype == "InpaintStitchImproved": + stitcher = resolved_inputs["stitcher"] + inpainted_image = resolved_inputs["inpainted_image"] + res = self.stitch_node.inpaint_stitch(stitcher, inpainted_image) + + elif ntype == "PreviewImage": + res = (resolved_inputs["images"],) + + elif ntype == "Note": + res = () + + else: + raise ValueError(f"WorkflowRunner: Unsupported node type '{ntype}' (id: {node_id})") + + memo[node_id] = res + return res + + for nid in nodes: + get_node_output(nid) + + return WorkflowRunResult(name, nodes, links, memo) + + +def validate_tensor(tensor, name, expected_ndim=None, min_val=0.0, max_val=1.0): + """Assert tensor validity: type, ndim, finite values, and range.""" + assert isinstance(tensor, torch.Tensor), f"{name}: Expected torch.Tensor, got {type(tensor)}" + if expected_ndim is not None: + assert tensor.ndim == expected_ndim, f"{name}: Expected {expected_ndim} dims, got {tensor.ndim} (shape: {tensor.shape})" + assert not torch.isnan(tensor).any(), f"{name}: Tensor contains NaN values." + assert not torch.isinf(tensor).any(), f"{name}: Tensor contains Inf values." + if min_val is not None and tensor.numel() > 0: + actual_min = tensor.min().item() + assert actual_min >= min_val - 1e-3, f"{name}: Value below minimum: {actual_min} < {min_val}" + if max_val is not None and tensor.numel() > 0: + actual_max = tensor.max().item() + assert actual_max <= max_val + 1e-3, f"{name}: Value above maximum: {actual_max} > {max_val}" + + +def validate_crop_outputs(outputs, node_id=None): + """ + Validates outputs of InpaintCropImproved node: + - Slot 0: stitcher dict with all required spatial & canvas metadata + - Slot 1: cropped_image [B, H, W, 3] + - Slot 2: cropped_mask [B, H, W] + - Slots 3..25: all 23 debug tensors + """ + prefix = f"Crop node {node_id}" if node_id is not None else "Crop node" + assert len(outputs) == 26, f"{prefix}: Expected 26 outputs, got {len(outputs)}" + + stitcher = outputs[0] + cropped_image = outputs[1] + cropped_mask = outputs[2] + + # 1. Validate stitcher dict + assert isinstance(stitcher, dict), f"{prefix}: Stitcher output is not a dict" + required_keys = [ + "cropped_to_canvas_x", + "cropped_to_canvas_y", + "cropped_to_canvas_w", + "cropped_to_canvas_h", + "canvas_image", + "cropped_mask_for_blend", + "canvas_to_orig_x", + "canvas_to_orig_y", + "canvas_to_orig_w", + "canvas_to_orig_h", + ] + for k in required_keys: + assert k in stitcher, f"{prefix}: Stitcher missing key '{k}'" + + # 2. Validate cropped_image & cropped_mask + validate_tensor(cropped_image, f"{prefix} cropped_image", expected_ndim=4, min_val=0.0, max_val=1.0) + validate_tensor(cropped_mask, f"{prefix} cropped_mask", expected_ndim=3, min_val=0.0, max_val=1.0) + assert cropped_image.shape[0] == cropped_mask.shape[0], f"{prefix}: Batch mismatch between image and mask" + assert cropped_image.shape[1:3] == cropped_mask.shape[1:3], f"{prefix}: Spatial mismatch between image and mask" + + # 3. Validate debug tensors + for i in range(3, 26): + out_tensor = outputs[i] + validate_tensor(out_tensor, f"{prefix} debug slot {i}", expected_ndim=None, min_val=0.0, max_val=1.0) + + +def validate_stitch_outputs(stitched_output, stitcher, inpainted_image, node_id=None): + """ + Validates outputs of InpaintStitchImproved node: + - Output is [B, H, W, C] + - Shape matches original canvas + - Values are within [0.0, 1.0], no NaNs or Infs + - Critical invariant: Outside of the modified/masked area, canvas pixels are preserved! + """ + prefix = f"Stitch node {node_id}" if node_id is not None else "Stitch node" + assert isinstance(stitched_output, tuple) and len(stitched_output) >= 1, f"{prefix}: Invalid return format" + stitched = stitched_output[0] + + validate_tensor(stitched, f"{prefix} stitched image", expected_ndim=4, min_val=0.0, max_val=1.0) + + # Validate that output shape matches canvas shape + canvas_imgs = stitcher["canvas_image"] + orig_h = stitcher["canvas_to_orig_h"][0] + orig_w = stitcher["canvas_to_orig_w"][0] + assert stitched.shape[1] == orig_h, f"{prefix}: Height mismatch {stitched.shape[1]} vs expected {orig_h}" + assert stitched.shape[2] == orig_w, f"{prefix}: Width mismatch {stitched.shape[2]} vs expected {orig_w}" + + # Verify unmasked area preservation: + # Where blend mask is 0 (outside inpainted region), stitched image should exactly equal original canvas + for i in range(min(stitched.shape[0], len(canvas_imgs))): + c_img = canvas_imgs[i] + if torch.is_tensor(c_img): + c_img = c_img.cpu() + if c_img.ndim == 3: + c_img = c_img.unsqueeze(0) + # Check corner pixel (0, 0) if crop box doesn't touch (0, 0) + top_x = stitcher["cropped_to_canvas_x"][i] + top_y = stitcher["cropped_to_canvas_y"][i] + if top_x > 0 and top_y > 0: + s_pixel = stitched[i, 0, 0, :3] + c_pixel = c_img[0, 0, 0, :3] + assert torch.allclose(s_pixel, c_pixel, atol=1e-3), ( + f"{prefix}: Unmasked pixel changed at (0, 0): {s_pixel} vs {c_pixel}" + ) + + +def validate_workflow_run(result): + """ + Performs comprehensive verification on all nodes of a completed workflow execution: + - Checks that all nodes were evaluated + - Validates all 105 InpaintCropImproved nodes + - Validates all 34 InpaintStitchImproved nodes + - Validates all 349 PreviewImage sink nodes + """ + assert len(result.outputs) == len(result.nodes), ( + f"Workflow {result.wf_name}: executed {len(result.outputs)} of {len(result.nodes)} nodes." + ) + + # 1. Validate Crop nodes + for nid in result.crop_node_ids: + outs = result.outputs[nid] + validate_crop_outputs(outs, node_id=nid) + + # 2. Validate Stitch nodes + for nid in result.stitch_node_ids: + stitch_node = result.nodes[nid] + stitcher_link_id = next(inp["link"] for inp in stitch_node["inputs"] if inp["name"] == "stitcher") + inpaint_link_id = next(inp["link"] for inp in stitch_node["inputs"] if inp["name"] == "inpainted_image") + s_link = result.links[stitcher_link_id] + i_link = result.links[inpaint_link_id] + stitcher = result.outputs[s_link[1]][s_link[2]] + inpainted = result.outputs[i_link[1]][i_link[2]] + validate_stitch_outputs(result.outputs[nid], stitcher, inpainted, node_id=nid) + + # 3. Validate Preview nodes + for nid in result.preview_node_ids: + outs = result.outputs[nid] + assert len(outs) == 1, f"Preview node {nid}: expected 1 output" + validate_tensor(outs[0], f"Preview node {nid}", expected_ndim=4, min_val=0.0, max_val=1.0)