From acc2d687d596bf82c2075f9a24003eacf18adfe7 Mon Sep 17 00:00:00 2001 From: bymyself Date: Tue, 14 May 2024 12:16:56 -0700 Subject: [PATCH] =?UTF-8?q?fix:=20=F0=9F=90=9B=20ImageCompare=20improvemen?= =?UTF-8?q?ts=20(#176)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * avoid unnecessary numpy conversion for diff and blend * add support for Batch * add support for input mismatch (RGB/RGBA) * fixes #175 --------- Co-authored-by: Mel Massadian --- nodes/image_processing.py | 53 +++++++++++++++++++++++++++++++++------ 1 file changed, 46 insertions(+), 7 deletions(-) diff --git a/nodes/image_processing.py b/nodes/image_processing.py index ba4a771..6b15eaf 100644 --- a/nodes/image_processing.py +++ b/nodes/image_processing.py @@ -216,16 +216,55 @@ class MTB_ImageCompare: CATEGORY = "mtb/image" def compare(self, imageA: torch.Tensor, imageB: torch.Tensor, mode): - imageA = imageA.numpy() - imageB = imageB.numpy() + if imageA.dim() == 4: + batch_count = imageA.size(0) + return ( + torch.cat( + tuple( + self.compare(imageA[i], imageB[i], mode)[0] + for i in range(batch_count) + ), + dim=0, + ), + ) - imageA = imageA.squeeze() - imageB = imageB.squeeze() + num_channels_A = imageA.size(2) + num_channels_B = imageB.size(2) - image = compare_images(imageA, imageB, method=mode) + # handle RGBA/RGB mismatch + if num_channels_A == 3 and num_channels_B == 4: + imageA = torch.cat( + (imageA, torch.ones_like(imageA[:, :, 0:1])), dim=2 + ) + elif num_channels_B == 3 and num_channels_A == 4: + imageB = torch.cat( + (imageB, torch.ones_like(imageB[:, :, 0:1])), dim=2 + ) + match mode: + case "diff": + compare_image = torch.abs(imageA - imageB) + case "blend": + compare_image = 0.5 * (imageA + imageB) + case "checkerboard": + imageA = imageA.numpy() + imageB = imageB.numpy() + compared_channels = [ + torch.from_numpy( + compare_images( + imageA[:, :, i], imageB[:, :, i], method=mode + ) + ) + for i in range(imageA.shape[2]) + ] - image = np.expand_dims(image, axis=0) - return (torch.from_numpy(image),) + compare_image = torch.stack(compared_channels, dim=2) + case _: + compare_image = None + raise ValueError(f"Unknown mode {mode}") + + compare_image = compare_image.unsqueeze(0) + + return (compare_image,) import requests