Add files via upload
This commit is contained in:
+441
-441
@@ -1,442 +1,442 @@
|
||||
# ComfyUI-RMBG v1.2.2
|
||||
# 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.
|
||||
|
||||
# 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:
|
||||
"""
|
||||
RMBG Node: Advanced Background Removal Suite
|
||||
|
||||
This node provides professional background removal capabilities using three state-of-the-art models:
|
||||
- RMBG-2.0: Latest model with excellent performance on complex backgrounds
|
||||
- INSPYRENET: Specialized for human portrait segmentation
|
||||
- BEN: Versatile model with good balance of speed and accuracy
|
||||
|
||||
Features:
|
||||
- Batch processing support
|
||||
- Multiple background options
|
||||
- Advanced mask refinement
|
||||
- High-quality edge preservation
|
||||
- Memory-efficient processing
|
||||
"""
|
||||
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", "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),
|
||||
"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)
|
||||
|
||||
# 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)
|
||||
processed_images.append(pil2tensor(composite_image))
|
||||
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": "RMBG (Background Remover)"
|
||||
# ComfyUI-RMBG v1.3.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.
|
||||
|
||||
# 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:
|
||||
"""
|
||||
RMBG Node: Advanced Background Removal Suite
|
||||
|
||||
This node provides professional background removal capabilities using three state-of-the-art models:
|
||||
- RMBG-2.0: Latest model with excellent performance on complex backgrounds
|
||||
- INSPYRENET: Specialized for human portrait segmentation
|
||||
- BEN: Versatile model with good balance of speed and accuracy
|
||||
|
||||
Features:
|
||||
- Batch processing support
|
||||
- Multiple background options
|
||||
- Advanced mask refinement
|
||||
- High-quality edge preservation
|
||||
- Memory-efficient processing
|
||||
"""
|
||||
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", "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),
|
||||
"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)
|
||||
|
||||
# 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)
|
||||
processed_images.append(pil2tensor(composite_image))
|
||||
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": "RMBG (Background Remover)"
|
||||
}
|
||||
@@ -0,0 +1,342 @@
|
||||
# ComfyUI-RMBG v1.3.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.
|
||||
|
||||
# License Notice:
|
||||
# - SAM: MIT License (https://github.com/facebookresearch/segment-anything)
|
||||
# - GroundingDINO: MIT License (https://github.com/IDEA-Research/GroundingDINO)
|
||||
|
||||
# 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 sys
|
||||
import copy
|
||||
import requests
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from PIL import ImageFilter
|
||||
from torch.hub import download_url_to_file
|
||||
|
||||
import folder_paths
|
||||
import comfy.model_management
|
||||
from segment_anything import sam_model_registry, SamPredictor
|
||||
|
||||
SAM_MODELS = {
|
||||
"sam_vit_h (2.56GB)": {
|
||||
"model_url": "https://huggingface.co/1038lab/sam/resolve/main/sam_vit_h.pth",
|
||||
"model_type": "vit_h"
|
||||
},
|
||||
"sam_vit_l (1.25GB)": {
|
||||
"model_url": "https://huggingface.co/1038lab/sam/resolve/main/sam_vit_l.pth",
|
||||
"model_type": "vit_l"
|
||||
},
|
||||
"sam_vit_b (375MB)": {
|
||||
"model_url": "https://huggingface.co/1038lab/sam/resolve/main/sam_vit_b.pth",
|
||||
"model_type": "vit_b"
|
||||
}
|
||||
}
|
||||
|
||||
DINO_MODELS = {
|
||||
"GroundingDINO_SwinT_OGC (694MB)": {
|
||||
"config_url": "https://huggingface.co/1038lab/GroundingDINO/resolve/main/GroundingDINO_SwinT_OGC.cfg.py",
|
||||
"model_url": "https://huggingface.co/1038lab/GroundingDINO/resolve/main/groundingdino_swint_ogc.pth",
|
||||
},
|
||||
"GroundingDINO_SwinB (938MB)": {
|
||||
"config_url": "https://huggingface.co/1038lab/GroundingDINO/resolve/main/GroundingDINO_SwinB.cfg.py",
|
||||
"model_url": "https://huggingface.co/1038lab/GroundingDINO/resolve/main/groundingdino_swinb_cogcoor.pth"
|
||||
}
|
||||
}
|
||||
|
||||
def normalize_array(arr):
|
||||
return arr.astype(np.float32) / 255.0
|
||||
|
||||
def denormalize_array(arr):
|
||||
return np.clip(255. * arr, 0, 255).astype(np.uint8)
|
||||
|
||||
def create_tensor_output(image_np, masks, boxes_filt):
|
||||
output_masks, output_images = [], []
|
||||
for mask in masks:
|
||||
image_np_copy = copy.deepcopy(image_np)
|
||||
image_np_copy[~np.any(mask, axis=0)] = np.array([0, 0, 0, 0])
|
||||
output_image, output_mask = split_image_mask(
|
||||
Image.fromarray(image_np_copy))
|
||||
output_masks.append(output_mask)
|
||||
output_images.append(output_image)
|
||||
return (torch.cat(output_images, dim=0), torch.cat(output_masks, dim=0))
|
||||
|
||||
def split_image_mask(image):
|
||||
image_rgb = image.convert("RGB")
|
||||
image_rgb = np.array(image_rgb).astype(np.float32) / 255.0
|
||||
image_rgb = torch.from_numpy(image_rgb)[None,]
|
||||
if 'A' in image.getbands():
|
||||
mask = np.array(image.getchannel('A')).astype(np.float32) / 255.0
|
||||
mask = torch.from_numpy(mask)[None,]
|
||||
else:
|
||||
mask = torch.zeros((image.height, image.width), dtype=torch.float32, device="cpu")[None,]
|
||||
return (image_rgb, mask)
|
||||
|
||||
def process_mask(mask_image: Image.Image, invert_output: bool = False,
|
||||
mask_blur: int = 0, mask_offset: int = 0) -> Image.Image:
|
||||
if invert_output:
|
||||
mask_np = np.array(mask_image)
|
||||
mask_image = Image.fromarray(255 - mask_np)
|
||||
|
||||
if mask_blur > 0:
|
||||
mask_image = mask_image.filter(ImageFilter.GaussianBlur(radius=mask_blur))
|
||||
|
||||
if mask_offset != 0:
|
||||
filter_type = ImageFilter.MaxFilter if mask_offset > 0 else ImageFilter.MinFilter
|
||||
size = abs(mask_offset) * 2 + 1
|
||||
for _ in range(abs(mask_offset)):
|
||||
mask_image = mask_image.filter(filter_type(size))
|
||||
|
||||
return mask_image
|
||||
|
||||
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):
|
||||
if image.mode != 'L':
|
||||
image = image.convert('L')
|
||||
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0)
|
||||
return image.squeeze()
|
||||
|
||||
def apply_background_color(image: Image.Image, mask_image: Image.Image,
|
||||
background_color: str = "Alpha") -> Image.Image:
|
||||
bg_colors = {
|
||||
"Alpha": None,
|
||||
"black": (0, 0, 0),
|
||||
"white": (255, 255, 255),
|
||||
"green": (0, 255, 0),
|
||||
"blue": (0, 0, 255),
|
||||
"red": (255, 0, 0)
|
||||
}
|
||||
|
||||
rgba_image = image.copy().convert('RGBA')
|
||||
rgba_image.putalpha(mask_image.convert('L'))
|
||||
|
||||
if background_color != "Alpha":
|
||||
bg_color = bg_colors[background_color]
|
||||
bg_image = Image.new('RGBA', image.size, (*bg_color, 255))
|
||||
composite_image = Image.alpha_composite(bg_image, rgba_image)
|
||||
return composite_image.convert('RGB')
|
||||
|
||||
return rgba_image
|
||||
|
||||
class Segment:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
tooltips = {
|
||||
"prompt": "Enter the object or scene you want to segment. Use tag-style or natural language for more detailed prompts.",
|
||||
"threshold": "Adjust mask detection strength (higher = more strict)",
|
||||
"mask_blur": "Apply Gaussian blur to mask edges (0 = disabled)",
|
||||
"mask_offset": "Expand/Shrink mask boundary (positive = expand, negative = shrink)",
|
||||
"background_color": "Choose background color (Alpha = transparent)",
|
||||
"invert_output": "Invert the mask output",
|
||||
}
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"prompt": ("STRING", {"default": "", "multiline": True, "placeholder": "Object to segment", "tooltip": tooltips["prompt"]}),
|
||||
"sam_model": (list(SAM_MODELS.keys()),),
|
||||
"dino_model": (list(DINO_MODELS.keys()),),
|
||||
},
|
||||
"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"]}),
|
||||
"background_color": (["Alpha", "black", "white", "green", "blue", "red"], {"default": "Alpha", "tooltip": tooltips["background_color"]}),
|
||||
"invert_output": ("BOOLEAN", {"default": False}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK")
|
||||
FUNCTION = "segment"
|
||||
CATEGORY = "🧪AILab/🧽RMBG"
|
||||
|
||||
def __init__(self):
|
||||
from groundingdino.datasets import transforms as T
|
||||
from groundingdino.util.utils import clean_state_dict
|
||||
from groundingdino.util.slconfig import SLConfig
|
||||
from groundingdino.models import build_model
|
||||
|
||||
self.T = T
|
||||
self.clean_state_dict = clean_state_dict
|
||||
self.SLConfig = SLConfig
|
||||
self.build_model = build_model
|
||||
|
||||
def segment(self, image, prompt, sam_model, dino_model, threshold=0.35,
|
||||
mask_blur=0, mask_offset=0, background_color="Alpha",
|
||||
invert_output=False):
|
||||
print(f'Processing create segment for: "{prompt}"...')
|
||||
|
||||
image = Image.fromarray(np.clip(255. * image[0].cpu().numpy(), 0, 255).astype(np.uint8)).convert('RGBA')
|
||||
dino_model = self.load_groundingdino(dino_model)
|
||||
sam_model = self.load_sam(sam_model)
|
||||
boxes = self.predict_boxes(dino_model, image, prompt, threshold)
|
||||
|
||||
if boxes is None or boxes.shape[0] == 0:
|
||||
print(f'No objects found for: "{prompt}"')
|
||||
width, height = image.size
|
||||
empty_mask = torch.zeros((1, height, width), dtype=torch.uint8, device="cpu")
|
||||
return (empty_mask, empty_mask)
|
||||
|
||||
masks = self.generate_masks(sam_model, image, boxes)
|
||||
if masks is None:
|
||||
print(f'Failed to generate mask for: "{prompt}"')
|
||||
width, height = image.size
|
||||
empty_mask = torch.zeros((1, height, width), dtype=torch.uint8, device="cpu")
|
||||
return (empty_mask, empty_mask)
|
||||
|
||||
mask_image = Image.fromarray((masks[1][0].numpy() * 255).astype(np.uint8))
|
||||
mask_image = process_mask(mask_image, invert_output, mask_blur, mask_offset)
|
||||
|
||||
result_image = apply_background_color(image, mask_image, background_color)
|
||||
|
||||
print(f'Successfully created segment for: "{prompt}"')
|
||||
return (pil2tensor(result_image), image2mask(mask_image))
|
||||
|
||||
def load_sam(self, model_name):
|
||||
sam_checkpoint_path = self.get_local_filepath(
|
||||
SAM_MODELS[model_name]["model_url"], "sam")
|
||||
model_type = SAM_MODELS[model_name]["model_type"]
|
||||
|
||||
sam = sam_model_registry[model_type](checkpoint=sam_checkpoint_path)
|
||||
sam_device = comfy.model_management.get_torch_device()
|
||||
sam.to(device=sam_device)
|
||||
sam.eval()
|
||||
return sam
|
||||
|
||||
def load_groundingdino(self, model_name):
|
||||
import sys
|
||||
from io import StringIO
|
||||
temp_stdout = StringIO()
|
||||
original_stdout = sys.stdout
|
||||
sys.stdout = temp_stdout
|
||||
|
||||
try:
|
||||
dino_model_args = self.SLConfig.fromfile(
|
||||
self.get_local_filepath(
|
||||
DINO_MODELS[model_name]["config_url"],
|
||||
"grounding-dino"
|
||||
)
|
||||
)
|
||||
dino = self.build_model(dino_model_args)
|
||||
checkpoint = torch.load(
|
||||
self.get_local_filepath(
|
||||
DINO_MODELS[model_name]["model_url"],
|
||||
"grounding-dino"
|
||||
)
|
||||
)
|
||||
dino.load_state_dict(self.clean_state_dict(checkpoint['model']), strict=False)
|
||||
device = comfy.model_management.get_torch_device()
|
||||
dino.to(device=device)
|
||||
dino.eval()
|
||||
return dino
|
||||
finally:
|
||||
output = temp_stdout.getvalue()
|
||||
sys.stdout = original_stdout
|
||||
|
||||
for line in output.split('\n'):
|
||||
if 'error' in line.lower():
|
||||
print(line)
|
||||
|
||||
def _load_dino_image(self, image_pil):
|
||||
transform = self.T.Compose([
|
||||
self.T.RandomResize([800], max_size=1333),
|
||||
self.T.ToTensor(),
|
||||
self.T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
|
||||
])
|
||||
image, _ = transform(image_pil, None)
|
||||
return image
|
||||
|
||||
def _get_grounding_output(self, model, image, caption, box_threshold):
|
||||
caption = caption.lower().strip()
|
||||
if not caption.endswith("."):
|
||||
caption = caption + "."
|
||||
device = comfy.model_management.get_torch_device()
|
||||
image = image.to(device)
|
||||
with torch.no_grad():
|
||||
outputs = model(image[None], captions=[caption])
|
||||
logits = outputs["pred_logits"].sigmoid()[0]
|
||||
boxes = outputs["pred_boxes"][0]
|
||||
logits_filt = logits.clone()
|
||||
boxes_filt = boxes.clone()
|
||||
filt_mask = logits_filt.max(dim=1)[0] > box_threshold
|
||||
logits_filt = logits_filt[filt_mask]
|
||||
boxes_filt = boxes_filt[filt_mask]
|
||||
return boxes_filt.cpu()
|
||||
|
||||
def predict_boxes(self, model, image, prompt, threshold):
|
||||
dino_image = self._load_dino_image(image.convert("RGB"))
|
||||
boxes_filt = self._get_grounding_output(model, dino_image, prompt, threshold)
|
||||
H, W = image.size[1], image.size[0]
|
||||
for i in range(boxes_filt.size(0)):
|
||||
boxes_filt[i] = boxes_filt[i] * torch.Tensor([W, H, W, H])
|
||||
boxes_filt[i][:2] -= boxes_filt[i][2:] / 2
|
||||
boxes_filt[i][2:] += boxes_filt[i][:2]
|
||||
return boxes_filt
|
||||
|
||||
def generate_masks(self, model, image, boxes):
|
||||
if boxes.shape[0] == 0:
|
||||
return None
|
||||
|
||||
if not hasattr(self, 'predictor'):
|
||||
self.predictor = SamPredictor(model)
|
||||
|
||||
image_np = np.array(image)
|
||||
image_np_rgb = image_np[..., :3]
|
||||
|
||||
self.predictor.set_image(image_np_rgb)
|
||||
|
||||
transformed_boxes = self.predictor.transform.apply_boxes_torch(boxes, image_np.shape[:2])
|
||||
masks, _, _ = self.predictor.predict_torch(
|
||||
point_coords=None,
|
||||
point_labels=None,
|
||||
boxes=transformed_boxes.to(comfy.model_management.get_torch_device()),
|
||||
multimask_output=False
|
||||
)
|
||||
|
||||
return create_tensor_output(image_np, masks.permute(1, 0, 2, 3).cpu().numpy(), boxes)
|
||||
|
||||
|
||||
def get_local_filepath(self, url, dirname, local_file_name=None):
|
||||
if not local_file_name:
|
||||
local_file_name = os.path.basename(urlparse(url).path)
|
||||
|
||||
destination = folder_paths.get_full_path(dirname, local_file_name)
|
||||
if destination:
|
||||
return destination
|
||||
|
||||
folder = os.path.join(folder_paths.models_dir, dirname)
|
||||
os.makedirs(folder, exist_ok=True)
|
||||
|
||||
destination = os.path.join(folder, local_file_name)
|
||||
if not os.path.exists(destination):
|
||||
try:
|
||||
download_url_to_file(url, destination)
|
||||
except Exception as e:
|
||||
if os.path.exists(destination):
|
||||
os.remove(destination)
|
||||
raise Exception(f'Failed to download model from {url}: {str(e)}')
|
||||
return destination
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"Segment": Segment
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"Segment": "Segment (RMBG)"
|
||||
}
|
||||
+3
-1
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "ComfyUI-RMBG"
|
||||
version = "1.2.1"
|
||||
version = "1.3.0
|
||||
description = "A ComfyUI custom node designed for advanced image background removal utilizing multiple models, including RMBG-2.0, INSPYRENET, and BEN."
|
||||
authors = [
|
||||
{name = "AILab", email = "ailab@mail.com"}
|
||||
@@ -14,6 +14,8 @@ dependencies = [
|
||||
"tqdm>=4.65.0",
|
||||
"transformers>=4.35.0",
|
||||
"transparent-background>=1.2.4",
|
||||
"segment-anything>=1.0.0",
|
||||
"groundingdino>=1.0.0",
|
||||
]
|
||||
requires-python = ">=3.8"
|
||||
readme = "README.md"
|
||||
|
||||
+4
-1
@@ -5,4 +5,7 @@ numpy>=1.22.0
|
||||
huggingface-hub>=0.19.0
|
||||
tqdm>=4.65.0
|
||||
transformers>=4.35.0
|
||||
transparent-background>=1.2.4
|
||||
transparent-background>=1.2.4
|
||||
groundingdino-py>=0.4.0
|
||||
segment-anything>=1.0
|
||||
opencv-python>=4.7.0
|
||||
@@ -1,139 +1,193 @@
|
||||
# ComfyUI-RMBG Update Log
|
||||
|
||||
## v1.2.2 (2024/12/12)
|
||||

|
||||
|
||||
### 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
|
||||
|
||||

|
||||
|
||||
- Improved mask processing
|
||||
- Better detail preservation
|
||||
- Enhanced edge quality
|
||||
- More accurate segmentation
|
||||
|
||||

|
||||
|
||||
- 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.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
|
||||
|
||||
### 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)
|
||||

|
||||
|
||||
### 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
|
||||
|
||||

|
||||
|
||||
- Improved mask processing
|
||||
- Better detail preservation
|
||||
- Enhanced edge quality
|
||||
- More accurate segmentation
|
||||
|
||||

|
||||
|
||||
- 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
|
||||
|
||||
Reference in New Issue
Block a user