From fe4a9dee5ebdc8c337fd07113fe1cb853059e752 Mon Sep 17 00:00:00 2001 From: matt3o Date: Mon, 4 Dec 2023 11:31:57 +0100 Subject: [PATCH] expand batches --- essentials.py | 102 ++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 102 insertions(+) diff --git a/essentials.py b/essentials.py index 558d9ed..f0a1899 100644 --- a/essentials.py +++ b/essentials.py @@ -242,6 +242,55 @@ class ImageEnhanceDifference: return(diff_image,) +class ImageExpandBatch: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "size": ("INT", { "default": 16, "min": 1, "step": 1, }), + "method": (["expand", "repeat all", "repeat first", "repeat last"],) + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "execute" + CATEGORY = "essentials" + + def execute(self, image, size, method): + orig_size = image.shape[0] + + if orig_size == size: + return (image,) + + if size <= 1: + return (image[:size],) + + if 'expand' in method: + out = torch.empty([size] + list(image.shape)[1:], dtype=image.dtype, device=image.device) + if size < orig_size: + scale = (orig_size - 1) / (size - 1) + for i in range(size): + out[i] = image[min(round(i * scale), orig_size - 1)] + else: + scale = orig_size / size + for i in range(size): + out[i] = image[min(math.floor((i + 0.5) * scale), orig_size - 1)] + elif 'all' in method: + out = image.repeat([math.ceil(size / image.shape[0])] + [1] * (len(image.shape) - 1))[:size] + elif 'first' in method: + if size < image.shape[0]: + out = image[:size] + else: + out = torch.cat([image[:1].repeat(size-image.shape[0], 1, 1, 1), image], dim=0) + elif 'last' in method: + if size < image.shape[0]: + out = image[:size] + else: + out = torch.cat((image, image[-1:].repeat((size-image.shape[0], 1, 1, 1))), dim=0) + + return (out,) + class MaskFlip: @classmethod def INPUT_TYPES(s): @@ -337,6 +386,55 @@ class MaskBatch: out = torch.cat((mask1, mask2), dim=0) return (out,) +class MaskExpandBatch: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "mask": ("MASK",), + "size": ("INT", { "default": 16, "min": 1, "step": 1, }), + "method": (["expand", "repeat all", "repeat first", "repeat last"],) + } + } + + RETURN_TYPES = ("MASK",) + FUNCTION = "execute" + CATEGORY = "essentials" + + def execute(self, mask, size, method): + orig_size = mask.shape[0] + + if orig_size == size: + return (mask,) + + if size <= 1: + return (mask[:size],) + + if 'expand' in method: + out = torch.empty([size] + list(mask.shape)[1:], dtype=mask.dtype, device=mask.device) + if size < orig_size: + scale = (orig_size - 1) / (size - 1) + for i in range(size): + out[i] = mask[min(round(i * scale), orig_size - 1)] + else: + scale = orig_size / size + for i in range(size): + out[i] = mask[min(math.floor((i + 0.5) * scale), orig_size - 1)] + elif 'all' in method: + out = mask.repeat([math.ceil(size / mask.shape[0])] + [1] * (len(mask.shape) - 1))[:size] + elif 'first' in method: + if size < mask.shape[0]: + out = mask[:size] + else: + out = torch.cat([mask[:1].repeat(size-mask.shape[0], 1, 1), mask], dim=0) + elif 'last' in method: + if size < mask.shape[0]: + out = mask[:size] + else: + out = torch.cat((mask, mask[-1:].repeat((size-mask.shape[0], 1, 1))), dim=0) + + return (out,) + def cubic_bezier(t, p): p0, p1, p2, p3 = p return (1 - t)**3 * p0 + 3 * (1 - t)**2 * t * p1 + 3 * (1 - t) * t**2 * p2 + t**3 * p3 @@ -613,11 +711,13 @@ NODE_CLASS_MAPPINGS = { "ImagePosterize+": ImagePosterize, "ImageCASharpening+": ImageCAS, "ImageEnhanceDifference+": ImageEnhanceDifference, + "ImageExpandBatch+": ImageExpandBatch, "MaskBlur+": MaskBlur, "MaskFlip+": MaskFlip, "MaskPreview+": MaskPreview, "MaskBatch+": MaskBatch, + "MaskExpandBatch+": MaskExpandBatch, "TransitionMask+": TransitionMask, "SimpleMath+": SimpleMath, @@ -637,11 +737,13 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ImagePosterize+": "🔧 Image Posterize", "ImageCASharpening+": "🔧 Image Contrast Adaptive Sharpening", "ImageEnhanceDifference+": "🔧 Image Enhance Difference", + "ImageExpandBatch+": "🔧 Image Expand Batch", "MaskBlur+": "🔧 Mask Blur", "MaskFlip+": "🔧 Mask Flip", "MaskPreview+": "🔧 Mask Preview", "MaskBatch+": "🔧 Mask Batch", + "MaskExpandBatch+": "🔧 Mask Expand Batch", "TransitionMask+": "🔧 Transition Mask", "SimpleMath+": "🔧 Simple Math",