Make fill mask holes in GPU progressive, actually fall it back to CPU behavior
This commit is contained in:
+16
-50
@@ -638,57 +638,23 @@ class GPUProcessorLogic(ProcessorLogic):
|
||||
#return samples
|
||||
|
||||
def fillholes_iterative_hipass_fill_m(self, samples):
|
||||
# samples shape: [B, H, W]
|
||||
B, H, W = samples.shape
|
||||
device = samples.device
|
||||
|
||||
# We find areas connected to the border in the inverted mask.
|
||||
# These are "outside" areas. Everything else is either mask or a hole.
|
||||
|
||||
# Invert: 1 where it's 0 (potential hole/outside), 0 where it's 1 (mask/blocker)
|
||||
inv_mask = 1.0 - (samples > 0.5).float()
|
||||
|
||||
# Pad to have a border for flood fill
|
||||
padded_inv = torch.zeros((B, H+2, W+2), device=device)
|
||||
padded_inv[:, 1:-1, 1:-1] = inv_mask
|
||||
|
||||
# Initial seeds: the padding border
|
||||
outside = torch.zeros((B, H+2, W+2), device=device)
|
||||
outside[:, 0, :] = 1
|
||||
outside[:, -1, :] = 1
|
||||
outside[:, :, 0] = 1
|
||||
outside[:, :, -1] = 1
|
||||
|
||||
# Propagate 'outside' status through inv_mask
|
||||
# We use a power-of-two growth for efficiency? No, simple iterative for now.
|
||||
# But wait, max(H, W) is too many iterations.
|
||||
# A better way is to use a large kernel or repeated doublings.
|
||||
|
||||
# Actually, for 512-1024px, 512 iterations of MaxPool are still quite fast compared to CPU.
|
||||
# But we can speed it up by using larger strides/kernels if we don't care about the exact shape?
|
||||
# No, we need exact.
|
||||
|
||||
curr_outside = outside
|
||||
for _ in range(max(H, W)):
|
||||
# Dilation
|
||||
next_outside = TF.max_pool2d(curr_outside.unsqueeze(1), kernel_size=3, stride=1, padding=1).squeeze(1)
|
||||
# Mask by inv_mask (only propagate to 0-areas)
|
||||
next_outside = next_outside * padded_inv
|
||||
# Also keep original border
|
||||
next_outside[:, 0, :] = 1
|
||||
next_outside[:, -1, :] = 1
|
||||
next_outside[:, :, 0] = 1
|
||||
next_outside[:, :, -1] = 1
|
||||
# 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.
|
||||
|
||||
if torch.all(next_outside == curr_outside):
|
||||
break
|
||||
curr_outside = next_outside
|
||||
|
||||
# Final mask: anything not 'outside'
|
||||
filled = 1.0 - curr_outside[:, 1:-1, 1:-1]
|
||||
|
||||
# Combine with original (ensure original mask pixels are kept)
|
||||
return torch.max(samples, filled)
|
||||
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()
|
||||
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)
|
||||
|
||||
def hipassfilter_m(self, samples, threshold):
|
||||
filtered_mask = samples.clone()
|
||||
|
||||
+1
-1
@@ -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.7"
|
||||
version = "3.0.8"
|
||||
license = { file = "LICENSE" }
|
||||
|
||||
[project.urls]
|
||||
|
||||
Reference in New Issue
Block a user