Also revert other rescaling modes to CPU, as it looks much better

This commit is contained in:
Luis Quesada
2026-02-05 22:20:40 +01:00
parent cfcac1b85c
commit 1ab691700c
2 changed files with 28 additions and 31 deletions
+27 -30
View File
@@ -601,44 +601,41 @@ class GPUProcessorLogic(ProcessorLogic):
# samples shape: [B, H, W, C]
mode = algorithm.lower()
# For algorithms not supported by torch.interpolate, use PIL on CPU
if mode in ('lanczos', 'box', 'hamming'):
original_device = samples.device
samples = samples.movedim(-1, 1) # [B, C, H, W]
algorithm_enum = getattr(Image, algorithm.upper())
results = []
for i in range(samples.shape[0]):
samples_pil: Image.Image = F.to_pil_image(samples[i].cpu()).resize((width, height), algorithm_enum)
results.append(F.to_tensor(samples_pil))
samples = torch.stack(results, dim=0).to(original_device)
samples = samples.movedim(1, -1)
return samples
# CPU works better, fallback to CPU for rescaling
original_device = samples.device
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)
algorithm_enum = getattr(Image, algorithm.upper())
results = []
for i in range(samples.shape[0]):
samples_pil: Image.Image = F.to_pil_image(samples[i].cpu()).resize((width, height), algorithm_enum)
results.append(F.to_tensor(samples_pil))
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]
mode = algorithm.lower()
# For algorithms not supported by torch.interpolate, use PIL on CPU
if mode in ('lanczos', 'box', 'hamming'):
original_device = samples.device
algorithm_enum = getattr(Image, algorithm.upper())
results = []
for i in range(samples.shape[0]):
samples_pil: Image.Image = F.to_pil_image(samples[i].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
# For torch-supported modes, use GPU interpolate
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)
# CPU works better, fallback to CPU for rescaling
original_device = samples.device
algorithm_enum = getattr(Image, algorithm.upper())
results = []
for i in range(samples.shape[0]):
samples_pil: Image.Image = F.to_pil_image(samples[i].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):
# samples shape: [B, H, W]
@@ -1234,7 +1231,7 @@ class InpaintCropImproved:
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")
+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.3"
version = "3.0.4"
license = { file = "LICENSE" }
[project.urls]