add ImageComposite compatible with animations

This commit is contained in:
cubiq
2024-07-19 14:37:00 +02:00
parent c84392ed27
commit b44bd5b1ba
+88 -1
View File
@@ -237,6 +237,91 @@ class ImageCompositeFromMaskBatch:
return (out, )
class ImageComposite:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"destination": ("IMAGE",),
"source": ("IMAGE",),
"x": ("INT", { "default": 0, "min": -MAX_RESOLUTION, "max": MAX_RESOLUTION, "step": 1 }),
"y": ("INT", { "default": 0, "min": -MAX_RESOLUTION, "max": MAX_RESOLUTION, "step": 1 }),
"offset_x": ("INT", { "default": 0, "min": -MAX_RESOLUTION, "max": MAX_RESOLUTION, "step": 1 }),
"offset_y": ("INT", { "default": 0, "min": -MAX_RESOLUTION, "max": MAX_RESOLUTION, "step": 1 }),
},
"optional": {
"mask": ("MASK",),
}
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "execute"
CATEGORY = "essentials/image manipulation"
def execute(self, destination, source, x, y, offset_x, offset_y, mask=None):
if mask is None:
mask = torch.ones_like(source)[:,:,:,0]
mask = mask.unsqueeze(-1).repeat(1, 1, 1, 3)
if mask.shape[1:3] != source.shape[1:3]:
mask = F.interpolate(mask.permute([0, 3, 1, 2]), size=(source.shape[1], source.shape[2]), mode='bicubic')
mask = mask.permute([0, 2, 3, 1])
if mask.shape[0] > source.shape[0]:
mask = mask[:source.shape[0]]
elif mask.shape[0] < source.shape[0]:
mask = torch.cat((mask, mask[-1:].repeat((source.shape[0]-mask.shape[0], 1, 1, 1))), dim=0)
if destination.shape[0] > source.shape[0]:
destination = destination[:source.shape[0]]
elif destination.shape[0] < source.shape[0]:
destination = torch.cat((destination, destination[-1:].repeat((source.shape[0]-destination.shape[0], 1, 1, 1))), dim=0)
if not isinstance(x, list):
x = [x]
if not isinstance(y, list):
y = [y]
if len(x) < destination.shape[0]:
x = x + [x[-1]] * (destination.shape[0] - len(x))
if len(y) < destination.shape[0]:
y = y + [y[-1]] * (destination.shape[0] - len(y))
x = [i + offset_x for i in x]
y = [i + offset_y for i in y]
output = []
for i in range(destination.shape[0]):
d = destination[i].clone()
s = source[i]
m = mask[i]
if x[i]+source.shape[2] > destination.shape[2]:
s = s[:, :, :destination.shape[2]-x[i], :]
m = m[:, :, :destination.shape[2]-x[i], :]
if y[i]+source.shape[1] > destination.shape[1]:
s = s[:, :destination.shape[1]-y[i], :, :]
m = m[:destination.shape[1]-y[i], :, :]
#output.append(s * m + d[y[i]:y[i]+s.shape[0], x[i]:x[i]+s.shape[1], :] * (1 - m))
d[y[i]:y[i]+s.shape[0], x[i]:x[i]+s.shape[1], :] = s * m + d[y[i]:y[i]+s.shape[0], x[i]:x[i]+s.shape[1], :] * (1 - m)
output.append(d)
output = torch.stack(output)
# apply the source to the destination at XY position using the mask
#for i in range(destination.shape[0]):
# output[i, y[i]:y[i]+source.shape[1], x[i]:x[i]+source.shape[2], :] = source * mask + destination[i, y[i]:y[i]+source.shape[1], x[i]:x[i]+source.shape[2], :] * (1 - mask)
#for x_, y_ in zip(x, y):
# output[:, y_:y_+source.shape[1], x_:x_+source.shape[2], :] = source * mask + destination[:, y_:y_+source.shape[1], x_:x_+source.shape[2], :] * (1 - mask)
#output[:, y:y+source.shape[1], x:x+source.shape[2], :] = source * mask + destination[:, y:y+source.shape[1], x:x+source.shape[2], :] * (1 - mask)
#output = destination * (1 - mask) + source * mask
return (output,)
class ImageResize:
@classmethod
def INPUT_TYPES(s):
@@ -1238,7 +1323,7 @@ class NoiseFromImage:
RETURN_TYPES = ("IMAGE",)
FUNCTION = "execute"
CATEGORY = "essentials"
CATEGORY = "essentials/image utils"
def execute(self, image, noise_size, color_noise, mask_strength, mask_scale_diff, mask_contrast, noise_strenght, saturation, contrast, blur, noise_mask=None):
torch.manual_seed(0)
@@ -1330,6 +1415,7 @@ IMAGE_CLASS_MAPPINGS = {
# Image manipulation
"ImageCompositeFromMaskBatch+": ImageCompositeFromMaskBatch,
"ImageComposite+": ImageComposite,
"ImageCrop+": ImageCrop,
"ImageFlip+": ImageFlip,
"ImageRandomTransform+": ImageRandomTransform,
@@ -1371,6 +1457,7 @@ IMAGE_NAME_MAPPINGS = {
# Image manipulation
"ImageCompositeFromMaskBatch+": "🔧 Image Composite From Mask Batch",
"ImageComposite+": "🔧 Image Composite",
"ImageCrop+": "🔧 Image Crop",
"ImageFlip+": "🔧 Image Flip",
"ImageRandomTransform+": "🔧 Image Random Transform",