diff --git a/AILab_ImageMaskTools.py b/AILab_ImageMaskTools.py index 3e36188..13cab5b 100644 --- a/AILab_ImageMaskTools.py +++ b/AILab_ImageMaskTools.py @@ -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) 🎨" } \ No newline at end of file diff --git a/AILab_Segment.py b/AILab_Segment.py index c10b543..6c6009d 100644 --- a/AILab_Segment.py +++ b/AILab_Segment.py @@ -373,5 +373,5 @@ NODE_CLASS_MAPPINGS = { } NODE_DISPLAY_NAME_MAPPINGS = { - "Segment": "Segment (RMBG)" + "Segment": "Segmentation V1 (RMBG)" } \ No newline at end of file diff --git a/AILab_SegmentV2.py b/AILab_SegmentV2.py new file mode 100644 index 0000000..0aa3a24 --- /dev/null +++ b/AILab_SegmentV2.py @@ -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)", +} + diff --git a/__init__.py b/__init__.py index 2e056ad..594b6aa 100644 --- a/__init__.py +++ b/__init__.py @@ -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