Add files via upload
This commit is contained in:
+283
-22
@@ -1,4 +1,4 @@
|
||||
# ComfyUI-RMBG v2.3.1
|
||||
# ComfyUI-RMBG v2.4.0
|
||||
#
|
||||
# This node facilitates background removal using various models, including RMBG-2.0, INSPYRENET, BEN, BEN2, and BIREFNET-HR.
|
||||
# It utilizes advanced deep learning techniques to process images and generate accurate masks for background removal.
|
||||
@@ -15,6 +15,7 @@
|
||||
#
|
||||
# 2. Conversion Node:
|
||||
# - ImageMaskConvert: Converts between image and mask formats and extracts masks from image channels.
|
||||
# - ColorInput: A node for inputting colors in various formats.
|
||||
#
|
||||
# 3. Mask Processing Nodes:
|
||||
# - MaskEnhancer: Refines masks through techniques such as blur, smoothing, expansion/contraction, and hole filling.
|
||||
@@ -25,6 +26,8 @@
|
||||
# - ImageStitch: Stitches multiple images together in various directions.
|
||||
# - ImageCrop: Crops an image to a specified size and position.
|
||||
# - ICLoRAConcat: Concatenates images with a mask using IC LoRA.
|
||||
# - CropObject: Crops an image to the object in the image.
|
||||
# - ImageCompare: Compares two images and returns a mask of the differences.
|
||||
|
||||
# These nodes are crafted to streamline common image and mask operations within ComfyUI workflows.
|
||||
|
||||
@@ -36,7 +39,7 @@ import hashlib
|
||||
import torch
|
||||
import cv2
|
||||
from nodes import MAX_RESOLUTION
|
||||
from PIL import Image, ImageFilter, ImageOps, ImageSequence, ImageChops
|
||||
from PIL import Image, ImageFilter, ImageOps, ImageSequence, ImageChops, ImageDraw, ImageFont
|
||||
import torchvision.transforms.functional as T
|
||||
from comfy.utils import common_upscale
|
||||
from scipy import ndimage
|
||||
@@ -92,6 +95,38 @@ def ensure_mask_shape(mask):
|
||||
return mask.squeeze(1)
|
||||
return mask
|
||||
|
||||
def resize_image(img: Image.Image, width: int, height: int) -> Image.Image:
|
||||
return img.resize((width, height), Image.Resampling.LANCZOS)
|
||||
|
||||
COLOR_PRESETS = {
|
||||
"black": "#000000", "white": "#FFFFFF", "red": "#FF0000", "green": "#00FF00", "blue": "#0000FF",
|
||||
"yellow": "#FFFF00", "cyan": "#00FFFF", "magenta": "#FF00FF", "gray": "#808080", "silver": "#C0C0C0",
|
||||
"maroon": "#800000", "olive": "#808000", "purple": "#800080", "teal": "#008080", "navy": "#000080",
|
||||
"orange": "#FFA500", "pink": "#FFC0CB", "brown": "#A52A2A", "violet": "#EE82EE", "indigo": "#4B0082",
|
||||
"light_gray": "#D3D3D3", "dark_gray": "#A9A9A9", "light_blue": "#ADD8E6", "dark_blue": "#00008B",
|
||||
"light_blue": "#ADD8E6", "dark_blue": "#00008B", "light_green": "#90EE90", "dark_green": "#006400"
|
||||
}
|
||||
|
||||
def fix_color_format(color: str) -> str:
|
||||
"""Fix color format to valid hex code"""
|
||||
if not color:
|
||||
return ""
|
||||
|
||||
color = color.strip().upper()
|
||||
if not color.startswith('#'):
|
||||
color = f"#{color}"
|
||||
|
||||
color = color[1:]
|
||||
if len(color) == 3:
|
||||
r, g, b = color[0], color[1], color[2]
|
||||
return f"#{r}{r}{g}{g}{b}{b}"
|
||||
elif len(color) < 6:
|
||||
raise ValueError(f"Invalid color format: {color}")
|
||||
elif len(color) > 6:
|
||||
color = color[:6]
|
||||
|
||||
return f"#{color}"
|
||||
|
||||
# Base class for preview
|
||||
class AILab_PreviewBase:
|
||||
def __init__(self):
|
||||
@@ -609,7 +644,8 @@ class AILab_ImageCombiner:
|
||||
}
|
||||
|
||||
CATEGORY = "🧪AILab/🖼️IMAGE"
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_TYPES = ("IMAGE", "INT", "INT")
|
||||
RETURN_NAMES = ("IMAGE", "WIDTH", "HEIGHT")
|
||||
FUNCTION = "combine_images"
|
||||
|
||||
def combine_images(self, foreground, background, mode="normal", foreground_opacity=1.0,
|
||||
@@ -695,7 +731,11 @@ class AILab_ImageCombiner:
|
||||
|
||||
output_images.append(pil2tensor(result))
|
||||
|
||||
return (torch.cat(output_images, dim=0),)
|
||||
final_image = torch.cat(output_images, dim=0)
|
||||
width = final_image.shape[2]
|
||||
height = final_image.shape[1]
|
||||
|
||||
return (final_image, width, height)
|
||||
|
||||
# Mask extractor node
|
||||
class AILab_MaskExtractor:
|
||||
@@ -704,9 +744,12 @@ class AILab_MaskExtractor:
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"mode": (["extract_masked_area", "apply_mask", "invert_mask"], {"default": "extract_masked_area"}),
|
||||
"background": (["Alpha", "original", "Color"], {"default": "Alpha", "tooltip": "Choose background type"}),
|
||||
"background_color": ("COLOR", {"default": "#FFFFFF", "tooltip": "Choose background color (Alpha = transparent)"})
|
||||
},
|
||||
"optional": {
|
||||
"mask": ("MASK",),
|
||||
"mode": (["extract_masked_area", "apply_mask", "invert_mask"], {"default": "invert_mask"}),
|
||||
"background": (["transparent", "black", "white", "original"], {"default": "transparent"})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -736,8 +779,22 @@ class AILab_MaskExtractor:
|
||||
print(f"Error in _prepare_mask: {str(e)}")
|
||||
raise e
|
||||
|
||||
def extract_masked_area(self, image, mask, mode="extract_masked_area", background="transparent"):
|
||||
def hex_to_rgb(self, hex_color):
|
||||
hex_color = hex_color.lstrip('#')
|
||||
r = int(hex_color[0:2], 16) / 255.0
|
||||
g = int(hex_color[2:4], 16) / 255.0
|
||||
b = int(hex_color[4:6], 16) / 255.0
|
||||
return (r, g, b)
|
||||
|
||||
def extract_masked_area(self, image, mode="extract_masked_area", background="Alpha", background_color="#FFFFFF", mask=None):
|
||||
try:
|
||||
if mask is None and image.shape[-1] == 4:
|
||||
alpha = image[..., 3]
|
||||
mask = 1.0 - alpha
|
||||
image = image[..., :3]
|
||||
elif mask is None:
|
||||
mask = torch.ones((image.shape[0], image.shape[1], image.shape[2]), dtype=torch.float32)
|
||||
|
||||
pil_image = tensor2pil(image)
|
||||
image_np = np.array(pil_image).astype(np.float32) / 255.0
|
||||
mask_np = self._prepare_mask(mask, image_np.shape)
|
||||
@@ -745,7 +802,7 @@ class AILab_MaskExtractor:
|
||||
|
||||
if mode == "extract_masked_area":
|
||||
result_np = image_np * mask_np
|
||||
if background == "transparent":
|
||||
if background == "Alpha":
|
||||
if pil_image.mode != "RGBA":
|
||||
pil_image = pil_image.convert("RGBA")
|
||||
result_rgba = np.zeros((*image_np.shape[:2], 4), dtype=np.float32)
|
||||
@@ -753,16 +810,15 @@ class AILab_MaskExtractor:
|
||||
result_rgba[:, :, 3] = mask_np[..., 0]
|
||||
result_pil = Image.fromarray((result_rgba * 255).astype(np.uint8), mode="RGBA")
|
||||
return (torch.from_numpy(np.array(result_pil).astype(np.float32) / 255.0).unsqueeze(0),)
|
||||
elif background == "black":
|
||||
pass # Already done with image_np * mask_np
|
||||
elif background == "white":
|
||||
result_np = result_np + (1 - mask_np)
|
||||
elif background == "original":
|
||||
result_np = image_np * mask_np
|
||||
elif background == "Color":
|
||||
r, g, b = self.hex_to_rgb(background_color)
|
||||
result_np = result_np + (1 - mask_np) * np.array([r, g, b])
|
||||
|
||||
elif mode == "apply_mask":
|
||||
result_np = image_np * mask_np
|
||||
if background == "transparent":
|
||||
if background == "Alpha":
|
||||
if pil_image.mode != "RGBA":
|
||||
pil_image = pil_image.convert("RGBA")
|
||||
result_rgba = np.zeros((*image_np.shape[:2], 4), dtype=np.float32)
|
||||
@@ -770,14 +826,15 @@ class AILab_MaskExtractor:
|
||||
result_rgba[:, :, 3] = mask_np[..., 0]
|
||||
result_pil = Image.fromarray((result_rgba * 255).astype(np.uint8), mode="RGBA")
|
||||
return (torch.from_numpy(np.array(result_pil).astype(np.float32) / 255.0).unsqueeze(0),)
|
||||
elif background == "white":
|
||||
result_np = result_np + (1 - mask_np)
|
||||
elif background == "original":
|
||||
result_np = image_np * mask_np + image_np * (1 - mask_np)
|
||||
elif background == "Color":
|
||||
r, g, b = self.hex_to_rgb(background_color)
|
||||
result_np = result_np + (1 - mask_np) * np.array([r, g, b])
|
||||
|
||||
elif mode == "invert_mask":
|
||||
result_np = image_np * (1 - mask_np)
|
||||
if background == "transparent":
|
||||
if background == "Alpha":
|
||||
if pil_image.mode != "RGBA":
|
||||
pil_image = pil_image.convert("RGBA")
|
||||
result_rgba = np.zeros((*image_np.shape[:2], 4), dtype=np.float32)
|
||||
@@ -785,10 +842,11 @@ class AILab_MaskExtractor:
|
||||
result_rgba[:, :, 3] = (1 - mask_np)[..., 0]
|
||||
result_pil = Image.fromarray((result_rgba * 255).astype(np.uint8), mode="RGBA")
|
||||
return (torch.from_numpy(np.array(result_pil).astype(np.float32) / 255.0).unsqueeze(0),)
|
||||
elif background == "white":
|
||||
result_np = result_np + mask_np
|
||||
elif background == "original":
|
||||
result_np = image_np * (1 - mask_np) + image_np * mask_np
|
||||
elif background == "Color":
|
||||
r, g, b = self.hex_to_rgb(background_color)
|
||||
result_np = result_np + mask_np * np.array([r, g, b])
|
||||
|
||||
result_pil = Image.fromarray(np.clip(result_np * 255, 0, 255).astype(np.uint8))
|
||||
return (pil2tensor(result_pil),)
|
||||
@@ -979,8 +1037,9 @@ class AILab_ICLoRAConcat:
|
||||
base_image = empty_image(object_image.shape[2], object_image.shape[1])
|
||||
base_mask = torch.full((1, object_image.shape[1], object_image.shape[2]), 1, dtype=torch.float32, device="cpu")
|
||||
elif base_image is not None and base_mask is None:
|
||||
raise ValueError("base_mask is required when base_image is provided")
|
||||
|
||||
# raise ValueError("base_mask is required when base_image is provided")
|
||||
base_mask = torch.full((1, object_image.shape[1], object_image.shape[2]), 1, dtype=torch.float32, device="cpu")
|
||||
|
||||
object_mask = ensure_mask_shape(object_mask)
|
||||
base_mask = ensure_mask_shape(base_mask)
|
||||
|
||||
@@ -1077,8 +1136,204 @@ class AILab_ICLoRAConcat:
|
||||
y = object_image.shape[1] if layout == 'top-bottom' else 0
|
||||
|
||||
return (image, OBJECT_MASK, BASE_MASK, out_w, out_h, x, y)
|
||||
|
||||
|
||||
class AILab_CropObject:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"optional": {
|
||||
"image": ("IMAGE",),
|
||||
"mask": ("MASK",),
|
||||
"padding": ("INT", {
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 256,
|
||||
"step": 1
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK")
|
||||
RETURN_NAMES = ("IMAGE", "MASK")
|
||||
FUNCTION = "crop_object"
|
||||
CATEGORY = "🧪AILab/🖼️IMAGE"
|
||||
|
||||
def get_bbox_from_tensor(self, tensor, padding):
|
||||
rows = torch.any(tensor > 0, dim=1)
|
||||
cols = torch.any(tensor > 0, dim=0)
|
||||
if not torch.any(rows) or not torch.any(cols):
|
||||
return None
|
||||
rmin, rmax = torch.where(rows)[0][[0, -1]]
|
||||
cmin, cmax = torch.where(cols)[0][[0, -1]]
|
||||
rmin = max(0, rmin - padding)
|
||||
rmax = min(tensor.shape[0] - 1, rmax + padding)
|
||||
cmin = max(0, cmin - padding)
|
||||
cmax = min(tensor.shape[1] - 1, cmax + padding)
|
||||
return rmin, rmax, cmin, cmax
|
||||
|
||||
def crop_object(self, image=None, mask=None, padding=0):
|
||||
if mask is None and image is None:
|
||||
raise ValueError("At least one of image or mask must be provided")
|
||||
bbox = None
|
||||
if mask is not None:
|
||||
mask_tensor = mask.squeeze()
|
||||
bbox = self.get_bbox_from_tensor(mask_tensor, padding)
|
||||
elif image is not None and image.shape[-1] == 4:
|
||||
alpha = image[0, :, :, 3]
|
||||
bbox = self.get_bbox_from_tensor(alpha, padding)
|
||||
if bbox is None:
|
||||
return (image, mask)
|
||||
rmin, rmax, cmin, cmax = bbox
|
||||
if mask is not None:
|
||||
cropped_mask = mask[:, rmin:rmax+1, cmin:cmax+1]
|
||||
else:
|
||||
if image is not None and image.shape[-1] == 4:
|
||||
alpha = image[0, rmin:rmax+1, cmin:cmax+1, 3]
|
||||
cropped_mask = alpha.unsqueeze(0)
|
||||
else:
|
||||
cropped_mask = None
|
||||
if image is not None:
|
||||
cropped_image = image[:, rmin:rmax+1, cmin:cmax+1, :]
|
||||
else:
|
||||
cropped_image = None
|
||||
return (
|
||||
cropped_image if image is not None else image,
|
||||
cropped_mask if mask is not None else mask
|
||||
)
|
||||
|
||||
# Image Compare node
|
||||
class AILab_ImageCompare:
|
||||
def __init__(self):
|
||||
self.font_size = 20
|
||||
self.padding = 10
|
||||
self.bg_color = "white"
|
||||
self.font_color = "black"
|
||||
self.text_align = "center"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image1": ("IMAGE",),
|
||||
"image2": ("IMAGE",),
|
||||
"text1": ("STRING", {"default": "image 1"}),
|
||||
"text2": ("STRING", {"default": "image 2"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "generate"
|
||||
CATEGORY = "🧪AILab/🖼️IMAGE"
|
||||
|
||||
def get_font(self) -> ImageFont.FreeTypeFont:
|
||||
try:
|
||||
if os.name == 'nt':
|
||||
return ImageFont.truetype("arial.ttf", self.font_size)
|
||||
else:
|
||||
return ImageFont.truetype("/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", self.font_size)
|
||||
except:
|
||||
base_font = ImageFont.load_default()
|
||||
scale_factor = self.font_size / 10
|
||||
return ImageFont.TransposedFont(base_font, scale=scale_factor)
|
||||
|
||||
def create_text_panel(self, width: int, text: str) -> Image.Image:
|
||||
font = self.get_font()
|
||||
|
||||
temp_img = Image.new('RGB', (width, self.font_size * 4), self.bg_color)
|
||||
temp_draw = ImageDraw.Draw(temp_img)
|
||||
|
||||
text_bbox = temp_draw.textbbox((0, self.font_size), text, font=font)
|
||||
text_width = text_bbox[2] - text_bbox[0]
|
||||
text_height = text_bbox[3] - text_bbox[1]
|
||||
|
||||
final_height = int(text_height * 1.5)
|
||||
panel = Image.new('RGB', (width, final_height), self.bg_color)
|
||||
draw = ImageDraw.Draw(panel)
|
||||
|
||||
x = (width - text_width) // 2
|
||||
y = (final_height - text_height) // 2
|
||||
|
||||
draw.text((x, y), text, font=font, fill=self.font_color)
|
||||
return panel
|
||||
|
||||
def process_image(self, img: Image.Image, target_size: tuple) -> Image.Image:
|
||||
target_width, target_height = target_size
|
||||
img_width, img_height = img.size
|
||||
|
||||
scale_width = target_width / img_width
|
||||
scale_height = target_height / img_height
|
||||
scale = max(scale_width, scale_height)
|
||||
|
||||
new_width = int(img_width * scale)
|
||||
new_height = int(img_height * scale)
|
||||
|
||||
resized = resize_image(img, new_width, new_height)
|
||||
left = (new_width - target_width) // 2
|
||||
top = (new_height - target_height) // 2
|
||||
right = left + target_width
|
||||
bottom = top + target_height
|
||||
|
||||
return resized.crop((left, top, right, bottom))
|
||||
|
||||
def generate(self, image1, image2, text1, text2):
|
||||
img1 = tensor2pil(image1)
|
||||
img2 = tensor2pil(image2)
|
||||
|
||||
if img2.size != img1.size:
|
||||
img2 = resize_image(img2, img1.size[0], img1.size[1])
|
||||
|
||||
panel1 = None if not text1.strip() else self.create_text_panel(img1.width, text1)
|
||||
panel2 = None if not text2.strip() else self.create_text_panel(img2.width, text2)
|
||||
|
||||
total_width = img1.width + img2.width + self.padding * 3
|
||||
img_height = img1.height
|
||||
panel_height = (panel1.height if panel1 else 0) if (panel1 or panel2) else 0
|
||||
total_height = img_height + (panel_height + self.padding if panel_height > 0 else 0) + self.padding * 2
|
||||
|
||||
result = Image.new('RGB', (total_width, total_height), self.bg_color)
|
||||
|
||||
x1 = self.padding
|
||||
x2 = x1 + img1.width + self.padding
|
||||
y = self.padding
|
||||
|
||||
result.paste(img1, (x1, y))
|
||||
result.paste(img2, (x2, y))
|
||||
|
||||
if panel1:
|
||||
result.paste(panel1, (x1, y + img_height + self.padding))
|
||||
if panel2:
|
||||
result.paste(panel2, (x2, y + img_height + self.padding))
|
||||
|
||||
return (pil2tensor(result),)
|
||||
|
||||
# Color Input node
|
||||
class AILab_ColorInput:
|
||||
@classmethod
|
||||
def INPUT_TYPES(self):
|
||||
return {
|
||||
"required": {
|
||||
"preset": (list(COLOR_PRESETS.keys()),),
|
||||
"color": ("STRING", {"default": "", "placeholder": "Enter color code (e.g. #FF0000 or #F00)"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("COLOR",)
|
||||
RETURN_NAMES = ("COLOR",)
|
||||
FUNCTION = 'get_color'
|
||||
CATEGORY = '🧪AILab/🛠️UTIL/🔄IO'
|
||||
|
||||
def get_color(self, preset, color):
|
||||
if not color:
|
||||
return (COLOR_PRESETS[preset],)
|
||||
|
||||
try:
|
||||
fixed_color = fix_color_format(color)
|
||||
if not all(c in '0123456789ABCDEFabcdef' for c in fixed_color[1:]):
|
||||
raise ValueError(f"Invalid hex characters in {color}")
|
||||
return (fixed_color,)
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Invalid color format: {color}. Please use format like #FF0000 or #F00")
|
||||
|
||||
# Node class mappings
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"AILab_LoadImage": AILab_LoadImage,
|
||||
@@ -1093,12 +1348,15 @@ NODE_CLASS_MAPPINGS = {
|
||||
"AILab_ImageStitch": AILab_ImageStitch,
|
||||
"AILab_ImageCrop": AILab_ImageCrop,
|
||||
"AILab_ICLoRAConcat": AILab_ICLoRAConcat,
|
||||
"AILab_CropObject": AILab_CropObject,
|
||||
"AILab_ImageCompare": AILab_ImageCompare,
|
||||
"AILab_ColorInput": AILab_ColorInput
|
||||
}
|
||||
|
||||
# Node display name mappings
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"AILab_LoadImage": "Load Image (RMBG) 🖼️",
|
||||
"AILab_Preview": "Preview (RMBG) 🖼️🎭",
|
||||
"AILab_Preview": "Image / Mask Preview (RMBG) 🖼️🎭",
|
||||
"AILab_ImagePreview": "Image Preview (RMBG) 🖼️",
|
||||
"AILab_MaskPreview": "Mask Preview (RMBG) 🎭",
|
||||
"AILab_ImageMaskConvert": "Image/Mask Converter (RMBG) 🖼️🎭",
|
||||
@@ -1109,4 +1367,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"AILab_ImageStitch": "Image Stitch (RMBG) 🖼️",
|
||||
"AILab_ImageCrop": "Image Crop (RMBG) 🖼️",
|
||||
"AILab_ICLoRAConcat": "IC LoRA Concat (RMBG) 🖼️🎭",
|
||||
"AILab_CropObject": "Crop To Object (RMBG) 🖼️🎭",
|
||||
"AILab_ImageCompare": "Image Compare (RMBG) 🖼️🖼️",
|
||||
"AILab_ColorInput": "Color Input (RMBG) 🎨"
|
||||
}
|
||||
+1
-1
@@ -373,5 +373,5 @@ NODE_CLASS_MAPPINGS = {
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"Segment": "Segment (RMBG)"
|
||||
"Segment": "Segmentation V1 (RMBG)"
|
||||
}
|
||||
@@ -0,0 +1,277 @@
|
||||
import os
|
||||
import sys
|
||||
import copy
|
||||
import torch
|
||||
import numpy as np
|
||||
from PIL import Image, ImageFilter
|
||||
from torch.hub import download_url_to_file
|
||||
|
||||
import folder_paths
|
||||
from segment_anything import sam_model_registry, SamPredictor
|
||||
from groundingdino.util.slconfig import SLConfig
|
||||
from groundingdino.models import build_model
|
||||
from groundingdino.util.utils import clean_state_dict
|
||||
from groundingdino.util import box_ops
|
||||
from transformers import AutoProcessor, AutoModelForZeroShotObjectDetection
|
||||
|
||||
from AILab_ImageMaskTools import pil2tensor, tensor2pil
|
||||
|
||||
# SAM model definitions (6 models)
|
||||
SAM_MODELS = {
|
||||
"sam_vit_h (2.56GB)": {
|
||||
"model_url": "https://huggingface.co/1038lab/sam/resolve/main/sam_vit_h.pth",
|
||||
"model_type": "vit_h",
|
||||
"filename": "sam_vit_h.pth"
|
||||
},
|
||||
"sam_vit_l (1.25GB)": {
|
||||
"model_url": "https://huggingface.co/1038lab/sam/resolve/main/sam_vit_l.pth",
|
||||
"model_type": "vit_l",
|
||||
"filename": "sam_vit_l.pth"
|
||||
},
|
||||
"sam_vit_b (375MB)": {
|
||||
"model_url": "https://huggingface.co/1038lab/sam/resolve/main/sam_vit_b.pth",
|
||||
"model_type": "vit_b",
|
||||
"filename": "sam_vit_b.pth"
|
||||
},
|
||||
"sam_hq_vit_h (2.57GB)": {
|
||||
"model_url": "https://huggingface.co/1038lab/sam/resolve/main/sam_hq_vit_h.pth",
|
||||
"model_type": "vit_h",
|
||||
"filename": "sam_hq_vit_h.pth"
|
||||
},
|
||||
"sam_hq_vit_l (1.25GB)": {
|
||||
"model_url": "https://huggingface.co/1038lab/sam/resolve/main/sam_hq_vit_l.pth",
|
||||
"model_type": "vit_l",
|
||||
"filename": "sam_hq_vit_l.pth"
|
||||
},
|
||||
"sam_hq_vit_b (379MB)": {
|
||||
"model_url": "https://huggingface.co/1038lab/sam/resolve/main/sam_hq_vit_b.pth",
|
||||
"model_type": "vit_b",
|
||||
"filename": "sam_hq_vit_b.pth"
|
||||
}
|
||||
}
|
||||
|
||||
# GroundingDINO model definitions (2 models)
|
||||
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",
|
||||
"config_filename": "GroundingDINO_SwinT_OGC.cfg.py",
|
||||
"model_filename": "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",
|
||||
"config_filename": "GroundingDINO_SwinB.cfg.py",
|
||||
"model_filename": "groundingdino_swinb_cogcoor.pth"
|
||||
}
|
||||
}
|
||||
|
||||
def get_or_download_model_file(filename, url, dirname):
|
||||
local_path = folder_paths.get_full_path(dirname, filename)
|
||||
if local_path:
|
||||
return local_path
|
||||
folder = os.path.join(folder_paths.models_dir, dirname)
|
||||
os.makedirs(folder, exist_ok=True)
|
||||
local_path = os.path.join(folder, filename)
|
||||
if not os.path.exists(local_path):
|
||||
print(f"Downloading {filename} from {url} ...")
|
||||
download_url_to_file(url, local_path)
|
||||
return local_path
|
||||
|
||||
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 apply_background_color(image: Image.Image, mask_image: Image.Image,
|
||||
background: str = "Alpha",
|
||||
background_color: str = "#222222") -> Image.Image:
|
||||
rgba_image = image.copy().convert('RGBA')
|
||||
rgba_image.putalpha(mask_image.convert('L'))
|
||||
if background == "Color":
|
||||
def hex_to_rgba(hex_color):
|
||||
hex_color = hex_color.lstrip('#')
|
||||
r, g, b = int(hex_color[0:2], 16), int(hex_color[2:4], 16), int(hex_color[4:6], 16)
|
||||
return (r, g, b, 255)
|
||||
rgba = hex_to_rgba(background_color)
|
||||
bg_image = Image.new('RGBA', image.size, rgba)
|
||||
composite_image = Image.alpha_composite(bg_image, rgba_image)
|
||||
return composite_image.convert('RGB')
|
||||
return rgba_image
|
||||
|
||||
def get_groundingdino_model(device):
|
||||
processor = AutoProcessor.from_pretrained("IDEA-Research/grounding-dino-tiny")
|
||||
model = AutoModelForZeroShotObjectDetection.from_pretrained("IDEA-Research/grounding-dino-tiny").to(device)
|
||||
return processor, model
|
||||
|
||||
def get_boxes(processor, model, img_pil, prompt, threshold):
|
||||
inputs = processor(images=img_pil, text=prompt, return_tensors="pt").to(model.device)
|
||||
with torch.no_grad():
|
||||
outputs = model(**inputs)
|
||||
results = processor.post_process_grounded_object_detection(
|
||||
outputs,
|
||||
inputs.input_ids,
|
||||
box_threshold=threshold,
|
||||
text_threshold=threshold,
|
||||
target_sizes=[img_pil.size[::-1]]
|
||||
)
|
||||
return results[0]["boxes"]
|
||||
|
||||
class SegmentV2:
|
||||
@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)",
|
||||
"invert_output": "Invert the mask output",
|
||||
"background": (["Alpha", "Color"], {"default": "Alpha", "tooltip": "Choose background type"}),
|
||||
"background_color": "Choose background color (Alpha = transparent)",
|
||||
}
|
||||
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.30, "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": -64, "max": 64, "step": 1, "tooltip": tooltips["mask_offset"]}),
|
||||
"invert_output": ("BOOLEAN", {"default": False, "tooltip": tooltips["invert_output"]}),
|
||||
"background": (["Alpha", "Color"], {"default": "Alpha", "tooltip": tooltips["background"]}),
|
||||
"background_color": ("COLOR", {"default": "#222222", "tooltip": tooltips["background_color"]}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK", "IMAGE")
|
||||
RETURN_NAMES = ("IMAGE", "MASK", "MASK_IMAGE")
|
||||
FUNCTION = "segment_v2"
|
||||
CATEGORY = "🧪AILab/🧽RMBG"
|
||||
|
||||
def __init__(self):
|
||||
self.dino_model_cache = {}
|
||||
self.sam_model_cache = {}
|
||||
|
||||
def segment_v2(self, image, prompt, sam_model, dino_model, threshold=0.30,
|
||||
mask_blur=0, mask_offset=0, background="Alpha",
|
||||
background_color="#222222", invert_output=False):
|
||||
img_pil = tensor2pil(image[0]) if image.ndim == 4 else tensor2pil(image)
|
||||
img_np = np.array(img_pil.convert("RGB"))
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
# Load GroundingDINO config and weights
|
||||
dino_info = DINO_MODELS[dino_model]
|
||||
config_path = get_or_download_model_file(dino_info["config_filename"], dino_info["config_url"], "grounding-dino")
|
||||
weights_path = get_or_download_model_file(dino_info["model_filename"], dino_info["model_url"], "grounding-dino")
|
||||
|
||||
# Load and cache GroundingDINO model
|
||||
dino_key = (config_path, weights_path, device)
|
||||
if dino_key not in self.dino_model_cache:
|
||||
args = SLConfig.fromfile(config_path)
|
||||
model = build_model(args)
|
||||
checkpoint = torch.load(weights_path, map_location="cpu")
|
||||
model.load_state_dict(clean_state_dict(checkpoint["model"]), strict=False)
|
||||
model.eval()
|
||||
model.to(device)
|
||||
self.dino_model_cache[dino_key] = model
|
||||
dino = self.dino_model_cache[dino_key]
|
||||
|
||||
# Preprocess image for DINO
|
||||
from groundingdino.datasets.transforms import Compose, RandomResize, ToTensor, Normalize
|
||||
transform = Compose([
|
||||
RandomResize([800], max_size=1333),
|
||||
ToTensor(),
|
||||
Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
|
||||
])
|
||||
image_tensor, _ = transform(img_pil.convert("RGB"), None)
|
||||
image_tensor = image_tensor.unsqueeze(0).to(device)
|
||||
|
||||
# Prepare text prompt
|
||||
text_prompt = prompt if prompt.endswith(".") else prompt + "."
|
||||
|
||||
# Forward pass
|
||||
with torch.no_grad():
|
||||
outputs = dino(image_tensor, captions=[text_prompt])
|
||||
logits = outputs["pred_logits"].sigmoid()[0]
|
||||
boxes = outputs["pred_boxes"][0]
|
||||
|
||||
# Filter boxes by threshold
|
||||
filt_mask = logits.max(dim=1)[0] > threshold
|
||||
boxes_filt = boxes[filt_mask]
|
||||
if boxes_filt.shape[0] == 0:
|
||||
width, height = img_pil.size
|
||||
empty_mask = torch.zeros((1, height, width), dtype=torch.float32, device="cpu")
|
||||
empty_mask_rgb = empty_mask.reshape((-1, 1, height, width)).movedim(1, -1).expand(-1, -1, -1, 3)
|
||||
result_image = apply_background_color(img_pil, Image.fromarray((empty_mask[0].numpy() * 255).astype(np.uint8)), background, background_color)
|
||||
return (pil2tensor(result_image), empty_mask, empty_mask_rgb)
|
||||
|
||||
# Convert boxes to xyxy
|
||||
H, W = img_pil.size[1], img_pil.size[0]
|
||||
boxes_xyxy = box_ops.box_cxcywh_to_xyxy(boxes_filt)
|
||||
boxes_xyxy = boxes_xyxy * torch.tensor([W, H, W, H], dtype=torch.float32, device=boxes_xyxy.device)
|
||||
boxes_xyxy = boxes_xyxy.cpu().numpy()
|
||||
|
||||
# Download/check SAM weights
|
||||
sam_info = SAM_MODELS[sam_model]
|
||||
sam_ckpt_path = get_or_download_model_file(sam_info["filename"], sam_info["model_url"], "SAM")
|
||||
|
||||
# Load SAM model (cache to avoid reloading)
|
||||
sam_key = (sam_info["model_type"], sam_ckpt_path, device)
|
||||
if sam_key not in self.sam_model_cache:
|
||||
sam = sam_model_registry[sam_info["model_type"]](checkpoint=sam_ckpt_path)
|
||||
sam.to(device)
|
||||
self.sam_model_cache[sam_key] = SamPredictor(sam)
|
||||
predictor = self.sam_model_cache[sam_key]
|
||||
|
||||
# Use SAM to get masks for each box
|
||||
predictor.set_image(img_np)
|
||||
boxes_tensor = torch.tensor(boxes_xyxy, dtype=torch.float32, device=predictor.device)
|
||||
transformed_boxes = predictor.transform.apply_boxes_torch(boxes_tensor, img_np.shape[:2])
|
||||
masks, _, _ = predictor.predict_torch(
|
||||
point_coords=None,
|
||||
point_labels=None,
|
||||
boxes=transformed_boxes,
|
||||
multimask_output=False
|
||||
)
|
||||
# Process mask following the original implementation
|
||||
print(f"Mask shape before processing: {masks.shape}")
|
||||
# Combine all masks into one
|
||||
combined_mask = torch.max(masks, dim=0)[0] # Take maximum across all masks
|
||||
mask = combined_mask.float().cpu().numpy()
|
||||
print(f"Mask shape after processing: {mask.shape}")
|
||||
# Squeeze out the extra dimension to get a 2D array
|
||||
mask = mask.squeeze(0)
|
||||
print(f"Final mask shape: {mask.shape}")
|
||||
mask = (mask * 255).astype(np.uint8)
|
||||
mask_pil = Image.fromarray(mask, mode="L")
|
||||
|
||||
mask_image = process_mask(mask_pil, invert_output, mask_blur, mask_offset)
|
||||
result_image = apply_background_color(img_pil, mask_image, background, background_color)
|
||||
if background == "Color":
|
||||
result_image = result_image.convert("RGB")
|
||||
else:
|
||||
result_image = result_image.convert("RGBA")
|
||||
mask_tensor = torch.from_numpy(np.array(mask_image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
mask_image_vis = mask_tensor.reshape((-1, 1, mask_image.height, mask_image.width)).movedim(1, -1).expand(-1, -1, -1, 3)
|
||||
return (pil2tensor(result_image), mask_tensor, mask_image_vis)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"AILab_SegmentV2": SegmentV2,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"AILab_SegmentV2": "Segmentation V2 (RMBG)",
|
||||
}
|
||||
|
||||
+1
-1
@@ -3,7 +3,7 @@ import sys
|
||||
import os
|
||||
import importlib.util
|
||||
|
||||
__version__ = "2.3.0"
|
||||
__version__ = "2.4.0"
|
||||
|
||||
# Add module directory to Python path
|
||||
current_dir = Path(__file__).parent
|
||||
|
||||
Reference in New Issue
Block a user