expand batches

This commit is contained in:
matt3o
2023-12-04 11:31:57 +01:00
parent 994e99b6d1
commit fe4a9dee5e
+102
View File
@@ -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",