From b07b203e05cd72fa515d619f91154a88b367369f Mon Sep 17 00:00:00 2001 From: AI Lab <129358391+1038lab@users.noreply.github.com> Date: Sat, 5 Apr 2025 02:01:12 -0700 Subject: [PATCH] Add files via upload --- AILab_ImageMaskTools.py | 562 +++++++++++++++- AILab_RMBG.py | 1336 ++++++++++++++++++++------------------- requirements.txt | 7 +- 3 files changed, 1228 insertions(+), 677 deletions(-) diff --git a/AILab_ImageMaskTools.py b/AILab_ImageMaskTools.py index 33659b8..c603105 100644 --- a/AILab_ImageMaskTools.py +++ b/AILab_ImageMaskTools.py @@ -1,4 +1,4 @@ -# ComfyUI-RMBG v2.0.0 +# ComfyUI-RMBG v2.2.0 # # This node facilitates background removal using various models, including RMBG-2.0, INSPYRENET, BEN, BEN2, and BIREFNET-HR. # It utilizes advanced deep learning techniques to process images and generate accurate masks for background removal. @@ -8,18 +8,23 @@ # It offers a collection of utility nodes for efficient handling of images and masks: # # 1. Preview Nodes: -# - AiLab_Preview: A universal preview tool for both images and masks. -# - AiLab_ImagePreview: A specialized preview tool for images. -# - AiLab_MaskPreview: A specialized preview tool for masks. -# - AiLab_LoadImage: A node for loading images with some Frequently used options. +# - Preview: A universal preview tool for both images and masks. +# - ImagePreview: A specialized preview tool for images. +# - MaskPreview: A specialized preview tool for masks. +# - LoadImage: A node for loading images with some Frequently used options. +# +# 2. Conversion Node: +# - ImageMaskConvert: Converts between image and mask formats and extracts masks from image channels. +# +# 3. Mask Processing Nodes: +# - MaskEnhancer: Refines masks through techniques such as blur, smoothing, expansion/contraction, and hole filling. +# - MaskCombiner: Combines multiple masks using union, intersection, or difference operations. +# +# 4. Image Processing Nodes: +# - ImageCombiner: Combines foreground and background images with various blending modes and positioning options. +# - ImageStitch: Stitches multiple images together in various directions. # # These nodes are crafted to streamline common image and mask operations within ComfyUI workflows. -# -# This integration script follows GPL-3.0 License. -# When using or modifying this code, please respect both the original model licenses -# and this integration's license terms. -# -# Source: https://github.com/1038lab/ComfyUI-RMBG import os import random @@ -30,6 +35,7 @@ import torch import cv2 from PIL import Image, ImageFilter, ImageOps, ImageSequence, ImageChops import torchvision.transforms.functional as T +from comfy.utils import common_upscale from scipy import ndimage # Utility functions @@ -52,7 +58,7 @@ def blend_overlay(img_1, img_2): return Image.fromarray(np.clip(result * 255, 0, 255).astype(np.uint8)) # Base class for preview -class AiLab_PreviewBase: +class AILab_PreviewBase: def __init__(self): self.output_dir = folder_paths.get_temp_directory() self.type = "temp" @@ -98,7 +104,7 @@ class AiLab_PreviewBase: return {"ui": {}} # Preview node -class AiLab_Preview(AiLab_PreviewBase): +class AILab_Preview(AILab_PreviewBase): def __init__(self): super().__init__() self.prefix_append = "_preview_" + ''.join(random.choice("abcdefghijklmnopqrstupvxyz") for x in range(5)) @@ -139,7 +145,7 @@ class AiLab_Preview(AiLab_PreviewBase): } # Mask preview node -class AiLab_MaskPreview(AiLab_PreviewBase): +class AILab_MaskPreview(AILab_PreviewBase): def __init__(self): super().__init__() self.prefix_append = "_mask_preview_" + ''.join(random.choice("abcdefghijklmnopqrstupvxyz") for x in range(5)) @@ -166,7 +172,7 @@ class AiLab_MaskPreview(AiLab_PreviewBase): } # Image preview node -class AiLab_ImagePreview(AiLab_PreviewBase): +class AILab_ImagePreview(AILab_PreviewBase): def __init__(self): super().__init__() self.prefix_append = "_image_preview_" + ''.join(random.choice("abcdefghijklmnopqrstupvxyz") for x in range(5)) @@ -191,8 +197,235 @@ class AiLab_ImagePreview(AiLab_PreviewBase): "result": (image,) } +# Image mask conversion node +class AILab_ImageMaskConvert: + @classmethod + def INPUT_TYPES(cls): + return { + "required": {}, + "optional": { + "image": ("IMAGE",), + "mask": ("MASK",), + "mask_channel": (["alpha", "red", "green", "blue"], {"default": "alpha"}) + } + } + + RETURN_TYPES = ("IMAGE", "MASK") + RETURN_NAMES = ("IMAGE", "MASK") + FUNCTION = "convert" + CATEGORY = "🧪AILab/🛠️UTIL/🖼️IMAGE" + + def convert(self, image=None, mask=None, mask_channel="alpha"): + # Case 1: No inputs + if image is None and mask is None: + empty_image = torch.zeros(1, 3, 64, 64) + empty_mask = torch.zeros(1, 64, 64) + return (empty_image, empty_mask) + + # Case 2: Only mask input + if image is None and mask is not None: + if mask.ndim == 4: + tensor = mask.permute(0, 2, 3, 1) + tensor_rgb = torch.cat([tensor] * 3, dim=-1) + return (tensor_rgb, mask) + elif mask.ndim == 3: + tensor = mask.unsqueeze(-1) + tensor_rgb = torch.cat([tensor] * 3, dim=-1) + return (tensor_rgb, mask) + elif mask.ndim == 2: + tensor = mask.unsqueeze(0).unsqueeze(-1) + tensor_rgb = torch.cat([tensor] * 3, dim=-1) + return (tensor_rgb, mask.unsqueeze(0)) + else: + print(f"Invalid mask shape: {mask.shape}") + empty_image = torch.zeros(1, 3, 64, 64) + return (empty_image, mask) + + # Case 3: Only image input + if image is not None and mask is None: + mask_list = [] + for img in image: + pil_img = tensor2pil(img) + pil_img = pil_img.convert("RGBA") + r, g, b, a = pil_img.split() + if mask_channel == "red": + channel_img = r + elif mask_channel == "green": + channel_img = g + elif mask_channel == "blue": + channel_img = b + elif mask_channel == "alpha": + channel_img = a + mask = np.array(channel_img.convert("L")).astype(np.float32) / 255.0 + mask_tensor = torch.from_numpy(mask) + mask_list.append(mask_tensor) + result_mask = torch.stack(mask_list) + return (image, result_mask) + + if image is not None and mask is not None: + if mask.ndim == 4: # [B,C,H,W] + mask = mask.squeeze(1) # Convert to [B,H,W] + return (image, mask) + +# Mask enhancer node +class AILab_MaskEnhancer: + @classmethod + def INPUT_TYPES(cls): + tooltips = { + "mask": "Input mask to be processed.", + "sensitivity": "Adjust the strength of mask detection (higher values result in more aggressive detection).", + "mask_blur": "Specify the amount of blur to apply to the mask edges (0 for no blur, higher values for more blur).", + "mask_offset": "Adjust the mask boundary (positive values expand the mask, negative values shrink it).", + "smooth": "Smooth the mask edges (0 for no smoothing, higher values create smoother edges).", + "fill_region": "Enable to fill holes in the mask.", + "invert_output": "Enable to invert the mask output (useful for certain effects)." + } + + return { + "required": { + "mask": ("MASK", {"tooltip": tooltips["mask"]}), + }, + "optional": { + "sensitivity": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": tooltips["sensitivity"]}), + "mask_blur": ("INT", {"default": 0, "min": 0, "max": 64, "step": 1, "tooltip": tooltips["mask_blur"]}), + "mask_offset": ("INT", {"default": 0, "min": -64, "max": 64, "step": 1, "tooltip": tooltips["mask_offset"]}), + "smooth": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 128.0, "step": 0.5, "tooltip": tooltips["smooth"]}), + "fill_region": ("BOOLEAN", {"default": False, "tooltip": tooltips["fill_region"]}), + "invert_output": ("BOOLEAN", {"default": False, "tooltip": tooltips["invert_output"]}), + } + } + + RETURN_TYPES = ("MASK",) + RETURN_NAMES = ("MASK",) + FUNCTION = "process_mask" + CATEGORY = "🧪AILab/🛠️UTIL/🖼️IMAGE" + + def fill_mask_region(self, mask_pil): + """Fill holes in the mask""" + mask_np = np.array(mask_pil) + contours, _ = cv2.findContours(mask_np, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) + filled_mask = np.zeros_like(mask_np) + for contour in contours: + cv2.drawContours(filled_mask, [contour], 0, 255, -1) # -1 means fill + return Image.fromarray(filled_mask) + + def process_mask(self, mask, sensitivity=1.0, mask_blur=0, mask_offset=0, smooth=0.0, + fill_region=False, invert_output=False): + processed_masks = [] + + for mask_item in mask: + m = mask_item * (1 + (1 - sensitivity)) + m = torch.clamp(m, 0, 1) + + if smooth > 0: + mask_np = m.cpu().numpy() + binary_mask = (mask_np > 0.5).astype(np.float32) + blurred_mask = ndimage.gaussian_filter(binary_mask, sigma=smooth) + final_mask = (blurred_mask > 0.5).astype(np.float32) + m = torch.from_numpy(final_mask) + + if fill_region: + mask_pil = tensor2pil(m) + mask_pil = self.fill_mask_region(mask_pil) + m = pil2tensor(mask_pil).squeeze(0) + + if mask_blur > 0: + mask_pil = tensor2pil(m) + mask_pil = mask_pil.filter(ImageFilter.GaussianBlur(radius=mask_blur)) + m = pil2tensor(mask_pil).squeeze(0) + + if mask_offset != 0: + mask_pil = tensor2pil(m) + if mask_offset > 0: + for _ in range(mask_offset): + mask_pil = mask_pil.filter(ImageFilter.MaxFilter(3)) + else: + for _ in range(-mask_offset): + mask_pil = mask_pil.filter(ImageFilter.MinFilter(3)) + m = pil2tensor(mask_pil).squeeze(0) + + if invert_output: + m = 1.0 - m + + processed_masks.append(m.unsqueeze(0)) + + return (torch.cat(processed_masks, dim=0),) + +# Mask combiner node +class AILab_MaskCombiner: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "mask_1": ("MASK",), + "mode": (["combine", "intersection", "difference"], {"default": "combine"}) + }, + "optional": { + "mask_2": ("MASK", {"default": None}), + "mask_3": ("MASK", {"default": None}), + "mask_4": ("MASK", {"default": None}) + } + } + + CATEGORY = "🧪AILab/🛠️UTIL/🖼️IMAGE" + RETURN_TYPES = ("MASK",) + FUNCTION = "combine_masks" + + def combine_masks(self, mask_1, mode="combine", mask_2=None, mask_3=None, mask_4=None): + try: + masks = [m for m in [mask_1, mask_2, mask_3, mask_4] if m is not None] + + if len(masks) <= 1: + return (masks[0] if masks else torch.zeros((1, 64, 64), dtype=torch.float32),) + + ref_shape = masks[0].shape + masks = [self._resize_if_needed(m, ref_shape) for m in masks] + + if mode == "combine": + result = torch.maximum(masks[0], masks[1]) + for mask in masks[2:]: + result = torch.maximum(result, mask) + elif mode == "intersection": + result = torch.minimum(masks[0], masks[1]) + else: + result = torch.abs(masks[0] - masks[1]) + + return (torch.clamp(result, 0, 1),) + except Exception as e: + print(f"Error in combine_masks: {str(e)}") + print(f"Mask shapes: {[m.shape for m in masks]}") + raise e + + def _resize_if_needed(self, mask, target_shape): + try: + if mask.shape == target_shape: + return mask + + if len(mask.shape) == 2: + mask = mask.unsqueeze(0) + elif len(mask.shape) == 4: + mask = mask.squeeze(1) + + target_height = target_shape[-2] if len(target_shape) >= 2 else target_shape[0] + target_width = target_shape[-1] if len(target_shape) >= 2 else target_shape[1] + + resized_masks = [] + for i in range(mask.shape[0]): + mask_np = mask[i].cpu().numpy() + img = Image.fromarray((mask_np * 255).astype(np.uint8)) + img_resized = img.resize((target_width, target_height), Image.LANCZOS) + mask_resized = np.array(img_resized).astype(np.float32) / 255.0 + resized_masks.append(torch.from_numpy(mask_resized)) + + return torch.stack(resized_masks) + + except Exception as e: + print(f"Error in _resize_if_needed: {str(e)}") + print(f"Input mask shape: {mask.shape}, Target shape: {target_shape}") + raise e + # Image loader node -class AiLab_LoadImage: +class AILab_LoadImage: @classmethod def INPUT_TYPES(cls): input_dir = folder_paths.get_input_directory() @@ -301,20 +534,301 @@ class AiLab_LoadImage: return True +# Image combiner node +class AILab_ImageCombiner: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "foreground": ("IMAGE",), + "background": ("IMAGE",), + "mode": (["normal", "multiply", "screen", "overlay", "add", "subtract"], + {"default": "normal"}), + "foreground_opacity": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), + "foreground_scale": ("FLOAT", {"default": 1.0, "min": 0.1, "max": 5.0, "step": 0.05}), + "position_x": ("INT", {"default": 50, "min": 0, "max": 100, "step": 1}), + "position_y": ("INT", {"default": 50, "min": 0, "max": 100, "step": 1}), + }, + "optional": { + "foreground_mask": ("MASK", {"default": None}), + } + } + CATEGORY = "🧪AILab/🛠️UTIL/🖼️IMAGE" + RETURN_TYPES = ("IMAGE",) + FUNCTION = "combine_images" + + def combine_images(self, foreground, background, mode="normal", foreground_opacity=1.0, + foreground_scale=1.0, position_x=50, position_y=50, foreground_mask=None): + if len(foreground.shape) == 3: + foreground = foreground.unsqueeze(0) + if len(background.shape) == 3: + background = background.unsqueeze(0) + + batch_size = foreground.shape[0] + output_images = [] + + for b in range(batch_size): + fg_pil = tensor2pil(foreground[b]) + bg_pil = tensor2pil(background[b]) + + if fg_pil.mode != 'RGBA': + fg_pil = fg_pil.convert('RGBA') + + if foreground_scale != 1.0: + new_width = int(fg_pil.width * foreground_scale) + new_height = int(fg_pil.height * foreground_scale) + fg_pil = fg_pil.resize((new_width, new_height), Image.LANCZOS) + + if foreground_mask is not None: + mask_tensor = foreground_mask[b] if len(foreground_mask.shape) > 2 else foreground_mask + mask_pil = Image.fromarray(np.uint8(mask_tensor.cpu().numpy() * 255)) + if mask_pil.size != fg_pil.size: + mask_pil = mask_pil.resize(fg_pil.size, Image.LANCZOS) + r, g, b, a = fg_pil.split() + a = ImageChops.multiply(a, mask_pil) + fg_pil = Image.merge('RGBA', (r, g, b, a)) + + fg_w, fg_h = fg_pil.size + bg_w, bg_h = bg_pil.size + + x = int(bg_w * position_x / 100 - fg_w / 2) + y = int(bg_h * position_y / 100 - fg_h / 2) + + new_fg = Image.new('RGBA', (bg_w, bg_h), (0, 0, 0, 0)) + new_fg.paste(fg_pil, (x, y), fg_pil) + fg_pil = new_fg + + if bg_pil.mode != 'RGBA': + bg_pil = bg_pil.convert('RGBA') + + if foreground_opacity < 1.0: + r, g, b, a = fg_pil.split() + a = Image.eval(a, lambda x: int(x * foreground_opacity)) + fg_pil = Image.merge('RGBA', (r, g, b, a)) + + if mode == "normal": + result = bg_pil.copy() + result = Image.alpha_composite(result, fg_pil) + else: + alpha = fg_pil.split()[3] + fg_rgb = fg_pil.convert('RGB') + bg_rgb = bg_pil.convert('RGB') + + if mode == "multiply": + blended = ImageChops.multiply(fg_rgb, bg_rgb) + elif mode == "screen": + blended = ImageChops.screen(fg_rgb, bg_rgb) + elif mode == "add": + blended = ImageChops.add(fg_rgb, bg_rgb, 1.0) + elif mode == "subtract": + blended = ImageChops.subtract(fg_rgb, bg_rgb, 1.0) + elif mode == "overlay": + blended = blend_overlay(fg_rgb, bg_rgb) + else: + blended = fg_rgb + + blended = blended.convert('RGBA') + r, g, b, _ = blended.split() + blended = Image.merge('RGBA', (r, g, b, alpha)) + result = bg_pil.copy() + result = Image.alpha_composite(result, blended) + + if result.mode != 'RGB': + white_bg = Image.new('RGB', result.size, 'white') + result = Image.alpha_composite(white_bg.convert('RGBA'), result) + result = result.convert('RGB') + + output_images.append(pil2tensor(result)) + + return (torch.cat(output_images, dim=0),) + +class AILab_MaskExtractor: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "mask": ("MASK",), + "mode": (["extract_masked_area", "apply_mask", "invert_mask"], {"default": "invert_mask"}), + "background": (["transparent", "black", "white", "original"], {"default": "transparent"}) + } + } + CATEGORY = "🧪AILab/🛠️UTIL/🖼️IMAGE" + RETURN_TYPES = ("IMAGE",) + FUNCTION = "extract_masked_area" + + def _prepare_mask(self, mask_np, image_shape): + try: + if isinstance(mask_np, torch.Tensor): + mask_np = mask_np.cpu().numpy() + mask_np = np.array(mask_np) + while len(mask_np.shape) > 2 and mask_np.shape[-1] == 1: + mask_np = mask_np.squeeze(-1) + while len(mask_np.shape) > 2 and mask_np.shape[0] == 1: + mask_np = mask_np.squeeze(0) + if len(mask_np.shape) > 2: + mask_np = mask_np.squeeze() + if mask_np.shape != image_shape[:2]: + mask_pil = Image.fromarray((mask_np * 255).astype(np.uint8)) + mask_pil = mask_pil.resize((image_shape[1], image_shape[0]), Image.LANCZOS) + mask_np = np.array(mask_pil).astype(np.float32) / 255.0 + mask_np = mask_np[..., np.newaxis] + mask_np = np.repeat(mask_np, image_shape[2], axis=2) + return mask_np + except Exception as e: + print(f"Error in _prepare_mask: {str(e)}") + raise e + + def extract_masked_area(self, image, mask, mode="extract_masked_area", background="transparent"): + try: + pil_image = tensor2pil(image) + image_np = np.array(pil_image).astype(np.float32) / 255.0 + mask_np = self._prepare_mask(mask, image_np.shape) + result_np = np.zeros_like(image_np) + + if mode == "extract_masked_area": + result_np = image_np * mask_np + if background == "transparent": + if pil_image.mode != "RGBA": + pil_image = pil_image.convert("RGBA") + result_rgba = np.zeros((*image_np.shape[:2], 4), dtype=np.float32) + result_rgba[:, :, :3] = image_np * mask_np + result_rgba[:, :, 3] = mask_np[..., 0] + result_pil = Image.fromarray((result_rgba * 255).astype(np.uint8), mode="RGBA") + return (torch.from_numpy(np.array(result_pil).astype(np.float32) / 255.0).unsqueeze(0),) + elif background == "black": + pass # Already done with image_np * mask_np + elif background == "white": + result_np = result_np + (1 - mask_np) + elif background == "original": + result_np = image_np * mask_np + + elif mode == "apply_mask": + result_np = image_np * mask_np + if background == "transparent": + if pil_image.mode != "RGBA": + pil_image = pil_image.convert("RGBA") + result_rgba = np.zeros((*image_np.shape[:2], 4), dtype=np.float32) + result_rgba[:, :, :3] = image_np * mask_np + result_rgba[:, :, 3] = mask_np[..., 0] + result_pil = Image.fromarray((result_rgba * 255).astype(np.uint8), mode="RGBA") + return (torch.from_numpy(np.array(result_pil).astype(np.float32) / 255.0).unsqueeze(0),) + elif background == "white": + result_np = result_np + (1 - mask_np) + elif background == "original": + result_np = image_np * mask_np + image_np * (1 - mask_np) + + elif mode == "invert_mask": + result_np = image_np * (1 - mask_np) + if background == "transparent": + if pil_image.mode != "RGBA": + pil_image = pil_image.convert("RGBA") + result_rgba = np.zeros((*image_np.shape[:2], 4), dtype=np.float32) + result_rgba[:, :, :3] = image_np * (1 - mask_np) + result_rgba[:, :, 3] = (1 - mask_np)[..., 0] + result_pil = Image.fromarray((result_rgba * 255).astype(np.uint8), mode="RGBA") + return (torch.from_numpy(np.array(result_pil).astype(np.float32) / 255.0).unsqueeze(0),) + elif background == "white": + result_np = result_np + mask_np + elif background == "original": + result_np = image_np * (1 - mask_np) + image_np * mask_np + + result_pil = Image.fromarray(np.clip(result_np * 255, 0, 255).astype(np.uint8)) + return (pil2tensor(result_pil),) + except Exception as e: + print(f"Error in extract_masked_area: {str(e)}") + raise e + +# Image Stitch node +class AILab_ImageStitch: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "image1": ("IMAGE",), + "image2": ("IMAGE",), + "concat_direction": (['right', 'top', 'left', 'bottom'], {"default": 'right'}), + }} + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "stitch_images" + CATEGORY = "🧪AILab/🛠️UTIL/🖼️IMAGE" + + def stitch_images(self, image1, image2, concat_direction): + if image1.shape[0] != image2.shape[0]: + max_batch = max(image1.shape[0], image2.shape[0]) + image1 = image1.repeat(max_batch // image1.shape[0], 1, 1, 1) + image2 = image2.repeat(max_batch // image2.shape[0], 1, 1, 1) + + if concat_direction in ['right', 'left']: + # Match heights for horizontal stitching + h1 = image1.shape[1] + h2, w2 = image2.shape[1:3] + aspect = w2 / h2 + + new_h = h1 + new_w = int(h1 * aspect) + + image2 = self._resize(image2, new_w, new_h) + else: + # Match widths for vertical stitching + w1 = image1.shape[2] + h2, w2 = image2.shape[1:3] + aspect = h2 / w2 + + new_w = w1 + new_h = int(w1 * aspect) + + image2 = self._resize(image2, new_w, new_h) + + ch1, ch2 = image1.shape[-1], image2.shape[-1] + if ch1 != ch2: + if ch1 < ch2: + image1 = torch.cat((image1, torch.ones((*image1.shape[:-1], ch2-ch1), device=image1.device)), dim=-1) + else: + image2 = torch.cat((image2, torch.ones((*image2.shape[:-1], ch1-ch2), device=image2.device)), dim=-1) + + if concat_direction == 'right': + result = torch.cat((image1, image2), dim=2) + elif concat_direction == 'bottom': + result = torch.cat((image1, image2), dim=1) + elif concat_direction == 'left': + result = torch.cat((image2, image1), dim=2) + elif concat_direction == 'top': + result = torch.cat((image2, image1), dim=1) + + return (result,) + + def _resize(self, image, width, height): + img = image.movedim(-1, 1) + resized = common_upscale(img, width, height, "lanczos", "disabled") + return resized.movedim(1, -1) + # Node class mappings NODE_CLASS_MAPPINGS = { - "AiLab_LoadImage": AiLab_LoadImage, - "AiLab_Preview": AiLab_Preview, - "AiLab_ImagePreview": AiLab_ImagePreview, - "AiLab_MaskPreview": AiLab_MaskPreview, + "AILab_LoadImage": AILab_LoadImage, + "AILab_Preview": AILab_Preview, + "AILab_ImagePreview": AILab_ImagePreview, + "AILab_MaskPreview": AILab_MaskPreview, + "AILab_ImageMaskConvert": AILab_ImageMaskConvert, + "AILab_MaskEnhancer": AILab_MaskEnhancer, + "AILab_MaskCombiner": AILab_MaskCombiner, + "AILab_ImageCombiner": AILab_ImageCombiner, + "AILab_MaskExtractor": AILab_MaskExtractor, + "AILab_ImageStitch": AILab_ImageStitch, } # Node display name mappings NODE_DISPLAY_NAME_MAPPINGS = { - "AiLab_LoadImage": "Load Image (RMBG) 🖼️", - "AiLab_Preview": "Preview (RMBG) 🖼️🎭", - "AiLab_ImagePreview": "Image Preview (RMBG) 🖼️", - "AiLab_MaskPreview": "Mask Preview (RMBG) 🎭", + "AILab_LoadImage": "Load Image (RMBG) 🖼️", + "AILab_Preview": "Preview (RMBG) 🖼️🎭", + "AILab_ImagePreview": "Image Preview (RMBG) 🖼️", + "AILab_MaskPreview": "Mask Preview (RMBG) 🎭", + "AILab_ImageMaskConvert": "Image/Mask Converter (RMBG) 🖼️🎭", + "AILab_MaskEnhancer": "Mask Enhancer (RMBG) 🎭", + "AILab_MaskCombiner": "Mask Combiner (RMBG) 🎭", + "AILab_ImageCombiner": "Image Combiner (RMBG) 🖼️", + "AILab_MaskExtractor": "Mask Extractor (RMBG) 🎭", + "AILab_ImageStitch": "Image Stitch (RMBG) 🖼️", } \ No newline at end of file diff --git a/AILab_RMBG.py b/AILab_RMBG.py index d82ad28..be84348 100644 --- a/AILab_RMBG.py +++ b/AILab_RMBG.py @@ -1,651 +1,687 @@ -# ComfyUI-RMBG v2.1.1 -# This custom node for ComfyUI provides functionality for background removal using various models, -# including RMBG-2.0, INSPYRENET, BEN, BEN2 and BIREFNET-HR. It leverages deep learning techniques -# to process images and generate masks for background removal. -# -# Models License Notice: -# - RMBG-2.0: Apache-2.0 License (https://huggingface.co/briaai/RMBG-2.0) -# - INSPYRENET: MIT License (https://github.com/plemeri/InSPyReNet) -# - BEN: Apache-2.0 License (https://huggingface.co/PramaLLC/BEN) -# - BEN2: Apache-2.0 License (https://huggingface.co/PramaLLC/BEN2) -# -# This integration script follows GPL-3.0 License. -# When using or modifying this code, please respect both the original model licenses -# and this integration's license terms. -# -# Source: https://github.com/1038lab/ComfyUI-RMBG - -import os -import torch -from PIL import Image -from torchvision import transforms -import numpy as np -import folder_paths -from PIL import ImageFilter -import torch.nn.functional as F -from huggingface_hub import hf_hub_download -import shutil -import sys -import importlib.util -from transformers import AutoModelForImageSegmentation -import cv2 - -device = "cuda" if torch.cuda.is_available() else "cpu" - -# Add model path -folder_paths.add_model_folder_path("rmbg", os.path.join(folder_paths.models_dir, "RMBG")) - -# Model configuration -AVAILABLE_MODELS = { - "RMBG-2.0": { - "type": "rmbg", - "repo_id": "1038lab/RMBG-2.0", - "files": { - "config.json": "config.json", - "model.safetensors": "model.safetensors", - "birefnet.py": "birefnet.py", - "BiRefNet_config.py": "BiRefNet_config.py" - }, - "cache_dir": "RMBG-2.0" - }, - "INSPYRENET": { - "type": "inspyrenet", - "repo_id": "1038lab/inspyrenet", - "files": { - "inspyrenet.safetensors": "inspyrenet.safetensors" - }, - "cache_dir": "INSPYRENET" - }, - "BEN": { - "type": "ben", - "repo_id": "1038lab/BEN", - "files": { - "model.py": "model.py", - "BEN_Base.pth": "BEN_Base.pth" - }, - "cache_dir": "BEN" - }, - "BEN2": { - "type": "ben2", - "repo_id": "1038lab/BEN2", - "files": { - "BEN2_Base.pth": "BEN2_Base.pth", - "BEN2.py": "BEN2.py" - }, - "cache_dir": "BEN2" - } -} - -# Utility functions -def tensor2pil(image): - return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8)) - -def pil2tensor(image): - return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0) - -def handle_model_error(message): - print(f"[RMBG ERROR] {message}") - raise RuntimeError(message) - -class BaseModelLoader: - def __init__(self): - self.model = None - self.current_model_version = None - self.base_cache_dir = os.path.join(folder_paths.models_dir, "RMBG") - - def get_cache_dir(self, model_name): - cache_path = os.path.join(self.base_cache_dir, AVAILABLE_MODELS[model_name]["cache_dir"]) - os.makedirs(cache_path, exist_ok=True) - return cache_path - - def check_model_cache(self, model_name): - model_info = AVAILABLE_MODELS[model_name] - cache_dir = self.get_cache_dir(model_name) - - if not os.path.exists(cache_dir): - return False, "Model directory not found" - - missing_files = [] - for filename in model_info["files"].keys(): - if not os.path.exists(os.path.join(cache_dir, model_info["files"][filename])): - missing_files.append(filename) - - if missing_files: - return False, f"Missing model files: {', '.join(missing_files)}" - - return True, "Model cache verified" - - def download_model(self, model_name): - model_info = AVAILABLE_MODELS[model_name] - cache_dir = self.get_cache_dir(model_name) - - try: - os.makedirs(cache_dir, exist_ok=True) - print(f"Downloading {model_name} model files...") - - for filename in model_info["files"].keys(): - print(f"Downloading {filename}...") - hf_hub_download( - repo_id=model_info["repo_id"], - filename=filename, - local_dir=cache_dir, - local_dir_use_symlinks=False - ) - - return True, "Model files downloaded successfully" - - except Exception as e: - return False, f"Error downloading model files: {str(e)}" - - def clear_model(self): - if self.model is not None: - self.model.cpu() - del self.model - - import gc - gc.collect() - if torch.cuda.is_available(): - torch.cuda.empty_cache() - self.model = None - self.current_model_version = None - -class RMBGModel(BaseModelLoader): - def __init__(self): - super().__init__() - - def load_model(self, model_name): - if self.current_model_version != model_name: - self.clear_model() - - cache_dir = self.get_cache_dir(model_name) - try: - self.model = AutoModelForImageSegmentation.from_pretrained( - cache_dir, - trust_remote_code=True, - local_files_only=True - ) - except Exception as e: - if "'Config' object has no attribute 'get_text_config'" in str(e): - print("[RMBG WARNING] Detected newer transformers version, attempting compatibility mode...") - try: - from transformers import PreTrainedModel - import json - - config_path = os.path.join(cache_dir, "config.json") - with open(config_path, 'r') as f: - config = json.load(f) - - birefnet_path = os.path.join(cache_dir, "birefnet.py") - module_name = f"custom_birefnet_model_{hash(birefnet_path)}" - spec = importlib.util.spec_from_file_location(module_name, birefnet_path) - birefnet_module = importlib.util.module_from_spec(spec) - sys.modules[module_name] = birefnet_module - spec.loader.exec_module(birefnet_module) - - for attr_name in dir(birefnet_module): - attr = getattr(birefnet_module, attr_name) - if isinstance(attr, type) and issubclass(attr, PreTrainedModel) and attr != PreTrainedModel: - config_class = getattr(birefnet_module, "BiRefNet_config", None) - if config_class: - model_config = config_class() - self.model = attr(model_config) - self.model.load_state_dict(torch.load(os.path.join(cache_dir, "model.safetensors"))) - break - - if self.model is None: - raise RuntimeError("Could not find suitable model class") - except Exception as custom_e: - handle_model_error(f"Failed to load model in compatibility mode: {str(custom_e)}\nConsider downgrading transformers to version 4.48.3: pip install transformers==4.48.3") - else: - raise e - - self.model.eval() - for param in self.model.parameters(): - param.requires_grad = False - - torch.set_float32_matmul_precision('high') - self.model.to(device) - self.current_model_version = model_name - - def process_image(self, images, model_name, params): - try: - self.load_model(model_name) - - # Prepare batch processing - transform_image = transforms.Compose([ - transforms.Resize((params["process_res"], params["process_res"])), - transforms.ToTensor(), - transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) - ]) - - # Ensure input is in list format - if isinstance(images, torch.Tensor): - if len(images.shape) == 3: - images = [images] - else: - images = [img for img in images] - - # Store original image sizes - original_sizes = [tensor2pil(img).size for img in images] - - # Batch process transformations - input_tensors = [transform_image(tensor2pil(img)).unsqueeze(0) for img in images] - input_batch = torch.cat(input_tensors, dim=0).to(device) - - with torch.no_grad(): - outputs = self.model(input_batch) - - if isinstance(outputs, list) and len(outputs) > 0: - results = outputs[-1].sigmoid().cpu() - elif isinstance(outputs, dict) and 'logits' in outputs: - results = outputs['logits'].sigmoid().cpu() - elif isinstance(outputs, torch.Tensor): - results = outputs.sigmoid().cpu() - else: - try: - if hasattr(outputs, 'last_hidden_state'): - results = outputs.last_hidden_state.sigmoid().cpu() - else: - for k, v in outputs.items(): - if isinstance(v, torch.Tensor): - results = v.sigmoid().cpu() - break - except: - handle_model_error("Unable to recognize model output format") - - masks = [] - - # Process each result and resize back to original dimensions - for i, (result, (orig_w, orig_h)) in enumerate(zip(results, original_sizes)): - result = result.squeeze() - result = result * (1 + (1 - params["sensitivity"])) - result = torch.clamp(result, 0, 1) - - # Resize back to original dimensions - result = F.interpolate(result.unsqueeze(0).unsqueeze(0), - size=(orig_h, orig_w), - mode='bilinear').squeeze() - - masks.append(tensor2pil(result)) - - return masks - - except Exception as e: - handle_model_error(f"Error in batch processing: {str(e)}") - -class InspyrenetModel(BaseModelLoader): - def __init__(self): - super().__init__() - - def load_model(self, model_name): - if self.current_model_version != model_name: - self.clear_model() - - try: - import transparent_background - self.model = transparent_background.Remover() - self.current_model_version = model_name - except ImportError: - try: - import pip - pip.main(['install', 'transparent_background']) - import transparent_background - self.model = transparent_background.Remover() - self.current_model_version = model_name - except Exception as e: - handle_model_error(f"Failed to install transparent_background: {str(e)}") - - def process_image(self, image, model_name, params): - try: - self.load_model(model_name) - - orig_image = tensor2pil(image) - w, h = orig_image.size - - # Resize for processing - aspect_ratio = h / w - new_w = params["process_res"] - new_h = int(params["process_res"] * aspect_ratio) - resized_image = orig_image.resize((new_w, new_h), Image.LANCZOS) - - # Process image - foreground = self.model.process(resized_image, type='rgba') - foreground = foreground.resize((w, h), Image.LANCZOS) - mask = foreground.split()[-1] - - return mask - - except Exception as e: - handle_model_error(f"Error in Inspyrenet processing: {str(e)}") - -class BENModel(BaseModelLoader): - def __init__(self): - super().__init__() - - def load_model(self, model_name): - if self.current_model_version != model_name: - self.clear_model() - - cache_dir = self.get_cache_dir(model_name) - model_path = os.path.join(cache_dir, "model.py") - module_name = f"custom_ben_model_{hash(model_path)}" - - spec = importlib.util.spec_from_file_location(module_name, model_path) - ben_module = importlib.util.module_from_spec(spec) - sys.modules[module_name] = ben_module - spec.loader.exec_module(ben_module) - - model_weights_path = os.path.join(cache_dir, "BEN_Base.pth") - self.model = ben_module.BEN_Base() - self.model.loadcheckpoints(model_weights_path) - - self.model.eval() - for param in self.model.parameters(): - param.requires_grad = False - - torch.set_float32_matmul_precision('high') - self.model.to(device) - self.current_model_version = model_name - - def process_image(self, image, model_name, params): - try: - self.load_model(model_name) - - orig_image = tensor2pil(image) - w, h = orig_image.size - - aspect_ratio = h / w - new_w = params["process_res"] - new_h = int(params["process_res"] * aspect_ratio) - resized_image = orig_image.resize((new_w, new_h), Image.LANCZOS) - - processed_input = resized_image.convert("RGBA") - - with torch.no_grad(): - _, foreground = self.model.inference(processed_input) - - foreground = foreground.resize((w, h), Image.LANCZOS) - mask = foreground.split()[-1] - - return mask - - except Exception as e: - handle_model_error(f"Error in BEN processing: {str(e)}") - -class BEN2Model(BaseModelLoader): - def __init__(self): - super().__init__() - - def load_model(self, model_name): - if self.current_model_version != model_name: - self.clear_model() - - try: - cache_dir = self.get_cache_dir(model_name) - model_path = os.path.join(cache_dir, "BEN2.py") - module_name = f"custom_ben2_model_{hash(model_path)}" - - spec = importlib.util.spec_from_file_location(module_name, model_path) - ben2_module = importlib.util.module_from_spec(spec) - sys.modules[module_name] = ben2_module - spec.loader.exec_module(ben2_module) - - model_weights_path = os.path.join(cache_dir, "BEN2_Base.pth") - self.model = ben2_module.BEN_Base() - self.model.loadcheckpoints(model_weights_path) - - self.model.eval() - for param in self.model.parameters(): - param.requires_grad = False - - torch.set_float32_matmul_precision('high') - self.model.to(device) - self.current_model_version = model_name - - except Exception as e: - handle_model_error(f"Error loading BEN2 model: {str(e)}") - - def process_image(self, images, model_name, params): - try: - self.load_model(model_name) - - if isinstance(images, torch.Tensor): - if len(images.shape) == 3: - images = [images] - else: - images = [img for img in images] - - batch_size = 3 - all_masks = [] - - for i in range(0, len(images), batch_size): - batch_images = images[i:i + batch_size] - batch_pil_images = [] - original_sizes = [] - - for img in batch_images: - orig_image = tensor2pil(img) - w, h = orig_image.size - original_sizes.append((w, h)) - - aspect_ratio = h / w - new_w = params["process_res"] - new_h = int(params["process_res"] * aspect_ratio) - resized_image = orig_image.resize((new_w, new_h), Image.LANCZOS) - processed_input = resized_image.convert("RGBA") - batch_pil_images.append(processed_input) - - with torch.no_grad(): - try: - foregrounds = self.model.inference(batch_pil_images) - if not isinstance(foregrounds, list): - foregrounds = [foregrounds] - except Exception as e: - handle_model_error(f"Error in BEN2 inference: {str(e)}") - - for foreground, (orig_w, orig_h) in zip(foregrounds, original_sizes): - foreground = foreground.resize((orig_w, orig_h), Image.LANCZOS) - mask = foreground.split()[-1] - all_masks.append(mask) - - if len(all_masks) == 1: - return all_masks[0] - return all_masks - - except Exception as e: - handle_model_error(f"Error in BEN2 processing: {str(e)}") - -def refine_foreground(image_bchw, masks_b1hw): - b, c, h, w = image_bchw.shape - if b != masks_b1hw.shape[0]: - raise ValueError("images and masks must have the same batch size") - - image_np = image_bchw.cpu().numpy() - mask_np = masks_b1hw.cpu().numpy() - - refined_fg = [] - for i in range(b): - mask = mask_np[i, 0] - thresh = 0.45 - mask_binary = (mask > thresh).astype(np.float32) - - edge_blur = cv2.GaussianBlur(mask_binary, (3, 3), 0) - transition_mask = np.logical_and(mask > 0.05, mask < 0.95) - - alpha = 0.85 - mask_refined = np.where(transition_mask, - alpha * mask + (1-alpha) * edge_blur, - mask_binary) - - edge_region = np.logical_and(mask > 0.2, mask < 0.8) - mask_refined = np.where(edge_region, - mask_refined * 0.98, - mask_refined) - - result = [] - for c in range(image_np.shape[1]): - channel = image_np[i, c] - refined = channel * mask_refined - result.append(refined) - - refined_fg.append(np.stack(result)) - - return torch.from_numpy(np.stack(refined_fg)) - -class RMBG: - def __init__(self): - self.models = { - "RMBG-2.0": RMBGModel(), - "INSPYRENET": InspyrenetModel(), - "BEN": BENModel(), - "BEN2": BEN2Model() - } - - @classmethod - def INPUT_TYPES(s): - tooltips = { - "image": "Input image to be processed for background removal.", - "model": "Select the background removal model to use (RMBG-2.0, INSPYRENET, BEN).", - "sensitivity": "Adjust the strength of mask detection (higher values result in more aggressive detection).", - "process_res": "Set the processing resolution (higher values require more VRAM and may increase processing time).", - "mask_blur": "Specify the amount of blur to apply to the mask edges (0 for no blur, higher values for more blur).", - "mask_offset": "Adjust the mask boundary (positive values expand the mask, negative values shrink it).", - "background": "Choose the background color for the final output (Alpha for transparent background).", - "invert_output": "Enable to invert both the image and mask output (useful for certain effects).", - "optimize": "Enable model optimization for faster processing (may affect output quality).", - "refine_foreground": "Use Fast Foreground Colour Estimation to optimize transparent background" - } - - return { - "required": { - "image": ("IMAGE", {"tooltip": tooltips["image"]}), - "model": (list(AVAILABLE_MODELS.keys()), {"tooltip": tooltips["model"]}), - }, - "optional": { - "sensitivity": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": tooltips["sensitivity"]}), - "process_res": ("INT", {"default": 1024, "min": 256, "max": 2048, "step": 8, "tooltip": tooltips["process_res"]}), - "mask_blur": ("INT", {"default": 0, "min": 0, "max": 64, "step": 1, "tooltip": tooltips["mask_blur"]}), - "mask_offset": ("INT", {"default": 0, "min": -64, "max": 64, "step": 1, "tooltip": tooltips["mask_offset"]}), - "background": (["Alpha", "black", "white", "gray", "green", "blue", "red"], {"default": "Alpha", "tooltip": tooltips["background"]}), - "invert_output": ("BOOLEAN", {"default": False, "tooltip": tooltips["invert_output"]}), - "optimize": (["default", "on"], {"default": "default", "tooltip": tooltips["optimize"]}), - "refine_foreground": ("BOOLEAN", {"default": False, "tooltip": tooltips["refine_foreground"]}) - } - } - - RETURN_TYPES = ("IMAGE", "MASK", "IMAGE") - RETURN_NAMES = ("IMAGE", "MASK", "MASK_IMAGE") - FUNCTION = "process_image" - CATEGORY = "🧪AILab/🧽RMBG" - - def process_image(self, image, model, **params): - try: - processed_images = [] - processed_masks = [] - - bg_colors = { - "Alpha": None, - "black": (0, 0, 0), - "white": (255, 255, 255), - "gray": (128, 128, 128), - "green": (0, 255, 0), - "blue": (0, 0, 255), - "red": (255, 0, 0) - } - - model_instance = self.models[model] - - # Check and download model if needed - cache_status, message = model_instance.check_model_cache(model) - if not cache_status: - print(f"Cache check: {message}") - print("Downloading required model files...") - download_status, download_message = model_instance.download_model(model) - if not download_status: - handle_model_error(download_message) - print("Model files downloaded successfully") - - for img in image: - # Get mask from specific model - mask = model_instance.process_image(img, model, params) - - # Ensure mask is in the correct format - if isinstance(mask, list): - masks = [m.convert("L") for m in mask if isinstance(m, Image.Image)] - mask = masks[0] if masks else None - elif isinstance(mask, Image.Image): - mask = mask.convert("L") - - # Post-process mask - mask_tensor = pil2tensor(mask) - mask_tensor = mask_tensor * (1 + (1 - params["sensitivity"])) - mask_tensor = torch.clamp(mask_tensor, 0, 1) - mask = tensor2pil(mask_tensor) - - if params["mask_blur"] > 0: - mask = mask.filter(ImageFilter.GaussianBlur(radius=params["mask_blur"])) - - if params["mask_offset"] != 0: - if params["mask_offset"] > 0: - for _ in range(params["mask_offset"]): - mask = mask.filter(ImageFilter.MaxFilter(3)) - else: - for _ in range(-params["mask_offset"]): - mask = mask.filter(ImageFilter.MinFilter(3)) - - if params["invert_output"]: - mask = Image.fromarray(255 - np.array(mask)) - - # Convert to tensors for refine_foreground - img_tensor = torch.from_numpy(np.array(tensor2pil(img))).permute(2, 0, 1).unsqueeze(0) / 255.0 - mask_tensor = torch.from_numpy(np.array(mask)).unsqueeze(0).unsqueeze(0) / 255.0 - - # Create final image - orig_image = tensor2pil(img) - - if params.get("refine_foreground", False): - refined_fg = refine_foreground(img_tensor, mask_tensor) - refined_fg = tensor2pil(refined_fg[0].permute(1, 2, 0)) - r, g, b = refined_fg.split() - foreground = Image.merge('RGBA', (r, g, b, mask)) - else: - orig_rgba = orig_image.convert("RGBA") - r, g, b, _ = orig_rgba.split() - foreground = Image.merge('RGBA', (r, g, b, mask)) - - if params["background"] != "Alpha": - bg_color = bg_colors[params["background"]] - bg_image = Image.new('RGBA', orig_image.size, (*bg_color, 255)) - composite_image = Image.alpha_composite(bg_image, foreground) - processed_images.append(pil2tensor(composite_image.convert("RGB"))) - else: - processed_images.append(pil2tensor(foreground)) - - processed_masks.append(pil2tensor(mask)) - - # Create mask image for visualization - mask_images = [] - for mask_tensor in processed_masks: - # Convert mask to RGB image format for visualization - mask_image = mask_tensor.reshape((-1, 1, mask_tensor.shape[-2], mask_tensor.shape[-1])).movedim(1, -1).expand(-1, -1, -1, 3) - mask_images.append(mask_image) - - mask_image_output = torch.cat(mask_images, dim=0) - - return (torch.cat(processed_images, dim=0), torch.cat(processed_masks, dim=0), mask_image_output) - - except Exception as e: - handle_model_error(f"Error in image processing: {str(e)}") - # Return original image and empty mask on error - empty_mask = torch.zeros((image.shape[0], image.shape[2], image.shape[3])) - empty_mask_image = empty_mask.reshape((-1, 1, empty_mask.shape[-2], empty_mask.shape[-1])).movedim(1, -1).expand(-1, -1, -1, 3) - return (image, empty_mask, empty_mask_image) - -# Node Mapping -NODE_CLASS_MAPPINGS = { - "RMBG": RMBG -} - -NODE_DISPLAY_NAME_MAPPINGS = { - "RMBG": "Remove Background (RMBG)" +# ComfyUI-RMBG v2.2.0 +# This custom node for ComfyUI provides functionality for background removal using various models, +# including RMBG-2.0, INSPYRENET, BEN, BEN2 and BIREFNET-HR. It leverages deep learning techniques +# to process images and generate masks for background removal. +# +# Models License Notice: +# - RMBG-2.0: Apache-2.0 License (https://huggingface.co/briaai/RMBG-2.0) +# - INSPYRENET: MIT License (https://github.com/plemeri/InSPyReNet) +# - BEN: Apache-2.0 License (https://huggingface.co/PramaLLC/BEN) +# - BEN2: Apache-2.0 License (https://huggingface.co/PramaLLC/BEN2) +# +# This integration script follows GPL-3.0 License. +# When using or modifying this code, please respect both the original model licenses +# and this integration's license terms. +# +# Source: https://github.com/1038lab/ComfyUI-RMBG + +import os +import torch +from PIL import Image +from torchvision import transforms +import numpy as np +import folder_paths +from PIL import ImageFilter +import torch.nn.functional as F +from huggingface_hub import hf_hub_download +import shutil +import sys +import importlib.util +from transformers import AutoModelForImageSegmentation +import cv2 +import types + +device = "cuda" if torch.cuda.is_available() else "cpu" + +# Add model path +folder_paths.add_model_folder_path("rmbg", os.path.join(folder_paths.models_dir, "RMBG")) + +# Model configuration +AVAILABLE_MODELS = { + "RMBG-2.0": { + "type": "rmbg", + "repo_id": "1038lab/RMBG-2.0", + "files": { + "config.json": "config.json", + "model.safetensors": "model.safetensors", + "birefnet.py": "birefnet.py", + "BiRefNet_config.py": "BiRefNet_config.py" + }, + "cache_dir": "RMBG-2.0" + }, + "INSPYRENET": { + "type": "inspyrenet", + "repo_id": "1038lab/inspyrenet", + "files": { + "inspyrenet.safetensors": "inspyrenet.safetensors" + }, + "cache_dir": "INSPYRENET" + }, + "BEN": { + "type": "ben", + "repo_id": "1038lab/BEN", + "files": { + "model.py": "model.py", + "BEN_Base.pth": "BEN_Base.pth" + }, + "cache_dir": "BEN" + }, + "BEN2": { + "type": "ben2", + "repo_id": "1038lab/BEN2", + "files": { + "BEN2_Base.pth": "BEN2_Base.pth", + "BEN2.py": "BEN2.py" + }, + "cache_dir": "BEN2" + } +} + +# Utility functions +def tensor2pil(image): + return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8)) + +def pil2tensor(image): + return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0) + +def handle_model_error(message): + print(f"[RMBG ERROR] {message}") + raise RuntimeError(message) + +class BaseModelLoader: + def __init__(self): + self.model = None + self.current_model_version = None + self.base_cache_dir = os.path.join(folder_paths.models_dir, "RMBG") + + def get_cache_dir(self, model_name): + cache_path = os.path.join(self.base_cache_dir, AVAILABLE_MODELS[model_name]["cache_dir"]) + os.makedirs(cache_path, exist_ok=True) + return cache_path + + def check_model_cache(self, model_name): + model_info = AVAILABLE_MODELS[model_name] + cache_dir = self.get_cache_dir(model_name) + + if not os.path.exists(cache_dir): + return False, "Model directory not found" + + missing_files = [] + for filename in model_info["files"].keys(): + if not os.path.exists(os.path.join(cache_dir, model_info["files"][filename])): + missing_files.append(filename) + + if missing_files: + return False, f"Missing model files: {', '.join(missing_files)}" + + return True, "Model cache verified" + + def download_model(self, model_name): + model_info = AVAILABLE_MODELS[model_name] + cache_dir = self.get_cache_dir(model_name) + + try: + os.makedirs(cache_dir, exist_ok=True) + print(f"Downloading {model_name} model files...") + + for filename in model_info["files"].keys(): + print(f"Downloading {filename}...") + hf_hub_download( + repo_id=model_info["repo_id"], + filename=filename, + local_dir=cache_dir, + local_dir_use_symlinks=False + ) + + return True, "Model files downloaded successfully" + + except Exception as e: + return False, f"Error downloading model files: {str(e)}" + + def clear_model(self): + if self.model is not None: + self.model.cpu() + del self.model + + import gc + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + self.model = None + self.current_model_version = None + +class RMBGModel(BaseModelLoader): + def __init__(self): + super().__init__() + + def load_model(self, model_name): + if self.current_model_version != model_name: + self.clear_model() + + cache_dir = self.get_cache_dir(model_name) + try: + # Try standard loading first + try: + self.model = AutoModelForImageSegmentation.from_pretrained( + cache_dir, + trust_remote_code=True, + local_files_only=True + ) + except AttributeError as ae: + if "'Config' object has no attribute 'get_text_config'" in str(ae): + print("[RMBG WARNING] Detected newer transformers version, using compatibility mode...") + try: + from transformers import PreTrainedModel + import json + + config_path = os.path.join(cache_dir, "config.json") + with open(config_path, 'r') as f: + config = json.load(f) + + birefnet_path = os.path.join(cache_dir, "birefnet.py") + BiRefNetConfig_path = os.path.join(cache_dir, "BiRefNet_config.py") + + # Load the BiRefNetConfig + config_spec = importlib.util.spec_from_file_location("BiRefNetConfig", BiRefNetConfig_path) + config_module = importlib.util.module_from_spec(config_spec) + sys.modules["BiRefNetConfig"] = config_module + config_spec.loader.exec_module(config_module) + + # Fix and load birefnet module + with open(birefnet_path, 'r') as f: + birefnet_content = f.read() + + birefnet_content = birefnet_content.replace( + "from .BiRefNet_config import BiRefNetConfig", + "from BiRefNetConfig import BiRefNetConfig" + ) + + module_name = f"custom_birefnet_model_{hash(birefnet_path)}" + module = types.ModuleType(module_name) + sys.modules[module_name] = module + exec(birefnet_content, module.__dict__) + + for attr_name in dir(module): + attr = getattr(module, attr_name) + if isinstance(attr, type) and issubclass(attr, PreTrainedModel) and attr != PreTrainedModel: + BiRefNetConfig = getattr(config_module, "BiRefNetConfig") + model_config = BiRefNetConfig() + self.model = attr(model_config) + + weights_path = os.path.join(cache_dir, "model.safetensors") + try: + try: + import safetensors.torch + self.model.load_state_dict(safetensors.torch.load_file(weights_path)) + except ImportError: + from transformers.modeling_utils import load_state_dict + state_dict = load_state_dict(weights_path) + self.model.load_state_dict(state_dict) + except Exception as load_error: + pytorch_weights = os.path.join(cache_dir, "pytorch_model.bin") + if os.path.exists(pytorch_weights): + self.model.load_state_dict(torch.load(pytorch_weights, map_location="cpu")) + else: + raise RuntimeError(f"Failed to load weights: {str(load_error)}") + break + + if self.model is None: + raise RuntimeError("Could not find suitable model class") + + except Exception as custom_e: + handle_model_error(f"Failed to load model in compatibility mode: {str(custom_e)}\nConsider downgrading transformers to version 4.48.3: pip install transformers==4.48.3") + else: + raise ae + except Exception as e: + handle_model_error(f"Error loading model: {str(e)}") + + self.model.eval() + for param in self.model.parameters(): + param.requires_grad = False + + torch.set_float32_matmul_precision('high') + self.model.to(device) + self.current_model_version = model_name + + def process_image(self, images, model_name, params): + try: + self.load_model(model_name) + + # Prepare batch processing + transform_image = transforms.Compose([ + transforms.Resize((params["process_res"], params["process_res"])), + transforms.ToTensor(), + transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) + ]) + + # Ensure input is in list format + if isinstance(images, torch.Tensor): + if len(images.shape) == 3: + images = [images] + else: + images = [img for img in images] + + # Store original image sizes + original_sizes = [tensor2pil(img).size for img in images] + + # Batch process transformations + input_tensors = [transform_image(tensor2pil(img)).unsqueeze(0) for img in images] + input_batch = torch.cat(input_tensors, dim=0).to(device) + + with torch.no_grad(): + outputs = self.model(input_batch) + + if isinstance(outputs, list) and len(outputs) > 0: + results = outputs[-1].sigmoid().cpu() + elif isinstance(outputs, dict) and 'logits' in outputs: + results = outputs['logits'].sigmoid().cpu() + elif isinstance(outputs, torch.Tensor): + results = outputs.sigmoid().cpu() + else: + try: + if hasattr(outputs, 'last_hidden_state'): + results = outputs.last_hidden_state.sigmoid().cpu() + else: + for k, v in outputs.items(): + if isinstance(v, torch.Tensor): + results = v.sigmoid().cpu() + break + except: + handle_model_error("Unable to recognize model output format") + + masks = [] + + # Process each result and resize back to original dimensions + for i, (result, (orig_w, orig_h)) in enumerate(zip(results, original_sizes)): + result = result.squeeze() + result = result * (1 + (1 - params["sensitivity"])) + result = torch.clamp(result, 0, 1) + + # Resize back to original dimensions + result = F.interpolate(result.unsqueeze(0).unsqueeze(0), + size=(orig_h, orig_w), + mode='bilinear').squeeze() + + masks.append(tensor2pil(result)) + + return masks + + except Exception as e: + handle_model_error(f"Error in batch processing: {str(e)}") + +class InspyrenetModel(BaseModelLoader): + def __init__(self): + super().__init__() + + def load_model(self, model_name): + if self.current_model_version != model_name: + self.clear_model() + + try: + import transparent_background + self.model = transparent_background.Remover() + self.current_model_version = model_name + except ImportError: + try: + import pip + pip.main(['install', 'transparent_background']) + import transparent_background + self.model = transparent_background.Remover() + self.current_model_version = model_name + except Exception as e: + handle_model_error(f"Failed to install transparent_background: {str(e)}") + + def process_image(self, image, model_name, params): + try: + self.load_model(model_name) + + orig_image = tensor2pil(image) + w, h = orig_image.size + + # Resize for processing + aspect_ratio = h / w + new_w = params["process_res"] + new_h = int(params["process_res"] * aspect_ratio) + resized_image = orig_image.resize((new_w, new_h), Image.LANCZOS) + + # Process image + foreground = self.model.process(resized_image, type='rgba') + foreground = foreground.resize((w, h), Image.LANCZOS) + mask = foreground.split()[-1] + + return mask + + except Exception as e: + handle_model_error(f"Error in Inspyrenet processing: {str(e)}") + +class BENModel(BaseModelLoader): + def __init__(self): + super().__init__() + + def load_model(self, model_name): + if self.current_model_version != model_name: + self.clear_model() + + cache_dir = self.get_cache_dir(model_name) + model_path = os.path.join(cache_dir, "model.py") + module_name = f"custom_ben_model_{hash(model_path)}" + + spec = importlib.util.spec_from_file_location(module_name, model_path) + ben_module = importlib.util.module_from_spec(spec) + sys.modules[module_name] = ben_module + spec.loader.exec_module(ben_module) + + model_weights_path = os.path.join(cache_dir, "BEN_Base.pth") + self.model = ben_module.BEN_Base() + self.model.loadcheckpoints(model_weights_path) + + self.model.eval() + for param in self.model.parameters(): + param.requires_grad = False + + torch.set_float32_matmul_precision('high') + self.model.to(device) + self.current_model_version = model_name + + def process_image(self, image, model_name, params): + try: + self.load_model(model_name) + + orig_image = tensor2pil(image) + w, h = orig_image.size + + aspect_ratio = h / w + new_w = params["process_res"] + new_h = int(params["process_res"] * aspect_ratio) + resized_image = orig_image.resize((new_w, new_h), Image.LANCZOS) + + processed_input = resized_image.convert("RGBA") + + with torch.no_grad(): + _, foreground = self.model.inference(processed_input) + + foreground = foreground.resize((w, h), Image.LANCZOS) + mask = foreground.split()[-1] + + return mask + + except Exception as e: + handle_model_error(f"Error in BEN processing: {str(e)}") + +class BEN2Model(BaseModelLoader): + def __init__(self): + super().__init__() + + def load_model(self, model_name): + if self.current_model_version != model_name: + self.clear_model() + + try: + cache_dir = self.get_cache_dir(model_name) + model_path = os.path.join(cache_dir, "BEN2.py") + module_name = f"custom_ben2_model_{hash(model_path)}" + + spec = importlib.util.spec_from_file_location(module_name, model_path) + ben2_module = importlib.util.module_from_spec(spec) + sys.modules[module_name] = ben2_module + spec.loader.exec_module(ben2_module) + + model_weights_path = os.path.join(cache_dir, "BEN2_Base.pth") + self.model = ben2_module.BEN_Base() + self.model.loadcheckpoints(model_weights_path) + + self.model.eval() + for param in self.model.parameters(): + param.requires_grad = False + + torch.set_float32_matmul_precision('high') + self.model.to(device) + self.current_model_version = model_name + + except Exception as e: + handle_model_error(f"Error loading BEN2 model: {str(e)}") + + def process_image(self, images, model_name, params): + try: + self.load_model(model_name) + + if isinstance(images, torch.Tensor): + if len(images.shape) == 3: + images = [images] + else: + images = [img for img in images] + + batch_size = 3 + all_masks = [] + + for i in range(0, len(images), batch_size): + batch_images = images[i:i + batch_size] + batch_pil_images = [] + original_sizes = [] + + for img in batch_images: + orig_image = tensor2pil(img) + w, h = orig_image.size + original_sizes.append((w, h)) + + aspect_ratio = h / w + new_w = params["process_res"] + new_h = int(params["process_res"] * aspect_ratio) + resized_image = orig_image.resize((new_w, new_h), Image.LANCZOS) + processed_input = resized_image.convert("RGBA") + batch_pil_images.append(processed_input) + + with torch.no_grad(): + try: + foregrounds = self.model.inference(batch_pil_images) + if not isinstance(foregrounds, list): + foregrounds = [foregrounds] + except Exception as e: + handle_model_error(f"Error in BEN2 inference: {str(e)}") + + for foreground, (orig_w, orig_h) in zip(foregrounds, original_sizes): + foreground = foreground.resize((orig_w, orig_h), Image.LANCZOS) + mask = foreground.split()[-1] + all_masks.append(mask) + + if len(all_masks) == 1: + return all_masks[0] + return all_masks + + except Exception as e: + handle_model_error(f"Error in BEN2 processing: {str(e)}") + +def refine_foreground(image_bchw, masks_b1hw): + b, c, h, w = image_bchw.shape + if b != masks_b1hw.shape[0]: + raise ValueError("images and masks must have the same batch size") + + image_np = image_bchw.cpu().numpy() + mask_np = masks_b1hw.cpu().numpy() + + refined_fg = [] + for i in range(b): + mask = mask_np[i, 0] + thresh = 0.45 + mask_binary = (mask > thresh).astype(np.float32) + + edge_blur = cv2.GaussianBlur(mask_binary, (3, 3), 0) + transition_mask = np.logical_and(mask > 0.05, mask < 0.95) + + alpha = 0.85 + mask_refined = np.where(transition_mask, + alpha * mask + (1-alpha) * edge_blur, + mask_binary) + + edge_region = np.logical_and(mask > 0.2, mask < 0.8) + mask_refined = np.where(edge_region, + mask_refined * 0.98, + mask_refined) + + result = [] + for c in range(image_np.shape[1]): + channel = image_np[i, c] + refined = channel * mask_refined + result.append(refined) + + refined_fg.append(np.stack(result)) + + return torch.from_numpy(np.stack(refined_fg)) + +class RMBG: + def __init__(self): + self.models = { + "RMBG-2.0": RMBGModel(), + "INSPYRENET": InspyrenetModel(), + "BEN": BENModel(), + "BEN2": BEN2Model() + } + + @classmethod + def INPUT_TYPES(s): + tooltips = { + "image": "Input image to be processed for background removal.", + "model": "Select the background removal model to use (RMBG-2.0, INSPYRENET, BEN).", + "sensitivity": "Adjust the strength of mask detection (higher values result in more aggressive detection).", + "process_res": "Set the processing resolution (higher values require more VRAM and may increase processing time).", + "mask_blur": "Specify the amount of blur to apply to the mask edges (0 for no blur, higher values for more blur).", + "mask_offset": "Adjust the mask boundary (positive values expand the mask, negative values shrink it).", + "background": "Choose the background color for the final output (Alpha for transparent background).", + "invert_output": "Enable to invert both the image and mask output (useful for certain effects).", + "optimize": "Enable model optimization for faster processing (may affect output quality).", + "refine_foreground": "Use Fast Foreground Colour Estimation to optimize transparent background" + } + + return { + "required": { + "image": ("IMAGE", {"tooltip": tooltips["image"]}), + "model": (list(AVAILABLE_MODELS.keys()), {"tooltip": tooltips["model"]}), + }, + "optional": { + "sensitivity": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": tooltips["sensitivity"]}), + "process_res": ("INT", {"default": 1024, "min": 256, "max": 2048, "step": 8, "tooltip": tooltips["process_res"]}), + "mask_blur": ("INT", {"default": 0, "min": 0, "max": 64, "step": 1, "tooltip": tooltips["mask_blur"]}), + "mask_offset": ("INT", {"default": 0, "min": -64, "max": 64, "step": 1, "tooltip": tooltips["mask_offset"]}), + "background": (["Alpha", "black", "white", "gray", "green", "blue", "red"], {"default": "Alpha", "tooltip": tooltips["background"]}), + "invert_output": ("BOOLEAN", {"default": False, "tooltip": tooltips["invert_output"]}), + "optimize": (["default", "on"], {"default": "default", "tooltip": tooltips["optimize"]}), + "refine_foreground": ("BOOLEAN", {"default": False, "tooltip": tooltips["refine_foreground"]}) + } + } + + RETURN_TYPES = ("IMAGE", "MASK", "IMAGE") + RETURN_NAMES = ("IMAGE", "MASK", "MASK_IMAGE") + FUNCTION = "process_image" + CATEGORY = "🧪AILab/🧽RMBG" + + def process_image(self, image, model, **params): + try: + processed_images = [] + processed_masks = [] + + bg_colors = { + "Alpha": None, + "black": (0, 0, 0), + "white": (255, 255, 255), + "gray": (128, 128, 128), + "green": (0, 255, 0), + "blue": (0, 0, 255), + "red": (255, 0, 0) + } + + model_instance = self.models[model] + + # Check and download model if needed + cache_status, message = model_instance.check_model_cache(model) + if not cache_status: + print(f"Cache check: {message}") + print("Downloading required model files...") + download_status, download_message = model_instance.download_model(model) + if not download_status: + handle_model_error(download_message) + print("Model files downloaded successfully") + + for img in image: + # Get mask from specific model + mask = model_instance.process_image(img, model, params) + + # Ensure mask is in the correct format + if isinstance(mask, list): + masks = [m.convert("L") for m in mask if isinstance(m, Image.Image)] + mask = masks[0] if masks else None + elif isinstance(mask, Image.Image): + mask = mask.convert("L") + + # Post-process mask + mask_tensor = pil2tensor(mask) + mask_tensor = mask_tensor * (1 + (1 - params["sensitivity"])) + mask_tensor = torch.clamp(mask_tensor, 0, 1) + mask = tensor2pil(mask_tensor) + + if params["mask_blur"] > 0: + mask = mask.filter(ImageFilter.GaussianBlur(radius=params["mask_blur"])) + + if params["mask_offset"] != 0: + if params["mask_offset"] > 0: + for _ in range(params["mask_offset"]): + mask = mask.filter(ImageFilter.MaxFilter(3)) + else: + for _ in range(-params["mask_offset"]): + mask = mask.filter(ImageFilter.MinFilter(3)) + + if params["invert_output"]: + mask = Image.fromarray(255 - np.array(mask)) + + # Convert to tensors for refine_foreground + img_tensor = torch.from_numpy(np.array(tensor2pil(img))).permute(2, 0, 1).unsqueeze(0) / 255.0 + mask_tensor = torch.from_numpy(np.array(mask)).unsqueeze(0).unsqueeze(0) / 255.0 + + # Create final image + orig_image = tensor2pil(img) + + if params.get("refine_foreground", False): + refined_fg = refine_foreground(img_tensor, mask_tensor) + refined_fg = tensor2pil(refined_fg[0].permute(1, 2, 0)) + r, g, b = refined_fg.split() + foreground = Image.merge('RGBA', (r, g, b, mask)) + else: + orig_rgba = orig_image.convert("RGBA") + r, g, b, _ = orig_rgba.split() + foreground = Image.merge('RGBA', (r, g, b, mask)) + + if params["background"] != "Alpha": + bg_color = bg_colors[params["background"]] + bg_image = Image.new('RGBA', orig_image.size, (*bg_color, 255)) + composite_image = Image.alpha_composite(bg_image, foreground) + processed_images.append(pil2tensor(composite_image.convert("RGB"))) + else: + processed_images.append(pil2tensor(foreground)) + + processed_masks.append(pil2tensor(mask)) + + # Create mask image for visualization + mask_images = [] + for mask_tensor in processed_masks: + # Convert mask to RGB image format for visualization + mask_image = mask_tensor.reshape((-1, 1, mask_tensor.shape[-2], mask_tensor.shape[-1])).movedim(1, -1).expand(-1, -1, -1, 3) + mask_images.append(mask_image) + + mask_image_output = torch.cat(mask_images, dim=0) + + return (torch.cat(processed_images, dim=0), torch.cat(processed_masks, dim=0), mask_image_output) + + except Exception as e: + handle_model_error(f"Error in image processing: {str(e)}") + # Return original image and empty mask on error + empty_mask = torch.zeros((image.shape[0], image.shape[2], image.shape[3])) + empty_mask_image = empty_mask.reshape((-1, 1, empty_mask.shape[-2], empty_mask.shape[-1])).movedim(1, -1).expand(-1, -1, -1, 3) + return (image, empty_mask, empty_mask_image) + +# Node Mapping +NODE_CLASS_MAPPINGS = { + "RMBG": RMBG +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "RMBG": "Remove Background (RMBG)" } \ No newline at end of file diff --git a/requirements.txt b/requirements.txt index 0b87111..4b01226 100644 --- a/requirements.txt +++ b/requirements.txt @@ -4,12 +4,13 @@ torchvision>=0.15.0 Pillow>=9.0.0 numpy>=1.22.0 huggingface-hub>=0.19.0 -# Note: We recommend transformers versions between 4.35.0 and 4.48.3, but higher versions are now supported. -# If you encounter issues, you can try: pip install transformers==4.48.3 transformers>=4.35.0 +safetensors>=0.3.0 transparent-background>=1.2.4 tqdm>=4.65.0 segment-anything>=1.0 groundingdino-py>=0.4.0 opencv-python>=4.7.0 -scipy>=1.10.0 \ No newline at end of file +scipy>=1.10.0 +onnxruntime>=1.15.0 +onnxruntime-gpu>=1.15.0 \ No newline at end of file