diff --git a/README.md b/README.md index 77c7661..d884190 100644 --- a/README.md +++ b/README.md @@ -2,6 +2,8 @@ ComfyUI-AutoCropBgTrim is a powerful tool designed to automatically clean up the background of your images. This tool trims unnecessary spaces and pixels, leaving only the main subject of the image. It generates both a mask and an image output, making it easy to focus on the essential elements. Perfect for enhancing your photos and preparing them for professional use. +![Demo](demo.png) + ## Features - Automatically trims backgrounds. - Leaves only the main subject. diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..945b7ef --- /dev/null +++ b/__init__.py @@ -0,0 +1,3 @@ +from .ron_layers_trim_bg_ultra_v2 import RonLayersTrimBgUltraV2, NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + +__all__ = ['RonLayersTrimBgUltraV2', 'NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] \ No newline at end of file diff --git a/demo.png b/demo.png new file mode 100644 index 0000000..e43502f Binary files /dev/null and b/demo.png differ diff --git a/imagefunc.py b/imagefunc.py new file mode 100644 index 0000000..0f918f2 --- /dev/null +++ b/imagefunc.py @@ -0,0 +1,33 @@ +import torch +import numpy as np +from PIL import Image, ImageDraw +from .imagefunc import * + +def tensor2pil(t_image: torch.Tensor) -> Image: + return Image.fromarray(np.clip(255.0 * t_image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8)) + +def pil2tensor(p_image: Image) -> torch.Tensor: + return torch.from_numpy(np.array(p_image).astype(np.float32) / 255.0).unsqueeze(0) + +def image2mask(p_image: Image) -> torch.Tensor: + if p_image.mode == "RGB": + p_image = p_image.convert("L") + return torch.from_numpy(np.array(p_image).astype(np.float32) / 255.0).unsqueeze(0) + +def mask2image(t_mask: torch.Tensor) -> Image: + return Image.fromarray(np.clip(255.0 * t_mask.cpu().numpy().squeeze(), 0, 255).astype(np.uint8), mode="L") + +def draw_rect(image, x, y, width, height, line_color, line_width): + draw = ImageDraw.Draw(image) + draw.rectangle((x, y, x + width, y + height), outline=line_color, width=line_width) + return image + +def log(message, message_type='info'): + if message_type == 'info': + print(f"INFO: {message}") + elif message_type == 'warning': + print(f"WARNING: {message}") + elif message_type == 'error': + print(f"ERROR: {message}") + elif message_type == 'finish': + print(f"FINISH: {message}") \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..a50efea --- /dev/null +++ b/requirements.txt @@ -0,0 +1,3 @@ +torch==1.10.0 +numpy==1.21.2 +Pillow==8.4.0 \ No newline at end of file diff --git a/ron_layers_trim_bg_ultra_v2.py b/ron_layers_trim_bg_ultra_v2.py new file mode 100644 index 0000000..a79a607 --- /dev/null +++ b/ron_layers_trim_bg_ultra_v2.py @@ -0,0 +1,82 @@ +import torch +import numpy as np +from PIL import Image +from .imagefunc import * + +NODE_NAME = 'RonLayersTrimBgUltraV2' + +class RonLayersTrimBgUltraV2: + def __init__(self): + self.input_image = None + self.input_mask = None + self.output_image = None + self.output_mask = None + self.crop_box = None + self.box_preview = None + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "mask": ("MASK",), + "padding": ("INT", {"default": 0, "min": 0, "max": 1000, "step": 1}), + } + } + + RETURN_TYPES = ("IMAGE", "MASK", "BOX", "IMAGE") + RETURN_NAMES = ("croped_image", "croped_mask", "crop_box", "box_preview") + FUNCTION = "trim_and_crop_by_mask" + CATEGORY = "RonLayers/TrimBg" + + def trim_and_crop_by_mask(self, image, mask, padding): + ret_images = [] + ret_masks = [] + + if mask.dim() == 2: + mask = torch.unsqueeze(mask, 0) + if mask.shape[0] > 1: + log(f"Warning: Multiple mask inputs, using the first.", message_type='warning') + mask = torch.unsqueeze(mask[0], 0) + + image_pil = tensor2pil(image).convert('RGBA') + mask_pil = tensor2pil(mask).convert('L') + + masked_image = Image.composite(image_pil, Image.new("RGBA", image_pil.size), mask_pil) + bbox = masked_image.getbbox() + + if bbox: + x1, y1, x2, y2 = bbox + x1 = x1 - padding if x1 - padding > 0 else 0 + y1 = y1 - padding if y1 - padding > 0 else 0 + x2 = x2 + padding if x2 + padding < image_pil.width else image_pil.width + y2 = y2 + padding if y2 + padding < image_pil.height else image_pil.height + crop_box = (x1, y1, x2, y2) + cropped_image = Image.composite(image_pil.crop(crop_box), Image.new("RGBA", (x2 - x1, y2 - y1), (0, 0, 0, 0)), mask_pil.crop(crop_box)) + cropped_mask = mask_pil.crop(crop_box) + preview_image = draw_rect(tensor2pil(mask).convert('RGB'), x1, y1, x2 - x1, y2 - y1, line_color="#00F000", + line_width=(x2 - x1 + y2 - y1) // 200) + ret_images.append(pil2tensor(cropped_image.convert("RGBA"))) + ret_masks.append(image2mask(cropped_mask)) + else: + ret_images.append(image) + ret_masks.append(mask) + crop_box = None + preview_image = tensor2pil(mask).convert('RGB') + + log(f"{NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish') + return (torch.cat(ret_images, dim=0), torch.cat(ret_masks, dim=0), crop_box, pil2tensor(preview_image)) + + def run(self, image, mask, padding): + self.input_image = image + self.input_mask = mask + self.output_image, self.output_mask, self.crop_box, self.box_preview = self.trim_and_crop_by_mask(image=self.input_image, mask=self.input_mask, padding=padding) + return self.output_image, self.output_mask, self.crop_box, self.box_preview + +NODE_CLASS_MAPPINGS = { + "RonLayers/TrimBg: RonLayersTrimBgUltraV2": RonLayersTrimBgUltraV2 +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "RonLayers/TrimBg: RonLayersTrimBgUltraV2": "RonLayers/TrimBg: RonLayersTrimBgUltraV2" +} \ No newline at end of file