diff --git a/AILab_BiRefNet.py b/AILab_BiRefNet.py index 0085de2..b4de64b 100644 --- a/AILab_BiRefNet.py +++ b/AILab_BiRefNet.py @@ -1,10 +1,14 @@ -# ComfyUI-RMBG v1.9.2 +# ComfyUI-RMBG v2.0.0 # This custom node for ComfyUI provides functionality for background removal using BiRefNet models. # # Model License Notice: # - BiRefNet Models: Apache-2.0 License (https://huggingface.co/ZhengPeng7) # # 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/AILab-AI/ComfyUI-RMBG import os import torch @@ -316,7 +320,7 @@ class BiRefNetModel: except Exception as e: handle_model_error(f"Error in BiRefNet processing: {str(e)}") -class BiRefNet: +class BiRefNetRMBG: def __init__(self): self.model = BiRefNetModel() @@ -347,7 +351,7 @@ class BiRefNet: } RETURN_TYPES = ("IMAGE", "MASK") - RETURN_NAMES = ("image", "mask") + RETURN_NAMES = ("IMAGE", "MASK") FUNCTION = "process_image" CATEGORY = "🧪AILab/🧽RMBG" @@ -450,9 +454,9 @@ class BiRefNet: # Node Mapping NODE_CLASS_MAPPINGS = { - "BiRefNet": BiRefNet + "BiRefNetRMBG": BiRefNetRMBG } NODE_DISPLAY_NAME_MAPPINGS = { - "BiRefNet": "BiRefNet (RMBG)" + "BiRefNetRMBG": "BiRefNet (RMBG)" } \ No newline at end of file diff --git a/AILab_BodySegment.py b/AILab_BodySegment.py new file mode 100644 index 0000000..5ae34db --- /dev/null +++ b/AILab_BodySegment.py @@ -0,0 +1,238 @@ +# ComfyUI-RMBG +# This custom node for ComfyUI provides functionality for background removal using various models, +# including RMBG-2.0, INSPYRENET, and BEN. It leverages deep learning techniques +# to process images and generate masks for background removal. +# +# 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/AILab-AI/ComfyUI-RMBG + + +import os +import torch +import torch.nn as nn +import numpy as np +from typing import Tuple, Union +from PIL import Image, ImageFilter +import onnxruntime +import folder_paths +from huggingface_hub import hf_hub_download +import shutil +from torchvision import transforms + +def pil2tensor(image: Image.Image) -> torch.Tensor: + return torch.from_numpy(np.array(image).astype(np.float32) / 255.0)[None,] + +def tensor2pil(image: torch.Tensor) -> Image.Image: + return Image.fromarray(np.clip(255. * image.cpu().numpy(), 0, 255).astype(np.uint8)) + +def image2mask(image: Image.Image) -> torch.Tensor: + if isinstance(image, Image.Image): + image = pil2tensor(image) + return image.squeeze()[..., 0] + +def mask2image(mask: torch.Tensor) -> Image.Image: + if len(mask.shape) == 2: + mask = mask.unsqueeze(0) + return tensor2pil(mask) + +def RGB2RGBA(image: Image.Image, mask: Union[Image.Image, torch.Tensor]) -> Image.Image: + if isinstance(mask, torch.Tensor): + mask = mask2image(mask) + if mask.size != image.size: + mask = mask.resize(image.size, Image.Resampling.LANCZOS) + return Image.merge('RGBA', (*image.convert('RGB').split(), mask.convert('L'))) + +device = "cuda" if torch.cuda.is_available() else "cpu" + +folder_paths.add_model_folder_path("rmbg", os.path.join(folder_paths.models_dir, "RMBG")) + +class BodySegment: + def __init__(self): + self.model = None + self.cache_dir = os.path.join(folder_paths.models_dir, "RMBG", "body_segment") + self.model_file = "deeplabv3p-resnet50-human.onnx" + + @classmethod + def INPUT_TYPES(cls): + available_classes = [ + "Hair", "Glasses", "Top-clothes", "Bottom-clothes", + "Torso-skin", "Face", "Left-arm", "Right-arm", + "Left-leg", "Right-leg", "Left-foot", "Right-foot" + ] + + tooltips = { + "process_res": "Processing resolution (fixed at 512x512)", + "mask_blur": "Blur amount for mask edges", + "mask_offset": "Expand/Shrink mask boundary", + "background_color": "Choose background color (Alpha = transparent)", + "invert_output": "Invert both image and mask output", + } + + return { + "required": { + "images": ("IMAGE",), + }, + "optional": { + **{cls_name: ("BOOLEAN", {"default": False}) + for cls_name in available_classes}, + "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_color": (["Alpha", "black", "white", "gray", "green", "blue", "red"], {"default": "Alpha", "tooltip": tooltips["background_color"]}), + "invert_output": ("BOOLEAN", {"default": False, "tooltip": tooltips["invert_output"]}), + }, + } + + RETURN_TYPES = ("IMAGE", "MASK") + RETURN_NAMES = ("IMAGE", "MASK") + FUNCTION = "segment_body" + CATEGORY = "🧪AILab/🧽RMBG" + + def check_model_cache(self): + model_path = os.path.join(self.cache_dir, self.model_file) + if not os.path.exists(model_path): + return False, "Model file not found" + return True, "Model cache verified" + + def clear_model(self): + if self.model is not None: + del self.model + self.model = None + + def download_model_files(self): + model_id = "Metal3d/deeplabv3p-resnet50-human" + os.makedirs(self.cache_dir, exist_ok=True) + print("Downloading body segmentation model...") + + try: + downloaded_path = hf_hub_download( + repo_id=model_id, + filename=self.model_file, + local_dir=self.cache_dir, + local_dir_use_symlinks=False + ) + + if os.path.dirname(downloaded_path) != self.cache_dir: + target_path = os.path.join(self.cache_dir, self.model_file) + shutil.move(downloaded_path, target_path) + return True, "Model file downloaded successfully" + except Exception as e: + return False, f"Error downloading model file: {str(e)}" + + def segment_body(self, images, mask_blur=0, mask_offset=0, background_color="Alpha", invert_output=False, **class_selections): + try: + # Check and download model if needed + cache_status, message = self.check_model_cache() + if not cache_status: + print(f"Cache check: {message}") + download_status, download_message = self.download_model_files() + if not download_status: + raise RuntimeError(download_message) + + # Load model if needed + if self.model is None: + self.model = onnxruntime.InferenceSession( + os.path.join(self.cache_dir, self.model_file) + ) + + # Class mapping + class_map = { + "Hair": 2, "Glasses": 4, "Top-clothes": 5, + "Bottom-clothes": 9, "Torso-skin": 10, "Face": 13, + "Left-arm": 14, "Right-arm": 15, "Left-leg": 16, + "Right-leg": 17, "Left-foot": 18, "Right-foot": 19 + } + + # Get selected classes + selected_classes = [name for name, selected in class_selections.items() if selected] + if not selected_classes: + selected_classes = ["Face", "Hair", "Top-clothes", "Bottom-clothes"] + + batch_tensor = [] + batch_masks = [] + + for image in images: + orig_image = tensor2pil(image) + w, h = orig_image.size + + # Resize to 512x512 (model requirement) + input_image = orig_image.resize((512, 512)) + input_array = np.array(input_image).astype(np.float32) / 127.5 - 1 + + # Add batch dimension + input_array = np.expand_dims(input_array, axis=0) + + # Run inference + input_name = self.model.get_inputs()[0].name + output_name = self.model.get_outputs()[0].name + result = self.model.run([output_name], {input_name: input_array}) + + # Process results + result = np.array(result[0]) + pred_seg = result.argmax(axis=3).squeeze(0) + + # Combine selected class masks + combined_mask = np.zeros_like(pred_seg, dtype=np.float32) + for class_name in selected_classes: + mask = (pred_seg == class_map[class_name]).astype(np.float32) + combined_mask = np.clip(combined_mask + mask, 0, 1) + + # Convert to PIL and resize back to original size + mask_image = Image.fromarray((combined_mask * 255).astype(np.uint8)) + mask_image = mask_image.resize((w, h), Image.Resampling.LANCZOS) + + if mask_blur > 0: + mask_image = mask_image.filter(ImageFilter.GaussianBlur(radius=mask_blur)) + + if mask_offset != 0: + if mask_offset > 0: + mask_image = mask_image.filter(ImageFilter.MaxFilter(size=mask_offset * 2 + 1)) + else: + mask_image = mask_image.filter(ImageFilter.MinFilter(size=-mask_offset * 2 + 1)) + + if invert_output: + mask_image = Image.fromarray(255 - np.array(mask_image)) + + # Handle background color + if background_color == "Alpha": + rgba_image = RGB2RGBA(orig_image, mask_image) + result_image = pil2tensor(rgba_image) + else: + bg_colors = { + "black": (0, 0, 0), + "white": (255, 255, 255), + "gray": (128, 128, 128), + "green": (0, 255, 0), + "blue": (0, 0, 255), + "red": (255, 0, 0) + } + + rgba_image = RGB2RGBA(orig_image, mask_image) + bg_image = Image.new('RGBA', orig_image.size, (*bg_colors[background_color], 255)) + composite_image = Image.alpha_composite(bg_image, rgba_image) + result_image = pil2tensor(composite_image.convert('RGB')) + + batch_tensor.append(result_image) + batch_masks.append(pil2tensor(mask_image)) + + # Prepare final output + batch_tensor = torch.cat(batch_tensor, dim=0) + batch_masks = torch.cat(batch_masks, dim=0) + + return (batch_tensor, batch_masks) + + except Exception as e: + self.clear_model() + raise RuntimeError(f"Error in Body Segmentation processing: {str(e)}") + finally: + self.clear_model() + +NODE_CLASS_MAPPINGS = { + "BodySegment": BodySegment +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "BodySegment": "Body Segment (RMBG)" +} \ No newline at end of file diff --git a/AILab_ClothSegment.py b/AILab_ClothSegment.py index b8a80b4..93a4b8f 100644 --- a/AILab_ClothSegment.py +++ b/AILab_ClothSegment.py @@ -2,9 +2,6 @@ # This custom node for ComfyUI provides functionality for background removal using various models, # including RMBG-2.0, INSPYRENET, and BEN. It leverages deep learning techniques # to process images and generate masks for background removal. - -# Models License Notice: -# - mattmdjaga/segformer_b2_clothes: MIT License (https://huggingface.co/mattmdjaga/segformer_b2_clothes) # # This integration script follows GPL-3.0 License. # When using or modifying this code, please respect both the original model licenses @@ -82,14 +79,14 @@ class ClothesSegment: for cls_name in available_classes}, "process_res": ("INT", {"default": 512, "min": 128, "max": 2048, "step": 32, "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": -20, "max": 20, "step": 1, "tooltip": tooltips["mask_offset"]}), + "mask_offset": ("INT", {"default": 0, "min": -64, "max": 64, "step": 1, "tooltip": tooltips["mask_offset"]}), "background_color": (["Alpha", "black", "white", "gray", "green", "blue", "red"], {"default": "Alpha", "tooltip": tooltips["background_color"]}), "invert_output": ("BOOLEAN", {"default": False, "tooltip": tooltips["invert_output"]}), }, } RETURN_TYPES = ("IMAGE", "MASK") - RETURN_NAMES = ("images", "mask") + RETURN_NAMES = ("IMAGE", "MASK") FUNCTION = "segment_clothes" CATEGORY = "🧪AILab/🧽RMBG" diff --git a/AILab_FaceSegment.py b/AILab_FaceSegment.py index ba8d2d5..2acf66f 100644 --- a/AILab_FaceSegment.py +++ b/AILab_FaceSegment.py @@ -85,14 +85,14 @@ class FaceSegment: for cls_name in available_classes}, "process_res": ("INT", {"default": 512, "min": 128, "max": 2048, "step": 32, "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": -20, "max": 20, "step": 1, "tooltip": tooltips["mask_offset"]}), + "mask_offset": ("INT", {"default": 0, "min": -64, "max": 64, "step": 1, "tooltip": tooltips["mask_offset"]}), "background_color": (["Alpha", "black", "white", "gray", "green", "blue", "red"], {"default": "Alpha", "tooltip": tooltips["background_color"]}), "invert_output": ("BOOLEAN", {"default": False, "tooltip": tooltips["invert_output"]}), }, } RETURN_TYPES = ("IMAGE", "MASK") - RETURN_NAMES = ("images", "mask") + RETURN_NAMES = ("IMAGE", "MASK") FUNCTION = "segment_face" CATEGORY = "🧪AILab/🧽RMBG" diff --git a/AILab_FashionSegment.py b/AILab_FashionSegment.py index d10a411..6125751 100644 --- a/AILab_FashionSegment.py +++ b/AILab_FashionSegment.py @@ -1,9 +1,6 @@ # ComfyUI-RMBG # This custom node for ComfyUI provides functionality for fashion segmentation using segformer-b3-fashion model. # It leverages deep learning techniques to process images and generate masks for fashion items segmentation. - -# Models License Notice: -# - sayeed99/segformer-b3-fashion: MIT License (https://huggingface.co/sayeed99/segformer-b3-fashion) # # This integration script follows GPL-3.0 License. # When using or modifying this code, please respect both the original model licenses @@ -166,7 +163,7 @@ class FashionSegmentClothing: for cls_name in clothing_classes}, "process_res": ("INT", {"default": 512, "min": 128, "max": 2048, "step": 32}), "mask_blur": ("INT", {"default": 0, "min": 0, "max": 64, "step": 1}), - "mask_offset": ("INT", {"default": 0, "min": -20, "max": 20, "step": 1}), + "mask_offset": ("INT", {"default": 0, "min": -64, "max": 64, "step": 1}), "background_color": (["Alpha", "black", "white", "gray", "green", "blue", "red"], {"default": "Alpha"}), "invert_output": ("BOOLEAN", {"default": False}), }, diff --git a/AILab_ImageMaskTools.py b/AILab_ImageMaskTools.py new file mode 100644 index 0000000..33659b8 --- /dev/null +++ b/AILab_ImageMaskTools.py @@ -0,0 +1,320 @@ +# ComfyUI-RMBG v2.0.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. +# +# AILab Image and Mask Tools +# This module is specifically designed for ComfyUI-RMBG, enhancing workflows within ComfyUI. +# 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. +# +# 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 +import folder_paths +import numpy as np +import hashlib +import torch +import cv2 +from PIL import Image, ImageFilter, ImageOps, ImageSequence, ImageChops +import torchvision.transforms.functional as T +from scipy import ndimage + +# 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 pil2mask(image): + return torch.from_numpy(np.array(image.convert("L")).astype(np.float32) / 255.0).unsqueeze(0) + +def blend_overlay(img_1, img_2): + arr1 = np.array(img_1).astype(float) / 255.0 + arr2 = np.array(img_2).astype(float) / 255.0 + mask = arr2 < 0.5 + result = np.zeros_like(arr1) + result[mask] = 2 * arr1[mask] * arr2[mask] + result[~mask] = 1 - 2 * (1 - arr1[~mask]) * (1 - arr2[~mask]) + return Image.fromarray(np.clip(result * 255, 0, 255).astype(np.uint8)) + +# Base class for preview +class AiLab_PreviewBase: + def __init__(self): + self.output_dir = folder_paths.get_temp_directory() + self.type = "temp" + self.prefix_append = "" + + def get_unique_filename(self, filename_prefix): + os.makedirs(self.output_dir, exist_ok=True) + filename = filename_prefix + self.prefix_append + counter = 1 + while True: + file = f"{filename}_{counter:04d}.png" + full_path = os.path.join(self.output_dir, file) + if not os.path.exists(full_path): + return full_path, file + counter += 1 + + def save_image(self, image, filename_prefix, prompt=None, extra_pnginfo=None): + results = [] + + try: + if isinstance(image, torch.Tensor): + if len(image.shape) == 4: # Batch of images + for i in range(image.shape[0]): + full_output_path, file = self.get_unique_filename(filename_prefix) + img = Image.fromarray(np.clip(image[i].cpu().numpy() * 255, 0, 255).astype(np.uint8)) + img.save(full_output_path) + results.append({"filename": full_output_path, "subfolder": "", "type": self.type}) + else: + full_output_path, file = self.get_unique_filename(filename_prefix) + img = Image.fromarray(np.clip(image.cpu().numpy() * 255, 0, 255).astype(np.uint8)) + img.save(full_output_path) + results.append({"filename": full_output_path, "subfolder": "", "type": self.type}) + else: + full_output_path, file = self.get_unique_filename(filename_prefix) + image.save(full_output_path) + results.append({"filename": full_output_path, "subfolder": "", "type": self.type}) + + return { + "ui": {"images": results}, + } + except Exception as e: + print(f"Error saving image: {e}") + return {"ui": {}} + +# Preview node +class AiLab_Preview(AiLab_PreviewBase): + def __init__(self): + super().__init__() + self.prefix_append = "_preview_" + ''.join(random.choice("abcdefghijklmnopqrstupvxyz") for x in range(5)) + + @classmethod + def INPUT_TYPES(s): + return { + "optional": { + "image": ("IMAGE", {"default": None}), + "mask": ("MASK", {"default": None}), + }, + "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"}, + } + + RETURN_TYPES = ("IMAGE", "MASK") + RETURN_NAMES = ("IMAGE", "MASK") + FUNCTION = "preview" + OUTPUT_NODE = True + CATEGORY = "🧪AILab/🛠️UTIL/🖼️IMAGE" + + def preview(self, image=None, mask=None, prompt=None, extra_pnginfo=None): + results = [] + + if image is not None: + image_result = self.save_image(image, "image_preview", prompt, extra_pnginfo) + if "ui" in image_result and "images" in image_result["ui"]: + results.extend(image_result["ui"]["images"]) + + if mask is not None: + preview = mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])).movedim(1, -1).expand(-1, -1, -1, 3) + mask_result = self.save_image(preview, "mask_preview", prompt, extra_pnginfo) + if "ui" in mask_result and "images" in mask_result["ui"]: + results.extend(mask_result["ui"]["images"]) + + return { + "ui": {"images": results}, + "result": (image if image is not None else None, mask if mask is not None else None) + } + +# Mask preview node +class AiLab_MaskPreview(AiLab_PreviewBase): + def __init__(self): + super().__init__() + self.prefix_append = "_mask_preview_" + ''.join(random.choice("abcdefghijklmnopqrstupvxyz") for x in range(5)) + + @classmethod + def INPUT_TYPES(s): + return { + "required": {"mask": ("MASK",),}, + "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"}, + } + + RETURN_TYPES = ("MASK",) + RETURN_NAMES = ("MASK",) + FUNCTION = "preview_mask" + OUTPUT_NODE = True + CATEGORY = "🧪AILab/🛠️UTIL/🖼️IMAGE" + + def preview_mask(self, mask, prompt=None, extra_pnginfo=None): + preview = mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])).movedim(1, -1).expand(-1, -1, -1, 3) + result = self.save_image(preview, "mask_preview", prompt, extra_pnginfo) + return { + "ui": result["ui"], + "result": (mask,) + } + +# Image preview node +class AiLab_ImagePreview(AiLab_PreviewBase): + def __init__(self): + super().__init__() + self.prefix_append = "_image_preview_" + ''.join(random.choice("abcdefghijklmnopqrstupvxyz") for x in range(5)) + + @classmethod + def INPUT_TYPES(s): + return { + "required": {"image": ("IMAGE",),}, + "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"}, + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("IMAGE",) + FUNCTION = "preview_image" + OUTPUT_NODE = True + CATEGORY = "🧪AILab/🛠️UTIL/🖼️IMAGE" + + def preview_image(self, image, prompt=None, extra_pnginfo=None): + result = self.save_image(image, "image_preview", prompt, extra_pnginfo) + return { + "ui": result["ui"], + "result": (image,) + } + +# Image loader node +class AiLab_LoadImage: + @classmethod + def INPUT_TYPES(cls): + input_dir = folder_paths.get_input_directory() + os.makedirs(input_dir, exist_ok=True) + files = [f for f in os.listdir(input_dir) if os.path.isfile(os.path.join(input_dir, f)) and f.lower().endswith(('.png', '.jpg', '.jpeg', '.webp', '.gif', '.bmp', '.tiff', '.tif'))] + return { + "required": { + "image": (sorted(files) or [""], {"image_upload": True}), + "mask_channel": (["alpha", "red", "green", "blue"], {"default": "alpha", "tooltip": "Select channel to extract mask from"}), + "scale_by": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 8.0, "step": 0.01, "tooltip": "Scale image by this factor (ignored if longest_side > 0)"}), + "longest_side": ("INT", {"default": 0, "min": 0, "max": 8192, "step": 8, "tooltip": "Resize image so longest side equals this value (0 = disabled)"}), + }, + "hidden": { + "extra_pnginfo": "EXTRA_PNGINFO", + }, + } + + CATEGORY = "🧪AILab/🛠️UTIL/🖼️IMAGE" + RETURN_TYPES = ("IMAGE", "MASK", "IMAGE", "INT", "INT") + RETURN_NAMES = ("IMAGE", "MASK", "MASK_IMAGE", "WIDTH", "HEIGHT") + FUNCTION = "load_image" + OUTPUT_NODE = False + + def load_image(self, image, mask_channel="alpha", scale_by=1.0, longest_side=0, extra_pnginfo=None): + try: + image_path = folder_paths.get_annotated_filepath(image) + img = Image.open(image_path) + + orig_width, orig_height = img.size + if longest_side > 0: + if orig_width >= orig_height: + new_width = longest_side + new_height = int(orig_height * (longest_side / orig_width)) + img = img.resize((new_width, new_height), Image.LANCZOS) + else: + new_height = longest_side + new_width = int(orig_width * (longest_side / orig_height)) + img = img.resize((new_width, new_height), Image.LANCZOS) + elif scale_by != 1.0: + new_width = int(orig_width * scale_by) + new_height = int(orig_height * scale_by) + img = img.resize((new_width, new_height), Image.LANCZOS) + + width, height = img.size + + output_images = [] + output_masks = [] + for i in ImageSequence.Iterator(img): + i = ImageOps.exif_transpose(i) + if i.mode == 'I': + i = i.point(lambda i: i * (1 / 255)) + image = i.convert("RGB") + image = np.array(image).astype(np.float32) / 255.0 + image = torch.from_numpy(image)[None,] + + if mask_channel == "alpha" and 'A' in i.getbands(): + mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0 + mask = 1. - torch.from_numpy(mask) + elif mask_channel == "red" and 'R' in i.getbands(): + mask = np.array(i.getchannel('R')).astype(np.float32) / 255.0 + mask = torch.from_numpy(mask) + elif mask_channel == "green" and 'G' in i.getbands(): + mask = np.array(i.getchannel('G')).astype(np.float32) / 255.0 + mask = torch.from_numpy(mask) + elif mask_channel == "blue" and 'B' in i.getbands(): + mask = np.array(i.getchannel('B')).astype(np.float32) / 255.0 + mask = torch.from_numpy(mask) + else: + mask = torch.ones((height, width), dtype=torch.float32, device="cpu") + + output_images.append(image) + output_masks.append(mask.unsqueeze(0)) + + if len(output_images) > 1: + output_image = torch.cat(output_images, dim=0) + output_mask = torch.cat(output_masks, dim=0) + else: + output_image = output_images[0] + output_mask = output_masks[0] + + mask_image = output_mask.reshape((-1, 1, output_mask.shape[-2], output_mask.shape[-1])).movedim(1, -1).expand(-1, -1, -1, 3) + + return (output_image, output_mask, mask_image, width, height) + + except Exception as e: + import traceback + traceback.print_exc() + print(f"Error loading image: {e}") + empty_image = torch.zeros(1, 3, 64, 64) + empty_mask = torch.zeros(1, 64, 64) + empty_mask_image = empty_mask.reshape((-1, 1, 64, 64)).movedim(1, -1).expand(-1, -1, -1, 3) + return (empty_image, empty_mask, empty_mask_image, 64, 64) + + @classmethod + def IS_CHANGED(cls, image, mask_channel="alpha", scale_by=1.0, longest_side=0, extra_pnginfo=None): + image_path = folder_paths.get_annotated_filepath(image) + m = hashlib.sha256() + with open(image_path, 'rb') as f: + m.update(f.read()) + return m.digest().hex() + + @classmethod + def VALIDATE_INPUTS(cls, image, mask_channel="alpha", scale_by=1.0, longest_side=0, extra_pnginfo=None): + if not folder_paths.exists_annotated_filepath(image): + return f"Invalid image file: {image}" + + return True + + + +# Node class mappings +NODE_CLASS_MAPPINGS = { + "AiLab_LoadImage": AiLab_LoadImage, + "AiLab_Preview": AiLab_Preview, + "AiLab_ImagePreview": AiLab_ImagePreview, + "AiLab_MaskPreview": AiLab_MaskPreview, +} + +# 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) 🎭", +} \ No newline at end of file diff --git a/AILab_RMBG.py b/AILab_RMBG.py index 7e29583..63d14c7 100644 --- a/AILab_RMBG.py +++ b/AILab_RMBG.py @@ -1,579 +1,597 @@ -# ComfyUI-RMBG v1.9.3 -# 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): - return os.path.join(self.base_cache_dir, AVAILABLE_MODELS[model_name]["cache_dir"]) - - 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 - self.model = None - self.current_model_version = None - torch.cuda.empty_cache() - print("Model cleared from memory") - -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) - self.model = AutoModelForImageSegmentation.from_pretrained( - cache_dir, - trust_remote_code=True, - local_files_only=True - ) - - 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(): - results = self.model(input_batch)[-1].sigmoid().cpu() - 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": 32, "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": -20, "max": 20, "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") - RETURN_NAMES = ("image", "mask") - 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)) - - return (torch.cat(processed_images, dim=0), torch.cat(processed_masks, dim=0)) - - except Exception as e: - handle_model_error(f"Error in image processing: {str(e)}") - -# Node Mapping -NODE_CLASS_MAPPINGS = { - "RMBG": RMBG -} - -NODE_DISPLAY_NAME_MAPPINGS = { - "RMBG": "Remove Background (RMBG)" +# ComfyUI-RMBG v2.0.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 + +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) + self.model = AutoModelForImageSegmentation.from_pretrained( + cache_dir, + trust_remote_code=True, + local_files_only=True + ) + + 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(): + results = self.model(input_batch)[-1].sigmoid().cpu() + 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/AILab_Segment.py b/AILab_Segment.py index cc0aaa9..92b6008 100644 --- a/AILab_Segment.py +++ b/AILab_Segment.py @@ -157,13 +157,14 @@ class Segment: "optional": { "threshold": ("FLOAT", {"default": 0.35, "min": 0.05, "max": 0.95, "step": 0.01, "tooltip": tooltips["threshold"]}), "mask_blur": ("INT", {"default": 0, "min": 0, "max": 64, "step": 1, "tooltip": tooltips["mask_blur"]}), - "mask_offset": ("INT", {"default": 0, "min": -20, "max": 20, "step": 1, "tooltip": tooltips["mask_offset"]}), + "mask_offset": ("INT", {"default": 0, "min": -64, "max": 64, "step": 1, "tooltip": tooltips["mask_offset"]}), "background_color": (["Alpha", "black", "white", "gray", "green", "blue", "red"], {"default": "Alpha", "tooltip": tooltips["background_color"]}), "invert_output": ("BOOLEAN", {"default": False}), } } RETURN_TYPES = ("IMAGE", "MASK") + RETURN_NAMES = ("IMAGE", "MASK") FUNCTION = "segment" CATEGORY = "🧪AILab/🧽RMBG" diff --git a/requirements.txt b/requirements.txt index 0776077..fcf196d 100644 --- a/requirements.txt +++ b/requirements.txt @@ -10,3 +10,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