Added tests and defensive conditions, accept more inputs without breaking

This commit is contained in:
Luis Quesada
2026-09-07 18:35:48 +02:00
parent 606e2b4fd8
commit cd953978e6
12 changed files with 2318 additions and 149 deletions
+70 -2
View File
@@ -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/
+391 -146
View File
@@ -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))
+1 -1
View File
@@ -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]
Executable
+30
View File
@@ -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
+1
View File
@@ -0,0 +1 @@
# Tests package
+127
View File
@@ -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()
+232
View File
@@ -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()
+366
View File
@@ -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()
+234
View File
@@ -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()
+217
View File
@@ -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()
+184
View File
@@ -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()
+465
View File
@@ -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)