From d39173b4d4a405566db4e719fd3d34bd25a2e5eb Mon Sep 17 00:00:00 2001 From: AI Lab <129358391+1038lab@users.noreply.github.com> Date: Tue, 4 Feb 2025 20:53:33 -0800 Subject: [PATCH] Add files via upload --- AILab_ClothSegment.py | 556 +++++++++++------------ AILab_FaceSegment.py | 2 +- AILab_FashionSegment.py | 2 +- AILab_RMBG.py | 961 ++++++++++++++++++++++------------------ AILab_Segment.py | 2 +- README.md | 6 +- pyproject.toml | 4 +- update.md | 520 ++++++++++++---------- 8 files changed, 1089 insertions(+), 964 deletions(-) diff --git a/AILab_ClothSegment.py b/AILab_ClothSegment.py index 77e0073..b8a80b4 100644 --- a/AILab_ClothSegment.py +++ b/AILab_ClothSegment.py @@ -1,279 +1,279 @@ -# ComfyUI-RMBG v1.6.0 -# 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 -# 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 -from transformers import SegformerImageProcessor, AutoModelForSemanticSegmentation -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")) - -AVAILABLE_MODELS = { - "segformer_b2_clothes": "1038lab/segformer_clothes" -} - -class ClothesSegment: - def __init__(self): - self.processor = None - self.model = None - self.cache_dir = os.path.join(folder_paths.models_dir, "RMBG", "segformer_clothes") - - @classmethod - def INPUT_TYPES(cls): - available_classes = ["Hat", "Hair", "Face", "Sunglasses", "Upper-clothes", "Skirt", "Dress", "Belt", "Pants", "Left-arm", "Right-arm", "Left-leg", "Right-leg", "Bag", "Scarf", "Left-shoe", "Right-shoe","Background"] - - tooltips = { - "process_res": "Processing resolution (higher = more VRAM)", - "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}, - "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"]}), - "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") - FUNCTION = "segment_clothes" - CATEGORY = "🧪AILab/🧽RMBG" - - def check_model_cache(self): - if not os.path.exists(self.cache_dir): - return False, "Model directory not found" - - required_files = [ - 'config.json', - 'model.safetensors', - 'preprocessor_config.json' - ] - - missing_files = [f for f in required_files if not os.path.exists(os.path.join(self.cache_dir, f))] - if missing_files: - return False, f"Required model files missing: {', '.join(missing_files)}" - return True, "Model cache verified" - - def clear_model(self): - if self.model is not None: - self.model.cpu() - del self.model - self.model = None - self.processor = None - torch.cuda.empty_cache() - - def download_model_files(self): - model_id = AVAILABLE_MODELS["segformer_b2_clothes"] - model_files = { - 'config.json': 'config.json', - 'model.safetensors': 'model.safetensors', - 'preprocessor_config.json': 'preprocessor_config.json' - } - - os.makedirs(self.cache_dir, exist_ok=True) - print(f"Downloading Clothes Segformer model files...") - - try: - for save_name, repo_path in model_files.items(): - print(f"Downloading {save_name}...") - downloaded_path = hf_hub_download( - repo_id=model_id, - filename=repo_path, - 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, save_name) - shutil.move(downloaded_path, target_path) - return True, "Model files downloaded successfully" - except Exception as e: - return False, f"Error downloading model files: {str(e)}" - - def segment_clothes(self, images, process_res=1024, 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.processor is None: - self.processor = SegformerImageProcessor.from_pretrained(self.cache_dir) - self.model = AutoModelForSemanticSegmentation.from_pretrained(self.cache_dir) - self.model.eval() - for param in self.model.parameters(): - param.requires_grad = False - self.model.to(device) - - # Class mapping for segmentation - class_map = { - "Background": 0, "Hat": 1, "Hair": 2, "Sunglasses": 3, - "Upper-clothes": 4, "Skirt": 5, "Pants": 6, "Dress": 7, - "Belt": 8, "Left-shoe": 9, "Right-shoe": 10, "Face": 11, - "Left-leg": 12, "Right-leg": 13, "Left-arm": 14, "Right-arm": 15, - "Bag": 16, "Scarf": 17 - } - - # Get selected classes - selected_classes = [name for name, selected in class_selections.items() if selected] - if not selected_classes: - selected_classes = ["Upper-clothes"] - - # Image preprocessing - transform_image = transforms.Compose([ - transforms.Resize((process_res, process_res)), - transforms.ToTensor(), - ]) - - batch_tensor = [] - batch_masks = [] - - for image in images: - orig_image = tensor2pil(image) - w, h = orig_image.size - - input_tensor = transform_image(orig_image) - - if input_tensor.shape[0] == 4: - input_tensor = input_tensor[:3] - - input_tensor = transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])(input_tensor) - - input_tensor = input_tensor.unsqueeze(0).to(device) - - with torch.no_grad(): - outputs = self.model(input_tensor) - logits = outputs.logits.cpu() - upsampled_logits = nn.functional.interpolate( - logits, - size=(h, w), - mode="bilinear", - align_corners=False, - ) - pred_seg = upsampled_logits.argmax(dim=1)[0] - - # Combine selected class masks - combined_mask = None - for class_name in selected_classes: - mask = (pred_seg == class_map[class_name]).float() - if combined_mask is None: - combined_mask = mask - else: - combined_mask = torch.clamp(combined_mask + mask, 0, 1) - - # Convert mask to PIL for processing - mask_image = Image.fromarray((combined_mask.numpy() * 255).astype(np.uint8)) - - 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 Clothes Segformer processing: {str(e)}") - finally: - - if not self.model.training: - self.clear_model() - -NODE_CLASS_MAPPINGS = { - "ClothesSegment": ClothesSegment -} - -NODE_DISPLAY_NAME_MAPPINGS = { - "ClothesSegment": "Clothes Segment (RMBG)" +# 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. + +# 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 +# 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 +from transformers import SegformerImageProcessor, AutoModelForSemanticSegmentation +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")) + +AVAILABLE_MODELS = { + "segformer_b2_clothes": "1038lab/segformer_clothes" +} + +class ClothesSegment: + def __init__(self): + self.processor = None + self.model = None + self.cache_dir = os.path.join(folder_paths.models_dir, "RMBG", "segformer_clothes") + + @classmethod + def INPUT_TYPES(cls): + available_classes = ["Hat", "Hair", "Face", "Sunglasses", "Upper-clothes", "Skirt", "Dress", "Belt", "Pants", "Left-arm", "Right-arm", "Left-leg", "Right-leg", "Bag", "Scarf", "Left-shoe", "Right-shoe","Background"] + + tooltips = { + "process_res": "Processing resolution (higher = more VRAM)", + "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}, + "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"]}), + "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") + FUNCTION = "segment_clothes" + CATEGORY = "🧪AILab/🧽RMBG" + + def check_model_cache(self): + if not os.path.exists(self.cache_dir): + return False, "Model directory not found" + + required_files = [ + 'config.json', + 'model.safetensors', + 'preprocessor_config.json' + ] + + missing_files = [f for f in required_files if not os.path.exists(os.path.join(self.cache_dir, f))] + if missing_files: + return False, f"Required model files missing: {', '.join(missing_files)}" + return True, "Model cache verified" + + def clear_model(self): + if self.model is not None: + self.model.cpu() + del self.model + self.model = None + self.processor = None + torch.cuda.empty_cache() + + def download_model_files(self): + model_id = AVAILABLE_MODELS["segformer_b2_clothes"] + model_files = { + 'config.json': 'config.json', + 'model.safetensors': 'model.safetensors', + 'preprocessor_config.json': 'preprocessor_config.json' + } + + os.makedirs(self.cache_dir, exist_ok=True) + print(f"Downloading Clothes Segformer model files...") + + try: + for save_name, repo_path in model_files.items(): + print(f"Downloading {save_name}...") + downloaded_path = hf_hub_download( + repo_id=model_id, + filename=repo_path, + 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, save_name) + shutil.move(downloaded_path, target_path) + return True, "Model files downloaded successfully" + except Exception as e: + return False, f"Error downloading model files: {str(e)}" + + def segment_clothes(self, images, process_res=1024, 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.processor is None: + self.processor = SegformerImageProcessor.from_pretrained(self.cache_dir) + self.model = AutoModelForSemanticSegmentation.from_pretrained(self.cache_dir) + self.model.eval() + for param in self.model.parameters(): + param.requires_grad = False + self.model.to(device) + + # Class mapping for segmentation + class_map = { + "Background": 0, "Hat": 1, "Hair": 2, "Sunglasses": 3, + "Upper-clothes": 4, "Skirt": 5, "Pants": 6, "Dress": 7, + "Belt": 8, "Left-shoe": 9, "Right-shoe": 10, "Face": 11, + "Left-leg": 12, "Right-leg": 13, "Left-arm": 14, "Right-arm": 15, + "Bag": 16, "Scarf": 17 + } + + # Get selected classes + selected_classes = [name for name, selected in class_selections.items() if selected] + if not selected_classes: + selected_classes = ["Upper-clothes"] + + # Image preprocessing + transform_image = transforms.Compose([ + transforms.Resize((process_res, process_res)), + transforms.ToTensor(), + ]) + + batch_tensor = [] + batch_masks = [] + + for image in images: + orig_image = tensor2pil(image) + w, h = orig_image.size + + input_tensor = transform_image(orig_image) + + if input_tensor.shape[0] == 4: + input_tensor = input_tensor[:3] + + input_tensor = transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])(input_tensor) + + input_tensor = input_tensor.unsqueeze(0).to(device) + + with torch.no_grad(): + outputs = self.model(input_tensor) + logits = outputs.logits.cpu() + upsampled_logits = nn.functional.interpolate( + logits, + size=(h, w), + mode="bilinear", + align_corners=False, + ) + pred_seg = upsampled_logits.argmax(dim=1)[0] + + # Combine selected class masks + combined_mask = None + for class_name in selected_classes: + mask = (pred_seg == class_map[class_name]).float() + if combined_mask is None: + combined_mask = mask + else: + combined_mask = torch.clamp(combined_mask + mask, 0, 1) + + # Convert mask to PIL for processing + mask_image = Image.fromarray((combined_mask.numpy() * 255).astype(np.uint8)) + + 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 Clothes Segformer processing: {str(e)}") + finally: + + if not self.model.training: + self.clear_model() + +NODE_CLASS_MAPPINGS = { + "ClothesSegment": ClothesSegment +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "ClothesSegment": "Clothes Segment (RMBG)" } \ No newline at end of file diff --git a/AILab_FaceSegment.py b/AILab_FaceSegment.py index 483d263..ba8d2d5 100644 --- a/AILab_FaceSegment.py +++ b/AILab_FaceSegment.py @@ -1,4 +1,4 @@ -# ComfyUI-RMBG v1.6.0 +# ComfyUI-RMBG # This custom node for ComfyUI provides functionality for face parsing using Segformer model. # # This integration script follows GPL-3.0 License. diff --git a/AILab_FashionSegment.py b/AILab_FashionSegment.py index e52420d..d10a411 100644 --- a/AILab_FashionSegment.py +++ b/AILab_FashionSegment.py @@ -1,4 +1,4 @@ -# ComfyUI-RMBG v1.6.0 +# 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. diff --git a/AILab_RMBG.py b/AILab_RMBG.py index f72fca8..4f1df92 100644 --- a/AILab_RMBG.py +++ b/AILab_RMBG.py @@ -1,438 +1,525 @@ -# ComfyUI-RMBG v1.6.0 -# 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: -# - 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) -# -# 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 -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 tqdm import tqdm -from transformers import AutoModelForImageSegmentation - -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": "briaai/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": "PramaLLC/BEN", - "files": { - "model.py": "model.py", - "BEN_Base.pth": "BEN_Base.pth" - }, - "cache_dir": "BEN" - } -} - -# 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 RMBG: - def __init__(self): - self.models = { - "RMBG-2.0": RMBGModel(), - "INSPYRENET": InspyrenetModel(), - "BEN": BENModel() - } - - @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)." - } - - 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"]}) - } - } - - 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)) - - # Create final image - orig_image = tensor2pil(img) - 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) - - # Convert to RGB if background is not Alpha - processed_images.append(pil2tensor(composite_image.convert("RGB"))) - else: - # Keep as RGBA if background is Alpha - 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 v1.7.0 +# This custom node for ComfyUI provides functionality for background removal using various models, +# including RMBG-2.0, INSPYRENET, BEN and BEN2. 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/AILab-AI/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 tqdm import tqdm +from transformers import AutoModelForImageSegmentation + +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": "briaai/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() + + 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 + + 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(): + foregrounds = self.model.inference(batch_pil_images, refine_foreground=False) + if not isinstance(foregrounds, list): + foregrounds = [foregrounds] + + 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)}") + +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)." + } + + 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"]}) + } + } + + 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)) + + # Create final image + orig_image = tensor2pil(img) + 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) + + # Convert to RGB if background is not Alpha + processed_images.append(pil2tensor(composite_image.convert("RGB"))) + else: + # Keep as RGBA if background is Alpha + 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)" } \ No newline at end of file diff --git a/AILab_Segment.py b/AILab_Segment.py index f1ceecf..cc0aaa9 100644 --- a/AILab_Segment.py +++ b/AILab_Segment.py @@ -1,4 +1,4 @@ -# ComfyUI-RMBG v1.6.0 +# 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. diff --git a/README.md b/README.md index 2c6c1df..b0151f3 100644 --- a/README.md +++ b/README.md @@ -6,6 +6,8 @@ $${\color{red}If\ this\ custom\ node\ helps\ you\ or\ you\ like\ my\ work,\ plea $${\color{red}It's\ a\ greatest\ encouragement\ for\ my\ efforts!}$$ ## News & Updates +- 2025/02/04: Update ComfyUI-RMBG to v1.7.0 with new BEN2 model ( [update.md](https://github.com/1038lab/ComfyUI-RMBG/blob/main/update.md#v170-20250204) ) + - 2025/01/22: Update ComfyUI-RMBG to v1.6.0 with new Face Segment custom node ( [update.md](https://github.com/1038lab/ComfyUI-RMBG/blob/main/update.md#v160-20250122) ) ![RMBG_v1 6 0](https://github.com/user-attachments/assets/9ccefec1-4370-4708-a12d-544c90888bf2) @@ -91,8 +93,10 @@ install requirment.txt in the ComfyUI-RMBG folder - The model will be automatically downloaded to `ComfyUI/models/RMBG/` when first time using the custom node. - Manually download the RMBG-2.0 model by visiting this [link](https://huggingface.co/briaai/RMBG-2.0/tree/main), then download the files and place them in the `/ComfyUI/models/RMBG/RMBG-2.0` folder. - Manually download the INSPYRENET models by visiting the [link](https://huggingface.co/1038lab/inspyrenet), then download the files and place them in the `/ComfyUI/models/RMBG/INSPYRENET` folder. -- Manually download the BEN model by visiting the [link](https://huggingface.co/PramaLLC/BEN), then download the files and place them in the `/ComfyUI/models/RMBG/BEN` folder. +- Manually download the BEN model by visiting the [link](https://huggingface.co/1038lab/BEN), then download the files and place them in the `/ComfyUI/models/RMBG/BEN` folder. +- Manually download the BEN2 model by visiting the [link](https://huggingface.co/1038lab/BEN2), then download the files and place them in the `/ComfyUI/models/RMBG/BEN2` folder. - Manually download the SAM models by visiting the [link](https://huggingface.co/1038lab/sam), then download the files and place them in the `/ComfyUI/models/SAM` folder. + - Manually download the GroundingDINO models by visiting the [link](https://huggingface.co/1038lab/GroundingDINO), then download the files and place them in the `/ComfyUI/models/grounding-dino` folder. - Manually download the Clothes Segment model by visiting the [link](https://huggingface.co/1038lab/segformer_clothes), then download the files and place them in the `/ComfyUI/models/RMBG/segformer_clothes` folder. - Manually download the Fashion Segment model by visiting the [link](https://huggingface.co/1038lab/segformer_fashion), then download the files and place them in the `/ComfyUI/models/RMBG/segformer_fashion` folder. diff --git a/pyproject.toml b/pyproject.toml index 23ceb05..85e5881 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui-rmbg" -description = "A ComfyUI custom node designed for advanced image background removal and object, face, clothes, and fashion segmentation, utilizing multiple models including RMBG-2.0, INSPYRENET, BEN, SAM, and GroundingDINO." -version = "1.6.0" +description = "A ComfyUI custom node designed for advanced image background removal and object, face, clothes, and fashion segmentation, utilizing multiple models including RMBG-2.0, INSPYRENET, BEN, BEN2, SAM, and GroundingDINO." +version = "1.7.0" license = {file = "LICENSE"} dependencies = ["torch>=2.0.0", "torchvision>=0.15.0", "Pillow>=9.0.0", "numpy>=1.22.0", "huggingface-hub>=0.19.0", "transformers>=4.35.0", "transparent-background>=1.2.4", "tqdm>=4.65.0", "segment-anything>=1.0", "groundingdino-py>=0.4.0", "opencv-python>=4.7.0"] diff --git a/update.md b/update.md index 3475c68..4a5c0c7 100644 --- a/update.md +++ b/update.md @@ -1,243 +1,277 @@ -# ComfyUI-RMBG Update Log - -## v1.6.0 (2025/01/22) - -### New Face Segment Custom Node -- Added a new custom node for face parsing and segmentation - - Support for 19 facial feature categories (Skin, Nose, Eyes, Eyebrows, etc.) - - Precise facial feature extraction and segmentation - - Multiple feature selection for combined segmentation - - Same parameter controls as other RMBG nodes - - Automatic model downloading and resource management - - Perfect for portrait editing and facial feature manipulation - -![RMBG_v1 6 0](https://github.com/user-attachments/assets/9ccefec1-4370-4708-a12d-544c90888bf2) - -## v1.5.0 (2025/01/05) - -### New Fashion and accessories Segment Custom Node -- Added a new custom node for fashion and accessories segmentation. - - Capable of identifying and segmenting various fashion items such as dresses, shoes, and accessories. - - Utilizes advanced machine learning techniques for accurate segmentation. - - Supports real-time processing for enhanced user experience. - - Ideal for fashion-related applications, including virtual try-ons and outfit recommendations. - - Support for gray background color. - -![RMBGv_1 5 0](https://github.com/user-attachments/assets/a250c1a6-8425-4902-b902-a6e1a8bfe959) - -## v1.4.0 (2025/01/02) - -### New Clothes Segment Node -- Added intelligent clothes segmentation functionality - - Support for 18 different clothing categories (Hat, Hair, Face, Sunglasses, Upper-clothes, etc.) - - Multiple item selection for combined segmentation - - Same parameter controls as other RMBG nodes (process_res, mask_blur, mask_offset, background options) - - Automatic model downloading and resource management - -![rmbg_v1 4 0](https://github.com/user-attachments/assets/978c168b-03a8-4937-aa03-06385f34b820) - -## v1.3.2 (2024/12/29) - -### Updates -- Enhanced background handling to support RGBA output when "Alpha" is selected. -- Ensured RGB output for all other background color selections. - -## v1.3.1 (2024/12/25) - -### Bug Fixes -- Fixed an issue with mask processing when the model returns a list of masks. -- Improved handling of image formats to prevent processing errors. - -## v1.3.0 (2024/12/23) - -### New Segment (RMBG) Node -- Text-Prompted Intelligent Object Segmentation - - Use natural language prompts (e.g., "a cat", "red car") to identify and segment target objects - - Support for multiple object detection and segmentation - - Perfect for precise object extraction and recognition tasks - -![rmbg v1.3.0](https://github.com/user-attachments/assets/7607546e-ffcb-45e2-ab90-83267292757e) - -### Supported Models -- SAM (Segment Anything Model) - - sam_vit_h: 2.56GB - Highest accuracy - - sam_vit_l: 1.25GB - Balanced performance - - sam_vit_b: 375MB - Lightweight option -- GroundingDINO - - SwinT: 694MB - Fast and efficient - - SwinB: 938MB - Higher precision - -### Key Features -- Intuitive Parameter Controls - - Threshold: Adjust detection precision - - Mask Blur: Smooth edges - - Mask Offset: Expand or shrink selection - - Background Options: Alpha/Black/White/Green/Blue/Red -- Automatic Model Management - - Auto-download models on first use - - Smart GPU memory handling - -### Usage Examples -1. Tag-Style Prompts - - Single object: "cat" - - Multiple objects: "cat, dog, person" - - With attributes: "red car, blue shirt" - - Format: Use commas to separate multiple objects (e.g., "a, b, c") - -2. Natural Language Prompts - - Simple sentence: "a person wearing a red jacket" - - Complex scene: "a woman in a blue dress standing next to a car" - - With location: "a cat sitting on the sofa" - - Format: Write a natural descriptive sentence - -3. Tips for Better Results - - For Tag Style: - - Separate objects with commas: "chair, table, lamp" - - Add attributes before objects: "wooden chair, glass table" - - Keep it simple and clear - - For Natural Language: - - Use complete sentences - - Include details like color, position, action - - Be as descriptive as needed - - Parameter Adjustments: - - Threshold: 0.25-0.35 for broad detection, 0.45-0.55 for precision - - Use mask blur for smoother edges - - Adjust mask offset to fine-tune selection - -## v1.2.2 (2024/12/12) -![RMBG1 2 2](https://github.com/user-attachments/assets/cb7b1ad0-a2ca-4369-9401-54957af6c636) - -### Improvements -- Changed INSPYRENET model format from .pth to .safetensors for: - - Better security - - Faster loading speed (2-3x faster) - - Improved memory efficiency - - Better cross-platform compatibility -- Simplified node display name for better UI integration - -## v1.2.1 (2024/12/02) - -### New Features -- ANPG (animated PNG), AWEBP (animated WebP) and GIF supported. - -https://github.com/user-attachments/assets/40ec0b27-4fa2-4c99-9aea-5afad9ca62a5 - -### Bug Fixes -- Fixed video processing issue - -### Performance Improvements -- Enhanced batch processing in RMBG-2.0 model -- Added support for proper batch image handling -- Improved memory efficiency by optimizing image size handling - -### Technical Details -- Added original size preservation for maintaining aspect ratios -- Implemented proper batch tensor processing -- Improved error handling and code robustness -- Performance gains: - - Single image processing: ~5-10% improvement - - Batch processing: up to 30-50% improvement (depending on batch size and GPU) - -## v1.2.0 (2024/11/29) - -### Major Changes -- Combined three background removal models into one unified node -- Added support for RMBG-2.0, INSPYRENET, and BEN models -- Implemented lazy loading for models (only downloads when first used) - -### Model Introduction -- RMBG-2.0 ([Homepage](https://huggingface.co/briaai/RMBG-2.0)) - - Latest version of RMBG model - - Excellent performance on complex backgrounds - - High accuracy in preserving fine details - - Best for general purpose background removal - -- INSPYRENET ([Homepage](https://github.com/plemeri/InSPyReNet)) - - Specialized in human portrait segmentation - - Fast processing speed - - Good edge detection capability - - Ideal for portrait photos and human subjects - -- BEN (Background Elimination Network) ([Homepage](https://huggingface.co/PramaLLC/BEN)) - - Robust performance on various image types - - Good balance between speed and accuracy - - Effective on both simple and complex scenes - - Suitable for batch processing - -### Features -- Unified interface for all three models -- Common parameters for all models: - - Sensitivity adjustment - - Processing resolution control - - Mask blur and offset options - - Multiple background color options - - Invert output option - - Model optimization toggle - -### Improvements -- Optimized memory usage with model clearing -- Enhanced error handling and user feedback -- Added detailed tooltips for all parameters -- Improved mask post-processing - -### Dependencies -- Updated all package dependencies to latest stable versions -- Added support for transparent-background package -- Optimized dependency management - -## v1.1.0 (2024/11/21) - -### New Features -- Added background color options - - Alpha (transparent background) - - Black, White, Green, Blue, Red - -![RMBG_v1 1 0](https://github.com/user-attachments/assets/b7cbadff-5386-4d96-bc34-a19ad34efb4b) - -- Improved mask processing - - Better detail preservation - - Enhanced edge quality - - More accurate segmentation - -![rmbg version compare](https://github.com/user-attachments/assets/8339aa8e-46db-4f11-aa7b-0a710f0a1711) - -- Added video batch processing - - Support for video file background removal - - Maintains original video framerate and resolution - - Multiple output format support (with Alpha channel) - - Efficient batch processing for video frames - -https://github.com/user-attachments/assets/259220d3-c148-4030-93d6-c17dd5bccee1 - -- Added model cache management - - Cache status checking - - Model memory cleanup - - Better error handling - -### Parameter Updates -- Renamed 'invert_mask' to 'invert_output' for clarity -- Added sensitivity adjustment for mask strength -- Updated tooltips for better clarity - -### Technical Improvements -- Optimized image processing pipeline -- Added proper model cache verification -- Improved memory management -- Better error handling and recovery -- Enhanced batch processing performance for videos - -### Dependencies -- Added timm>=0.6.12,<1.0.0 for model support -- Updated requirements.txt with version constraints - -### Bug Fixes -- Fixed mask detail preservation issues -- Improved mask edge quality -- Fixed memory leaks in model handling - -### Usage Notes -- The 'Alpha' background option provides transparent background -- Sensitivity parameter now controls mask strength -- Model cache is checked before each operation -- Memory is automatically cleaned when switching models -- Video processing supports various formats and maintains quality +# ComfyUI-RMBG Update Log + +## v1.7.0 (2024/01/05) + +### New Model Added: BEN2 +- Added support for BEN2 (Background Elimination Network 2) + - Improved performance over original BEN model + - Better edge detection and detail preservation + - Enhanced batch processing capabilities (up to 3 images per batch) + - Optimized memory usage and processing speed + +### Model Changes +- Updated model repository paths for BEN and BEN2 +- Switched to 1038lab repositories for better maintenance and updates +- Maintained full compatibility with existing workflows + +### Technical Improvements +- Implemented efficient batch processing for BEN2 +- Optimized memory management for large batches +- Enhanced error handling and model loading +- Improved model switching and resource cleanup + +### Comparison with Previous Models +- BEN2 vs BEN: + - Better edge detection + - Improved handling of complex backgrounds + - More efficient batch processing + - Enhanced detail preservation + - Faster processing speed + +### Repository Updates +- Updated documentation to include BEN2 model +- Added new model license information +- Improved installation instructions +- Updated version number to 1.7.0 + +## v1.6.0 (2025/01/22) + +### New Face Segment Custom Node +- Added a new custom node for face parsing and segmentation + - Support for 19 facial feature categories (Skin, Nose, Eyes, Eyebrows, etc.) + - Precise facial feature extraction and segmentation + - Multiple feature selection for combined segmentation + - Same parameter controls as other RMBG nodes + - Automatic model downloading and resource management + - Perfect for portrait editing and facial feature manipulation + +![RMBG_v1 6 0](https://github.com/user-attachments/assets/9ccefec1-4370-4708-a12d-544c90888bf2) + +## v1.5.0 (2025/01/05) + +### New Fashion and accessories Segment Custom Node +- Added a new custom node for fashion and accessories segmentation. + - Capable of identifying and segmenting various fashion items such as dresses, shoes, and accessories. + - Utilizes advanced machine learning techniques for accurate segmentation. + - Supports real-time processing for enhanced user experience. + - Ideal for fashion-related applications, including virtual try-ons and outfit recommendations. + - Support for gray background color. + +![RMBGv_1 5 0](https://github.com/user-attachments/assets/a250c1a6-8425-4902-b902-a6e1a8bfe959) + +## v1.4.0 (2025/01/02) + +### New Clothes Segment Node +- Added intelligent clothes segmentation functionality + - Support for 18 different clothing categories (Hat, Hair, Face, Sunglasses, Upper-clothes, etc.) + - Multiple item selection for combined segmentation + - Same parameter controls as other RMBG nodes (process_res, mask_blur, mask_offset, background options) + - Automatic model downloading and resource management + +![rmbg_v1 4 0](https://github.com/user-attachments/assets/978c168b-03a8-4937-aa03-06385f34b820) + +## v1.3.2 (2024/12/29) + +### Updates +- Enhanced background handling to support RGBA output when "Alpha" is selected. +- Ensured RGB output for all other background color selections. + +## v1.3.1 (2024/12/25) + +### Bug Fixes +- Fixed an issue with mask processing when the model returns a list of masks. +- Improved handling of image formats to prevent processing errors. + +## v1.3.0 (2024/12/23) + +### New Segment (RMBG) Node +- Text-Prompted Intelligent Object Segmentation + - Use natural language prompts (e.g., "a cat", "red car") to identify and segment target objects + - Support for multiple object detection and segmentation + - Perfect for precise object extraction and recognition tasks + +![rmbg v1.3.0](https://github.com/user-attachments/assets/7607546e-ffcb-45e2-ab90-83267292757e) + +### Supported Models +- SAM (Segment Anything Model) + - sam_vit_h: 2.56GB - Highest accuracy + - sam_vit_l: 1.25GB - Balanced performance + - sam_vit_b: 375MB - Lightweight option +- GroundingDINO + - SwinT: 694MB - Fast and efficient + - SwinB: 938MB - Higher precision + +### Key Features +- Intuitive Parameter Controls + - Threshold: Adjust detection precision + - Mask Blur: Smooth edges + - Mask Offset: Expand or shrink selection + - Background Options: Alpha/Black/White/Green/Blue/Red +- Automatic Model Management + - Auto-download models on first use + - Smart GPU memory handling + +### Usage Examples +1. Tag-Style Prompts + - Single object: "cat" + - Multiple objects: "cat, dog, person" + - With attributes: "red car, blue shirt" + - Format: Use commas to separate multiple objects (e.g., "a, b, c") + +2. Natural Language Prompts + - Simple sentence: "a person wearing a red jacket" + - Complex scene: "a woman in a blue dress standing next to a car" + - With location: "a cat sitting on the sofa" + - Format: Write a natural descriptive sentence + +3. Tips for Better Results + - For Tag Style: + - Separate objects with commas: "chair, table, lamp" + - Add attributes before objects: "wooden chair, glass table" + - Keep it simple and clear + - For Natural Language: + - Use complete sentences + - Include details like color, position, action + - Be as descriptive as needed + - Parameter Adjustments: + - Threshold: 0.25-0.35 for broad detection, 0.45-0.55 for precision + - Use mask blur for smoother edges + - Adjust mask offset to fine-tune selection + +## v1.2.2 (2024/12/12) +![RMBG1 2 2](https://github.com/user-attachments/assets/cb7b1ad0-a2ca-4369-9401-54957af6c636) + +### Improvements +- Changed INSPYRENET model format from .pth to .safetensors for: + - Better security + - Faster loading speed (2-3x faster) + - Improved memory efficiency + - Better cross-platform compatibility +- Simplified node display name for better UI integration + +## v1.2.1 (2024/12/02) + +### New Features +- ANPG (animated PNG), AWEBP (animated WebP) and GIF supported. + +https://github.com/user-attachments/assets/40ec0b27-4fa2-4c99-9aea-5afad9ca62a5 + +### Bug Fixes +- Fixed video processing issue + +### Performance Improvements +- Enhanced batch processing in RMBG-2.0 model +- Added support for proper batch image handling +- Improved memory efficiency by optimizing image size handling + +### Technical Details +- Added original size preservation for maintaining aspect ratios +- Implemented proper batch tensor processing +- Improved error handling and code robustness +- Performance gains: + - Single image processing: ~5-10% improvement + - Batch processing: up to 30-50% improvement (depending on batch size and GPU) + +## v1.2.0 (2024/11/29) + +### Major Changes +- Combined three background removal models into one unified node +- Added support for RMBG-2.0, INSPYRENET, and BEN models +- Implemented lazy loading for models (only downloads when first used) + +### Model Introduction +- RMBG-2.0 ([Homepage](https://huggingface.co/briaai/RMBG-2.0)) + - Latest version of RMBG model + - Excellent performance on complex backgrounds + - High accuracy in preserving fine details + - Best for general purpose background removal + +- INSPYRENET ([Homepage](https://github.com/plemeri/InSPyReNet)) + - Specialized in human portrait segmentation + - Fast processing speed + - Good edge detection capability + - Ideal for portrait photos and human subjects + +- BEN (Background Elimination Network) ([Homepage](https://huggingface.co/PramaLLC/BEN)) + - Robust performance on various image types + - Good balance between speed and accuracy + - Effective on both simple and complex scenes + - Suitable for batch processing + +### Features +- Unified interface for all three models +- Common parameters for all models: + - Sensitivity adjustment + - Processing resolution control + - Mask blur and offset options + - Multiple background color options + - Invert output option + - Model optimization toggle + +### Improvements +- Optimized memory usage with model clearing +- Enhanced error handling and user feedback +- Added detailed tooltips for all parameters +- Improved mask post-processing + +### Dependencies +- Updated all package dependencies to latest stable versions +- Added support for transparent-background package +- Optimized dependency management + +## v1.1.0 (2024/11/21) + +### New Features +- Added background color options + - Alpha (transparent background) + - Black, White, Green, Blue, Red + +![RMBG_v1 1 0](https://github.com/user-attachments/assets/b7cbadff-5386-4d96-bc34-a19ad34efb4b) + +- Improved mask processing + - Better detail preservation + - Enhanced edge quality + - More accurate segmentation + +![rmbg version compare](https://github.com/user-attachments/assets/8339aa8e-46db-4f11-aa7b-0a710f0a1711) + +- Added video batch processing + - Support for video file background removal + - Maintains original video framerate and resolution + - Multiple output format support (with Alpha channel) + - Efficient batch processing for video frames + +https://github.com/user-attachments/assets/259220d3-c148-4030-93d6-c17dd5bccee1 + +- Added model cache management + - Cache status checking + - Model memory cleanup + - Better error handling + +### Parameter Updates +- Renamed 'invert_mask' to 'invert_output' for clarity +- Added sensitivity adjustment for mask strength +- Updated tooltips for better clarity + +### Technical Improvements +- Optimized image processing pipeline +- Added proper model cache verification +- Improved memory management +- Better error handling and recovery +- Enhanced batch processing performance for videos + +### Dependencies +- Added timm>=0.6.12,<1.0.0 for model support +- Updated requirements.txt with version constraints + +### Bug Fixes +- Fixed mask detail preservation issues +- Improved mask edge quality +- Fixed memory leaks in model handling + +### Usage Notes +- The 'Alpha' background option provides transparent background +- Sensitivity parameter now controls mask strength +- Model cache is checked before each operation +- Memory is automatically cleaned when switching models +- Video processing supports various formats and maintains quality