diff --git a/AILab_BiRefNet.py b/AILab_BiRefNet.py index 29800f9..6efd491 100644 --- a/AILab_BiRefNet.py +++ b/AILab_BiRefNet.py @@ -445,7 +445,8 @@ class BiRefNetRMBG: else: raise ValueError("Invalid color format") return (r, g, b, a) - rgba = hex_to_rgba(params["background_color"]) + background_color = params.get("background_color", "#222222") + rgba = hex_to_rgba(background_color) bg_image = Image.new('RGBA', orig_image.size, rgba) composite_image = Image.alpha_composite(bg_image, foreground) processed_images.append(pil2tensor(composite_image.convert("RGB"))) diff --git a/AILab_BodySegment.py b/AILab_BodySegment.py index 9b2ac84..aefe6a3 100644 --- a/AILab_BodySegment.py +++ b/AILab_BodySegment.py @@ -213,6 +213,7 @@ class BodySegment: raise ValueError("Invalid color format") return (r, g, b, a) rgba_image = RGB2RGBA(orig_image, mask_image) + background_color = params.get("background_color", "#222222") rgba = hex_to_rgba(background_color) bg_image = Image.new('RGBA', orig_image.size, rgba) composite_image = Image.alpha_composite(bg_image, rgba_image) diff --git a/AILab_ClothSegment.py b/AILab_ClothSegment.py index 742431b..eaf1e49 100644 --- a/AILab_ClothSegment.py +++ b/AILab_ClothSegment.py @@ -249,6 +249,7 @@ class ClothesSegment: raise ValueError("Invalid color format") return (r, g, b, a) rgba_image = RGB2RGBA(orig_image, mask_image) + background_color = params.get("background_color", "#222222") rgba = hex_to_rgba(background_color) bg_image = Image.new('RGBA', orig_image.size, rgba) composite_image = Image.alpha_composite(bg_image, rgba_image) diff --git a/AILab_FaceSegment.py b/AILab_FaceSegment.py index 3d53a7b..aa45313 100644 --- a/AILab_FaceSegment.py +++ b/AILab_FaceSegment.py @@ -254,6 +254,7 @@ class FaceSegment: raise ValueError("Invalid color format") return (r, g, b, a) rgba_image = RGB2RGBA(orig_image, mask_image) + background_color = params.get("background_color", "#222222") rgba = hex_to_rgba(background_color) bg_image = Image.new('RGBA', orig_image.size, rgba) composite_image = Image.alpha_composite(bg_image, rgba_image) diff --git a/AILab_FashionSegment.py b/AILab_FashionSegment.py index 3a2f264..93856e7 100644 --- a/AILab_FashionSegment.py +++ b/AILab_FashionSegment.py @@ -333,6 +333,7 @@ class FashionSegmentClothing: raise ValueError("Invalid color format") return (r, g, b, a) rgba_image = RGB2RGBA(orig_image, mask_image) + background_color = params.get("background_color", "#222222") rgba = hex_to_rgba(background_color) bg_image = Image.new('RGBA', orig_image.size, rgba) composite_image = Image.alpha_composite(bg_image, rgba_image) diff --git a/AILab_ImageMaskTools.py b/AILab_ImageMaskTools.py index 88c3918..135f74f 100644 --- a/AILab_ImageMaskTools.py +++ b/AILab_ImageMaskTools.py @@ -1,4 +1,4 @@ -# ComfyUI-RMBG v2.6.0 +# ComfyUI-RMBG v2.7.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. @@ -11,17 +11,21 @@ # - Preview: A universal preview tool for both images and masks. # - ImagePreview: A specialized preview tool for images. # - MaskPreview: A specialized preview tool for masks. -# -# 2. Image and Mask Processing Nodes: +# +# 2. Load Image Nodes: +# - LoadImage: A node for loading images with some frequently used options. +# - LoadImageSimple: A node for loading images with some frequently used options. +# - LoadImageAdvanced: A node for loading images with advanced options. +# +# 3. Image and Mask Processing Nodes: # - MaskOverlay: A node for overlaying a mask on an image. -# - LoadImage: A node for loading images with some Frequently used options. # - ImageMaskConvert: Converts between image and mask formats and extracts masks from image channels. # -# 3. Mask Processing Nodes: +# 4. Mask Processing Nodes: # - MaskEnhancer: Refines masks through techniques such as blur, smoothing, expansion/contraction, and hole filling. # - MaskCombiner: Combines multiple masks using union, intersection, or difference operations. # -# 4. Image Processing Nodes: +# 5. Image Processing Nodes: # - ImageCombiner: Combines foreground and background images with various blending modes and positioning options. # - ImageStitch: Stitches multiple images together in various directions. # - ImageCrop: Crops an image to a specified size and position. @@ -586,18 +590,311 @@ class AILab_MaskCombiner: print(f"Input mask shape: {mask.shape}, Target shape: {target_shape}") raise e -# Image loader node -class AILab_LoadImage: +# Base class for image loaders +class AILab_BaseImageLoader: + @classmethod + def get_image_files(cls): + input_dir = folder_paths.get_input_directory() + os.makedirs(input_dir, exist_ok=True) + return [f for f in os.listdir(input_dir) if os.path.isfile(os.path.join(input_dir, f)) and + f.lower().endswith(('.png', '.jpg', '.jpeg', '.webp', '.gif', '.bmp', '.tiff', '.tif'))] + + def download_image(self, url): + try: + import requests + from io import BytesIO + + headers = { + 'User-Agent': 'Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/91.0.4472.124 Safari/537.36' + } + + response = requests.get(url, stream=True, timeout=10, headers=headers) + if response.status_code != 200: + raise ValueError(f"Failed to download image from URL: {url}, status code: {response.status_code}") + + return Image.open(BytesIO(response.content)) + except Exception as e: + print(f"Error downloading image from URL: {str(e)}") + raise e + + def get_image(self, image_path_or_URL="", image=""): + """Get image from path, URL or selected file""" + if not image_path_or_URL and (not image or image == ""): + return None + + if image_path_or_URL: + if image_path_or_URL.startswith(('http://', 'https://')): + return self.download_image(image_path_or_URL) + else: + if os.path.isfile(image_path_or_URL): + return Image.open(image_path_or_URL) + else: + input_dir = folder_paths.get_input_directory() + full_path = os.path.join(input_dir, image_path_or_URL) + if os.path.isfile(full_path): + return Image.open(full_path) + else: + raise ValueError(f"Image file not found: {image_path_or_URL}") + else: + image_path = folder_paths.get_annotated_filepath(image) + return Image.open(image_path) + + @classmethod + def calculate_hash(cls, image_path_or_URL="", image=""): + """Calculate hash for IS_CHANGED method""" + if not image_path_or_URL and (not image or image == ""): + return "no_input" + + if image_path_or_URL: + try: + if image_path_or_URL.startswith(('http://', 'https://')): + m = hashlib.sha256() + m.update(image_path_or_URL.encode('utf-8')) + return m.digest().hex() + else: + if os.path.isfile(image_path_or_URL): + file_path = image_path_or_URL + else: + input_dir = folder_paths.get_input_directory() + file_path = os.path.join(input_dir, image_path_or_URL) + if not os.path.isfile(file_path): + return None + + m = hashlib.sha256() + with open(file_path, 'rb') as f: + m.update(f.read()) + return m.digest().hex() + except: + return None + else: + image_path = folder_paths.get_annotated_filepath(image) + m = hashlib.sha256() + with open(image_path, 'rb') as f: + m.update(f.read()) + return m.digest().hex() + + @classmethod + def validate_inputs(cls, image_path_or_URL="", image=""): + """Validate inputs for VALIDATE_INPUTS method""" + if not image_path_or_URL and (not image or image == ""): + return True + + if image_path_or_URL: + return True + + if not folder_paths.exists_annotated_filepath(image): + return f"Invalid image file: {image}" + + return True + + def process_image_to_tensor(self, img): + """Convert PIL image to tensor with proper format""" + if img is None: + return None + + img_rgb = img.convert('RGB') + output_images = [] + + for i in ImageSequence.Iterator(img_rgb): + i = ImageOps.exif_transpose(i) + if i.mode == 'I': + i = i.point(lambda i: i * (1 / 255)) + + if i.mode != 'RGB': + i = i.convert('RGB') + + image = np.array(i).astype(np.float32) / 255.0 + if len(image.shape) == 3: + image = torch.from_numpy(image)[None,] + else: + image = torch.from_numpy(image).unsqueeze(0) # Add batch dimension + output_images.append(image) + + if len(output_images) > 1: + return torch.cat(output_images, dim=0) + else: + return output_images[0] + +# Simple image loader node (basic functionality) +class AILab_LoadImageSimple(AILab_BaseImageLoader): + @classmethod + def INPUT_TYPES(cls): + files = cls.get_image_files() + return { + "required": { + "image_path_or_URL": ("STRING", {"default": "", "placeholder": "Local path, network path or URL"}), + "image": ([""] + sorted(files) if files else [""], {"image_upload": True}), + }, + "hidden": { + "extra_pnginfo": "EXTRA_PNGINFO", + }, + } + + CATEGORY = "🧪AILab/🖼️IMAGE" + RETURN_TYPES = ("IMAGE", "INT", "INT") + RETURN_NAMES = ("IMAGE", "WIDTH", "HEIGHT") + FUNCTION = "load_image" + OUTPUT_NODE = False + + def load_image(self, image_path_or_URL="", image="", extra_pnginfo=None): + try: + img = self.get_image(image_path_or_URL, image) + + if img is None: + print("No image input provided, returning empty image") + empty_image = torch.zeros((1, 64, 64, 3), dtype=torch.float32) + return (empty_image, 64, 64) + + width, height = img.size + output_image = self.process_image_to_tensor(img) + + return (output_image, width, height) + + except Exception as e: + import traceback + traceback.print_exc() + print(f"Error loading image: {e}") + empty_image = torch.zeros((1, 64, 64, 3), dtype=torch.float32) + return (empty_image, 64, 64) + + @classmethod + def IS_CHANGED(cls, image_path_or_URL="", image="", extra_pnginfo=None): + return cls.calculate_hash(image_path_or_URL, image) + + @classmethod + def VALIDATE_INPUTS(cls, image_path_or_URL="", image="", extra_pnginfo=None): + return cls.validate_inputs(image_path_or_URL, image) + +# Standard image loader node (with resize and basic mask) +class AILab_LoadImage(AILab_BaseImageLoader): upscale_methods = ["nearest-exact", "bilinear", "area", "bicubic", "lanczos"] @classmethod def INPUT_TYPES(cls): - input_dir = folder_paths.get_input_directory() - os.makedirs(input_dir, exist_ok=True) - files = [f for f in os.listdir(input_dir) if os.path.isfile(os.path.join(input_dir, f)) and f.lower().endswith(('.png', '.jpg', '.jpeg', '.webp', '.gif', '.bmp', '.tiff', '.tif'))] + files = cls.get_image_files() return { "required": { - "image": (sorted(files) or [""], {"image_upload": True}), + "image_path_or_URL": ("STRING", {"default": "","placeholder": "Local path, network path or URL"}), + "image": ([""] + sorted(files) if files else [""], {"image_upload": True}), + "upscale_method": (cls.upscale_methods, {"default": "lanczos", "tooltip": "Method used for resizing the image"}), + "scale_by": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 8.0, "step": 0.01, "tooltip": "Scale image by this factor (ignored if size > 0)"}), + "resize_mode": (["longest_side", "shortest_side", "width", "height"], {"default": "longest_side", "tooltip": "Choose how to resize the image"}), + "size": ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 1, "tooltip": "Target size for the selected resize mode (0 = keep original size)"}), + }, + "hidden": { + "extra_pnginfo": "EXTRA_PNGINFO", + }, + } + + CATEGORY = "🧪AILab/🖼️IMAGE" + RETURN_TYPES = ("IMAGE", "MASK", "INT", "INT") + RETURN_NAMES = ("IMAGE", "MASK", "WIDTH", "HEIGHT") + FUNCTION = "load_image" + OUTPUT_NODE = False + + def load_image(self, image_path_or_URL="", image="", upscale_method="lanczos", scale_by=1.0, + resize_mode="longest_side", size=0, extra_pnginfo=None): + try: + img = self.get_image(image_path_or_URL, image) + + if img is None: + empty_image = torch.zeros((1, 64, 64, 3), dtype=torch.float32) + empty_mask = torch.zeros((1, 64, 64), dtype=torch.float32) + return (empty_image, empty_mask, 64, 64) + + orig_width, orig_height = img.size + + resampling_map = { + "nearest-exact": Image.NEAREST, + "bilinear": Image.BILINEAR, + "area": Image.BOX, + "bicubic": Image.BICUBIC, + "lanczos": Image.LANCZOS + } + resampling = resampling_map.get(upscale_method, Image.LANCZOS) + + has_alpha = 'A' in img.getbands() + if has_alpha: + original_alpha = img.getchannel('A') + + img_rgb = img.convert('RGB') + + if size > 0: + if resize_mode == "longest_side": + if orig_width >= orig_height: + new_width = size + new_height = int(orig_height * (size / orig_width)) + else: + new_height = size + new_width = int(orig_width * (size / orig_height)) + img_rgb = img_rgb.resize((new_width, new_height), resampling) + elif resize_mode == "shortest_side": + if orig_width <= orig_height: + new_width = size + new_height = int(orig_height * (size / orig_width)) + else: + new_height = size + new_width = int(orig_width * (size / orig_height)) + img_rgb = img_rgb.resize((new_width, new_height), resampling) + elif resize_mode == "width": + new_width = size + new_height = int(orig_height * (size / orig_width)) + img_rgb = img_rgb.resize((new_width, new_height), resampling) + elif resize_mode == "height": + new_height = size + new_width = int(orig_width * (size / orig_height)) + img_rgb = img_rgb.resize((new_width, new_height), resampling) + elif scale_by != 1.0: + new_width = int(orig_width * scale_by) + new_height = int(orig_height * scale_by) + img_rgb = img_rgb.resize((new_width, new_height), resampling) + + width, height = img_rgb.size + + mask = None + if has_alpha: + if size > 0 or scale_by != 1.0: + mask_img = original_alpha.resize((width, height), resampling) + else: + mask_img = original_alpha + mask = np.array(mask_img).astype(np.float32) / 255.0 + mask = torch.from_numpy(mask) + if len(mask.shape) == 2: + mask = mask.unsqueeze(0) + else: + mask = torch.ones((1, height, width), dtype=torch.float32) + + output_image = self.process_image_to_tensor(img_rgb) + + return (output_image, mask, width, height) + + except Exception as e: + import traceback + traceback.print_exc() + print(f"Error loading image: {e}") + empty_image = torch.zeros((1, 64, 64, 3), dtype=torch.float32) + empty_mask = torch.zeros((1, 64, 64), dtype=torch.float32) + return (empty_image, empty_mask, 64, 64) + + @classmethod + def IS_CHANGED(cls, image_path_or_URL="", image="", upscale_method="lanczos", scale_by=1.0, resize_mode="longest_side", size=0, extra_pnginfo=None): + return cls.calculate_hash(image_path_or_URL, image) + + @classmethod + def VALIDATE_INPUTS(cls, image_path_or_URL="", image="", upscale_method="lanczos", scale_by=1.0, resize_mode="longest_side", size=0, extra_pnginfo=None): + return cls.validate_inputs(image_path_or_URL, image) + +# Advanced image loader node (with full mask processing) +class AILab_LoadImageAdvanced(AILab_BaseImageLoader): + upscale_methods = ["nearest-exact", "bilinear", "area", "bicubic", "lanczos"] + + @classmethod + def INPUT_TYPES(cls): + files = cls.get_image_files() + return { + "required": { + "image_path_or_URL": ("STRING", {"default": "","placeholder": "Local path, network path or URL"}), + "image": ([""] + sorted(files) if files else [""], {"image_upload": True}), "mask_channel": (["alpha", "red", "green", "blue"], {"default": "alpha", "tooltip": "Select channel to extract mask from"}), "upscale_method": (cls.upscale_methods, {"default": "lanczos", "tooltip": "Method used for resizing the image"}), "scale_by": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 8.0, "step": 0.01, "tooltip": "Scale image by this factor (ignored if size > 0)"}), @@ -615,10 +912,16 @@ class AILab_LoadImage: FUNCTION = "load_image" OUTPUT_NODE = False - def load_image(self, image, mask_channel="alpha", upscale_method="lanczos", scale_by=1.0, resize_mode="longest_side", size=0, extra_pnginfo=None): + def load_image(self, image_path_or_URL="", image="", mask_channel="alpha", upscale_method="lanczos", scale_by=1.0, + resize_mode="longest_side", size=0, extra_pnginfo=None): try: - image_path = folder_paths.get_annotated_filepath(image) - img = Image.open(image_path) + img = self.get_image(image_path_or_URL, image) + + if img is None: + empty_image = torch.zeros((1, 64, 64, 3), dtype=torch.float32) + empty_mask = torch.zeros((1, 64, 64), dtype=torch.float32) + empty_mask_image = empty_mask.reshape((-1, 1, 64, 64)).movedim(1, -1).expand(-1, -1, -1, 3) + return (empty_image, empty_mask, empty_mask_image, 64, 64) orig_width, orig_height = img.size @@ -721,25 +1024,18 @@ class AILab_LoadImage: import traceback traceback.print_exc() print(f"Error loading image: {e}") - empty_image = torch.zeros(1, 3, 64, 64) - empty_mask = torch.zeros(1, 64, 64) + empty_image = torch.zeros((1, 64, 64, 3), dtype=torch.float32) + empty_mask = torch.zeros((1, 64, 64), dtype=torch.float32) empty_mask_image = empty_mask.reshape((-1, 1, 64, 64)).movedim(1, -1).expand(-1, -1, -1, 3) return (empty_image, empty_mask, empty_mask_image, 64, 64) @classmethod - def IS_CHANGED(cls, image, mask_channel="alpha", upscale_method="lanczos", scale_by=1.0, resize_mode="longest_side", size=0, extra_pnginfo=None): - image_path = folder_paths.get_annotated_filepath(image) - m = hashlib.sha256() - with open(image_path, 'rb') as f: - m.update(f.read()) - return m.digest().hex() + def IS_CHANGED(cls, image_path_or_URL="", image="", mask_channel="alpha", upscale_method="lanczos", scale_by=1.0, resize_mode="longest_side", size=0, extra_pnginfo=None): + return cls.calculate_hash(image_path_or_URL, image) @classmethod - def VALIDATE_INPUTS(cls, image, mask_channel="alpha", upscale_method="lanczos", scale_by=1.0, resize_mode="longest_side", size=0, extra_pnginfo=None): - if not folder_paths.exists_annotated_filepath(image): - return f"Invalid image file: {image}" - - return True + def VALIDATE_INPUTS(cls, image_path_or_URL="", image="", mask_channel="alpha", upscale_method="lanczos", scale_by=1.0, resize_mode="longest_side", size=0, extra_pnginfo=None): + return cls.validate_inputs(image_path_or_URL, image) # Image combiner node class AILab_ImageCombiner: @@ -976,66 +1272,290 @@ class AILab_MaskExtractor: class AILab_ImageStitch: @classmethod def INPUT_TYPES(s): - return {"required": { - "image1": ("IMAGE",), - "image2": ("IMAGE",), - "concat_direction": (['right', 'top', 'left', 'bottom'], {"default": 'right'}), - }} + tooltips = { + "image1": "First image to stitch", + "direction": "Direction to stitch the second image", + "match_image_size": "If True, resize image2 to match image1's aspect ratio", + "max_width": "Maximum width of output image (0 = no limit)", + "max_height": "Maximum height of output image (0 = no limit)", + "spacing_width": "Width of spacing between images", + "background_color": "Color for spacing between images and padding background", + "kontext_mode": "Special mode that arranges 3 images in a specific layout (image1 and image2 stacked vertically, image3 on the right)" + } + + return { + "required": { + "image1": ("IMAGE",), + "direction": (["right", "down", "left", "up", "kontext_mode"], {"default": "right", "tooltip": tooltips["direction"]}), + "match_image_size": ("BOOLEAN", {"default": True, "tooltip": tooltips["match_image_size"]}), + "max_width": ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 8, "tooltip": tooltips["max_width"]}), + "max_height": ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 8, "tooltip": tooltips["max_height"]}), + "spacing_width": ("INT", {"default": 0, "min": 0, "max": 512, "step": 1, "tooltip": tooltips["spacing_width"]}), + "background_color": ("COLOR", {"default": "#FFFFFF", "tooltip": tooltips["background_color"]}), + }, + "optional": { + "image2": ("IMAGE",), + "image3": ("IMAGE",), + }, + } RETURN_TYPES = ("IMAGE",) - FUNCTION = "stitch_images" + FUNCTION = "stitch" CATEGORY = "🧪AILab/🖼️IMAGE" - def stitch_images(self, image1, image2, concat_direction): - if image1.shape[0] != image2.shape[0]: - max_batch = max(image1.shape[0], image2.shape[0]) - image1 = image1.repeat(max_batch // image1.shape[0], 1, 1, 1) - image2 = image2.repeat(max_batch // image2.shape[0], 1, 1, 1) + 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) - if concat_direction in ['right', 'left']: - # Match heights for horizontal stitching - h1 = image1.shape[1] - h2, w2 = image2.shape[1:3] - aspect = w2 / h2 + def pad_with_color(self, image, padding, color_val): + """Pad image with specified color""" + batch, height, width, channels = image.shape + r, g, b = color_val + + pad_top, pad_bottom, pad_left, pad_right = padding + + new_height = height + pad_top + pad_bottom + new_width = width + pad_left + pad_right + + result = torch.zeros((batch, new_height, new_width, channels), device=image.device) + + if channels >= 3: + result[..., 0] = r + result[..., 1] = g + result[..., 2] = b + if channels == 4: + result[..., 3] = 1.0 + + result[:, pad_top:pad_top+height, pad_left:pad_left+width, :] = image + + return result + + def match_dimensions(self, image1, image2, direction, color_val): + h1, w1 = image1.shape[1:3] + h2, w2 = image2.shape[1:3] + + if direction in ["left", "right"]: + if h1 != h2: + target_h = max(h1, h2) + if h1 < target_h: + pad_h = target_h - h1 + pad_top, pad_bottom = pad_h // 2, pad_h - pad_h // 2 + image1 = self.pad_with_color(image1, (pad_top, pad_bottom, 0, 0), color_val) + if h2 < target_h: + pad_h = target_h - h2 + pad_top, pad_bottom = pad_h // 2, pad_h - pad_h // 2 + image2 = self.pad_with_color(image2, (pad_top, pad_bottom, 0, 0), color_val) + else: + if w1 != w2: + target_w = max(w1, w2) + if w1 < target_w: + pad_w = target_w - w1 + pad_left, pad_right = pad_w // 2, pad_w - pad_w // 2 + image1 = self.pad_with_color(image1, (0, 0, pad_left, pad_right), color_val) + if w2 < target_w: + pad_w = target_w - w2 + pad_left, pad_right = pad_w // 2, pad_w - pad_w // 2 + image2 = self.pad_with_color(image2, (0, 0, pad_left, pad_right), color_val) + + return image1, image2 + + def ensure_same_channels(self, image1, image2): + if image1.shape[-1] != image2.shape[-1]: + max_channels = max(image1.shape[-1], image2.shape[-1]) + if image1.shape[-1] < max_channels: + image1 = torch.cat([ + image1, + torch.ones(*image1.shape[:-1], max_channels - image1.shape[-1], device=image1.device), + ], dim=-1) + if image2.shape[-1] < max_channels: + image2 = torch.cat([ + image2, + torch.ones(*image2.shape[:-1], max_channels - image2.shape[-1], device=image2.device), + ], dim=-1) + return image1, image2 + + def create_spacing(self, image1, image2, spacing_width, direction, color_val): + if spacing_width <= 0: + return None - new_h = h1 - new_w = int(h1 * aspect) - - image2 = self._resize(image2, new_w, new_h) + spacing_width = spacing_width + (spacing_width % 2) + + if direction in ["left", "right"]: + spacing_shape = ( + image1.shape[0], + max(image1.shape[1], image2.shape[1]), + spacing_width, + image1.shape[-1], + ) else: - # Match widths for vertical stitching + spacing_shape = ( + image1.shape[0], + spacing_width, + max(image1.shape[2], image2.shape[2]), + image1.shape[-1], + ) + + spacing = torch.zeros(spacing_shape, device=image1.device) + + r, g, b = color_val + if spacing.shape[-1] >= 3: + spacing[..., 0] = r + spacing[..., 1] = g + spacing[..., 2] = b + if spacing.shape[-1] == 4: + spacing[..., 3] = 1.0 + + return spacing + + def stitch_kontext_mode(self, image1, image2, image3, match_image_size, spacing_width, color_val): + if image1 is None or image2 is None or image3 is None: + if image3 is None: + return self.stitch_two_images(image1, image2, "down", match_image_size, spacing_width, color_val) + elif image2 is None: + return self.stitch_two_images(image1, image3, "right", match_image_size, spacing_width, color_val) + else: + return image1 + + max_batch = max(image1.shape[0], image2.shape[0], image3.shape[0]) + if image1.shape[0] < max_batch: + image1 = torch.cat([image1, image1[-1:].repeat(max_batch - image1.shape[0], 1, 1, 1)]) + if image2.shape[0] < max_batch: + image2 = torch.cat([image2, image2[-1:].repeat(max_batch - image2.shape[0], 1, 1, 1)]) + if image3.shape[0] < max_batch: + image3 = torch.cat([image3, image3[-1:].repeat(max_batch - image3.shape[0], 1, 1, 1)]) + + if match_image_size: w1 = image1.shape[2] h2, w2 = image2.shape[1:3] - aspect = h2 / w2 + aspect_ratio = h2 / w2 + target_w = w1 + target_h = int(w1 * aspect_ratio) - new_w = w1 - new_h = int(w1 * aspect) + image2 = common_upscale( + image2.movedim(-1, 1), target_w, target_h, "lanczos", "disabled" + ).movedim(1, -1) + else: + image1, image2 = self.match_dimensions(image1, image2, "down", color_val) + + image1, image2 = self.ensure_same_channels(image1, image2) + + v_spacing = self.create_spacing(image1, image2, spacing_width, "down", color_val) + + v_images = [image1, image2] + if v_spacing is not None: + v_images.insert(1, v_spacing) + + left_column = torch.cat(v_images, dim=1) + + if match_image_size: + h_left = left_column.shape[1] + h3, w3 = image3.shape[1:3] + aspect_ratio = w3 / h3 + target_h = h_left + target_w = int(h_left * aspect_ratio) - image2 = self._resize(image2, new_w, new_h) + image3 = common_upscale( + image3.movedim(-1, 1), target_w, target_h, "lanczos", "disabled" + ).movedim(1, -1) + else: + left_column, image3 = self.match_dimensions(left_column, image3, "right", color_val) + + left_column, image3 = self.ensure_same_channels(left_column, image3) + + h_spacing = self.create_spacing(left_column, image3, spacing_width, "right", color_val) + + h_images = [left_column, image3] + if h_spacing is not None: + h_images.insert(1, h_spacing) + + result = torch.cat(h_images, dim=2) + + return result - ch1, ch2 = image1.shape[-1], image2.shape[-1] - if ch1 != ch2: - if ch1 < ch2: - image1 = torch.cat((image1, torch.ones((*image1.shape[:-1], ch2-ch1), device=image1.device)), dim=-1) + def stitch_two_images(self, image1, image2, direction, match_image_size, spacing_width, color_val): + if image2 is None: + return image1 + + if image1.shape[0] != image2.shape[0]: + max_batch = max(image1.shape[0], image2.shape[0]) + if image1.shape[0] < max_batch: + image1 = torch.cat( + [image1, image1[-1:].repeat(max_batch - image1.shape[0], 1, 1, 1)] + ) + if image2.shape[0] < max_batch: + image2 = torch.cat( + [image2, image2[-1:].repeat(max_batch - image2.shape[0], 1, 1, 1)] + ) + + if match_image_size: + h1, w1 = image1.shape[1:3] + h2, w2 = image2.shape[1:3] + aspect_ratio = w2 / h2 + + if direction in ["left", "right"]: + target_h, target_w = h1, int(h1 * aspect_ratio) else: - image2 = torch.cat((image2, torch.ones((*image2.shape[:-1], ch1-ch2), device=image2.device)), dim=-1) + target_w, target_h = w1, int(w1 / aspect_ratio) - if concat_direction == 'right': - result = torch.cat((image1, image2), dim=2) - elif concat_direction == 'bottom': - result = torch.cat((image1, image2), dim=1) - elif concat_direction == 'left': - result = torch.cat((image2, image1), dim=2) - elif concat_direction == 'top': - result = torch.cat((image2, image1), dim=1) + image2 = common_upscale( + image2.movedim(-1, 1), target_w, target_h, "lanczos", "disabled" + ).movedim(1, -1) + else: + image1, image2 = self.match_dimensions(image1, image2, direction, color_val) + + image1, image2 = self.ensure_same_channels(image1, image2) + + spacing = self.create_spacing(image1, image2, spacing_width, direction, color_val) + + images = [image2, image1] if direction in ["left", "up"] else [image1, image2] + if spacing is not None: + images.insert(1, spacing) + + concat_dim = 2 if direction in ["left", "right"] else 1 + result = torch.cat(images, dim=concat_dim) + + return result + + def stitch(self, image1, direction, match_image_size, max_width, max_height, spacing_width, background_color, image2=None, image3=None,): + + if image1 is None: + return (torch.zeros((1, 64, 64, 3)),) + color_val = self.hex_to_rgb(background_color) + + if direction == "kontext_mode": + result = self.stitch_kontext_mode(image1, image2, image3, match_image_size, spacing_width, color_val) + else: + result = self.stitch_two_images(image1, image2, direction, match_image_size, spacing_width, color_val) + + if max_width > 0 or max_height > 0: + h, w = result.shape[1:3] + need_resize = False + + if max_width > 0 and w > max_width: + scale_factor = max_width / w + target_w = max_width + target_h = int(h * scale_factor) + need_resize = True + else: + target_w, target_h = w, h + + if max_height > 0 and (target_h > max_height or (target_h == h and h > max_height)): + scale_factor = max_height / target_h + target_h = max_height + target_w = int(target_w * scale_factor) + need_resize = True + + if need_resize: + result = common_upscale( + result.movedim(-1, 1), target_w, target_h, "lanczos", "disabled" + ).movedim(1, -1) + return (result,) - def _resize(self, image, width, height): - img = image.movedim(-1, 1) - resized = common_upscale(img, width, height, "lanczos", "disabled") - return resized.movedim(1, -1) - # Image Crop node class AILab_ImageCrop: @classmethod @@ -1666,6 +2186,8 @@ class AILab_ImageMaskResize: # Node class mappings NODE_CLASS_MAPPINGS = { "AILab_LoadImage": AILab_LoadImage, + "AILab_LoadImageSimple": AILab_LoadImageSimple, + "AILab_LoadImageAdvanced": AILab_LoadImageAdvanced, "AILab_Preview": AILab_Preview, "AILab_MaskOverlay": AILab_MaskOverlay, "AILab_ImagePreview": AILab_ImagePreview, @@ -1687,6 +2209,8 @@ NODE_CLASS_MAPPINGS = { # Node display name mappings NODE_DISPLAY_NAME_MAPPINGS = { "AILab_LoadImage": "Load Image (RMBG) 🖼️", + "AILab_LoadImageSimple": "Load Image Simple (RMBG) 🖼️", + "AILab_LoadImageAdvanced": "Load Image Advanced (RMBG) 🖼️", "AILab_Preview": "Image / Mask Preview (RMBG) 🖼️🎭", "AILab_MaskOverlay": "Mask Overlay (RMBG) 🖼️🎭", "AILab_ImagePreview": "Image Preview (RMBG) 🖼️", diff --git a/AILab_RMBG.py b/AILab_RMBG.py index 88a5b4e..fd39cd4 100644 --- a/AILab_RMBG.py +++ b/AILab_RMBG.py @@ -633,7 +633,8 @@ class RMBG: else: raise ValueError("Invalid color format") return (r, g, b, a) - rgba = hex_to_rgba(params["background_color"]) + background_color = params.get("background_color", "#222222") + rgba = hex_to_rgba(background_color) bg_image = Image.new('RGBA', orig_image.size, rgba) composite_image = Image.alpha_composite(bg_image, foreground) processed_images.append(pil2tensor(composite_image.convert("RGB"))) diff --git a/AILab_Segment.py b/AILab_Segment.py index 6c6009d..2317908 100644 --- a/AILab_Segment.py +++ b/AILab_Segment.py @@ -188,6 +188,8 @@ class Segment: self.clean_state_dict = clean_state_dict self.SLConfig = SLConfig self.build_model = build_model + self._sam_model_cache = {} + self._dino_model_cache = {} def segment(self, image, prompt, sam_model, dino_model, threshold=0.35, mask_blur=0, mask_offset=0, background="Alpha", @@ -241,6 +243,8 @@ class Segment: return (pil2tensor(result_image), mask_tensor, mask_image_output) def load_sam(self, model_name): + if model_name in self._sam_model_cache: + return self._sam_model_cache[model_name] sam_checkpoint_path = self.get_local_filepath( SAM_MODELS[model_name]["model_url"], "sam") model_type = SAM_MODELS[model_name]["model_type"] @@ -252,9 +256,12 @@ class Segment: sam_device = comfy.model_management.get_torch_device() sam.to(device=sam_device) sam.eval() + self._sam_model_cache[model_name] = sam return sam def load_groundingdino(self, model_name): + if model_name in self._dino_model_cache: + return self._dino_model_cache[model_name] import sys from io import StringIO temp_stdout = StringIO() @@ -279,6 +286,7 @@ class Segment: device = comfy.model_management.get_torch_device() dino.to(device=device) dino.eval() + self._dino_model_cache[model_name] = dino return dino finally: output = temp_stdout.getvalue() diff --git a/AILab_SegmentV2.py b/AILab_SegmentV2.py index ad781da..a11d79b 100644 --- a/AILab_SegmentV2.py +++ b/AILab_SegmentV2.py @@ -97,11 +97,14 @@ def apply_background_color(image: Image.Image, mask_image: Image.Image, 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) + params = {"background_color": background_color} + background_color = params.get("background_color", "#222222") rgba = hex_to_rgba(background_color) bg_image = Image.new('RGBA', image.size, rgba) composite_image = Image.alpha_composite(bg_image, rgba_image) @@ -146,7 +149,7 @@ class SegmentV2: "dino_model": (list(DINO_MODELS.keys()),), }, "optional": { - "threshold": ("FLOAT", {"default": 0.30, "min": 0.05, "max": 0.95, "step": 0.01, "tooltip": tooltips["threshold"]}), + "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": -64, "max": 64, "step": 1, "tooltip": tooltips["mask_offset"]}), "invert_output": ("BOOLEAN", {"default": False, "tooltip": tooltips["invert_output"]}), @@ -167,105 +170,150 @@ class SegmentV2: 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") + # 处理批量图像 + batch_size = image.shape[0] if len(image.shape) == 4 else 1 + if len(image.shape) == 3: + image = image.unsqueeze(0) + + result_images = [] + result_masks = [] + result_mask_images = [] + + for b in range(batch_size): + img_pil = tensor2pil(image[b]) + img_np = np.array(img_pil.convert("RGB")) - # 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] + # 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") - # 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) + # 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] - # Prepare text prompt - text_prompt = prompt if prompt.endswith(".") else prompt + "." + # 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") - # Forward pass - with torch.no_grad(): - outputs = dino(image_tensor, captions=[text_prompt]) - logits = outputs["pred_logits"].sigmoid()[0] - boxes = outputs["pred_boxes"][0] + # 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: + try: + sam = sam_model_registry[sam_info["model_type"]](checkpoint=sam_ckpt_path) + sam.to(device) + self.sam_model_cache[sam_key] = SamPredictor(sam) + except RuntimeError as e: + if "Unexpected key(s) in state_dict" in str(e): + print("Warning: SAM model loading issue detected, please try using SegmentV1 node instead") + print(f"Error details: {str(e)}") + 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) + result_images.append(pil2tensor(result_image)) + result_masks.append(empty_mask) + result_mask_images.append(empty_mask_rgb) + continue + else: + raise e + predictor = self.sam_model_cache[sam_key] - # 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") + # 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] + + # Handle case with no detected boxes + 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) + result_images.append(pil2tensor(result_image)) + result_masks.append(empty_mask) + result_mask_images.append(empty_mask_rgb) + continue + + # 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() + + # 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 + ) + + # Combine all masks into one + combined_mask = torch.max(masks, dim=0)[0] # Take maximum across all masks + mask = combined_mask.float().cpu().numpy() + mask = mask.squeeze(0) + mask = (mask * 255).astype(np.uint8) + mask_pil = Image.fromarray(mask, mode="L") + + # Process mask and apply background + 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") + + # Convert to tensors + 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) + + result_images.append(pil2tensor(result_image)) + result_masks.append(mask_tensor) + result_mask_images.append(mask_image_vis) + + # 如果没有成功处理任何图像,返回空结果 + if len(result_images) == 0: + width, height = tensor2pil(image[0]).size + empty_mask = torch.zeros((batch_size, 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) + return (image, empty_mask, empty_mask_rgb) + + # 合并所有批次的结果 + return (torch.cat(result_images, dim=0), + torch.cat(result_masks, dim=0), + torch.cat(result_mask_images, dim=0)) NODE_CLASS_MAPPINGS = { "SegmentV2": SegmentV2,