diff --git a/image.py b/image.py index 15ee324..1d7cdd0 100644 --- a/image.py +++ b/image.py @@ -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",