From 1af83f6c7ad65c5e96e1ca8287d18fd5046263fd Mon Sep 17 00:00:00 2001 From: AI Lab <129358391+1038lab@users.noreply.github.com> Date: Tue, 1 Jul 2025 23:11:06 -0700 Subject: [PATCH] Add files via upload --- AILab_ImageMaskTools.py | 318 +++++++++++++++++++++++++++++++++++++++- AILab_LamaRemover.py | 191 ++++++++++++++++++++++++ AILab_RMBG.py | 164 +++++++++------------ __init__.py | 2 +- 4 files changed, 576 insertions(+), 99 deletions(-) create mode 100644 AILab_LamaRemover.py diff --git a/AILab_ImageMaskTools.py b/AILab_ImageMaskTools.py index 13cab5b..a4cd54f 100644 --- a/AILab_ImageMaskTools.py +++ b/AILab_ImageMaskTools.py @@ -1,4 +1,4 @@ -# ComfyUI-RMBG v2.4.0 +# ComfyUI-RMBG v2.5.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,11 +11,11 @@ # - 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: +# - MaskOverlay: A node for overlaying a mask on an image. # - LoadImage: A node for loading images with some Frequently used options. -# -# 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. @@ -28,6 +28,9 @@ # - 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. +# +# 5. Input Nodes: +# - ColorInput: A node for inputting colors in various formats. # These nodes are crafted to streamline common image and mask operations within ComfyUI workflows. @@ -42,6 +45,9 @@ from nodes import MAX_RESOLUTION from PIL import Image, ImageFilter, ImageOps, ImageSequence, ImageChops, ImageDraw, ImageFont import torchvision.transforms.functional as T from comfy.utils import common_upscale +import torch.nn.functional as F +from comfy import model_management +from comfy_extras.nodes_mask import ImageCompositeMasked from scipy import ndimage # Utility functions @@ -214,6 +220,91 @@ class AILab_Preview(AILab_PreviewBase): "result": (image if image is not None else None, mask if mask is not None else None) } +# Mask overlay node +class AILab_MaskOverlay(AILab_PreviewBase): + def __init__(self): + super().__init__() + self.prefix_append = "_preview_" + ''.join(random.choice("abcdefghijklmnopqrstupvxyz") for x in range(5)) + self.compress_level = 4 + + @classmethod + def INPUT_TYPES(s): + tooltips = { + "mask_opacity": "Control mask opacity (0.0-1.0)", + "mask_color": "Color for the mask overlay", + "image": "Input image (RGBA will be converted to RGB)", + "mask": "Input mask" + } + + return { + "required": { + "mask_opacity": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": tooltips["mask_opacity"]}), + "mask_color": ("COLOR", {"default": "#0000FF", "tooltip": tooltips["mask_color"]}), + }, + "optional": { + "image": ("IMAGE", {"tooltip": tooltips["image"]}), + "mask": ("MASK", {"tooltip": tooltips["mask"]}), + }, + "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"}, + } + RETURN_TYPES = ("IMAGE", "MASK") + RETURN_NAMES = ("IMAGE", "MASK") + FUNCTION = "execute" + CATEGORY = "🧪AILab/🖼️IMAGE" + OUTPUT_NODE = True + + def hex_to_rgb(self, hex_color): + """Convert hex color code to RGB values (0-1 range)""" + 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 ensure_rgb(self, image): + """Ensure image is RGB format, convert from RGBA if needed""" + if image.shape[-1] == 4: + rgb_image = image[..., :3] + return rgb_image + return image + + def execute(self, mask_opacity, mask_color, filename_prefix="ComfyUI", image=None, mask=None, prompt=None, extra_pnginfo=None): + """Execute image and mask composition""" + if image is not None: + image = self.ensure_rgb(image) + + preview = None + + if mask is not None and image is None: + preview = mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])).movedim(1, -1).expand(-1, -1, -1, 3) + elif mask is None and image is not None: + preview = image + elif mask is not None and image is not None: + mask_adjusted = mask * mask_opacity + mask_image = mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])).movedim(1, -1).expand(-1, -1, -1, 3).clone() + + r, g, b = self.hex_to_rgb(mask_color) + mask_image[:, :, :, 0] = r + mask_image[:, :, :, 1] = g + mask_image[:, :, :, 2] = b + + preview, = ImageCompositeMasked.composite(self, image, mask_image, 0, 0, True, mask_adjusted) + + if preview is None: + preview = empty_image(64, 64) + + if mask is None: + mask = torch.zeros((1, 64, 64)) + + # Save preview for display + result = self.save_image(preview, filename_prefix, prompt, extra_pnginfo) + + # Return both the image and mask for further processing + return { + "ui": result["ui"] if "ui" in result else {}, + "result": (preview, mask) + } + # Mask preview node class AILab_MaskPreview(AILab_PreviewBase): def __init__(self): @@ -1334,10 +1425,222 @@ class AILab_ColorInput: except Exception as e: raise RuntimeError(f"Invalid color format: {color}. Please use format like #FF0000 or #F00") +# Image Mask Resize node +class AILab_ImageMaskResize: + upscale_methods = ["nearest-exact", "bilinear", "area", "bicubic", "lanczos"] + @classmethod + def INPUT_TYPES(s): + tooltips = { + "image": "Input image to resize", + "width": "Target width in pixels (0 to keep original width)", + "height": "Target height in pixels (0 to keep original height)", + "scale_by": "Scale image by this factor (ignored if width or height > 0)", + "upscale_method": "Method used for resizing the image", + "resize_mode": "How to handle aspect ratio: stretch (ignore ratio), resize (maintain ratio by scaling), pad/pad_edge (maintain ratio with padding), crop (maintain ratio by cropping)", + "pad_color": "Color to use for padding when resize_mode is set to pad", + "crop_position": "Position to crop from when resize_mode is set to crop", + "divisible_by": "Make dimensions divisible by this value (useful for some models that require specific dimensions)", + "mask": "Optional mask to resize along with the image", + "device": "Device to perform resizing on (CPU or GPU)" + } + + return { + "required": { + "image": ("IMAGE", {"tooltip": tooltips["image"]}), + "width": ("INT", { "default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 1, "tooltip": tooltips["width"] }), + "height": ("INT", { "default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 1, "tooltip": tooltips["height"] }), + "scale_by": ("FLOAT", { "default": 1.0, "min": 0.01, "max": 8.0, "step": 0.01, "tooltip": tooltips["scale_by"] }), + "upscale_method": (s.upscale_methods, {"tooltip": tooltips["upscale_method"]}), + "resize_mode": (["stretch", "resize", "pad", "pad_edge", "crop"], { "default": "stretch", "tooltip": tooltips["resize_mode"] }), + "pad_color": ("COLOR", { "default": "#FFFFFF", "tooltip": tooltips["pad_color"] }), + "crop_position": (["center", "top", "bottom", "left", "right"], { "default": "center", "tooltip": tooltips["crop_position"] }), + "divisible_by": ("INT", { "default": 2, "min": 0, "max": 512, "step": 1, "tooltip": tooltips["divisible_by"] }), + }, + "optional" : { + "mask": ("MASK", {"tooltip": tooltips["mask"]}), + "device": (["cpu", "gpu"], {"default": "cpu", "tooltip": tooltips["device"]}), + } + } + + RETURN_TYPES = ("IMAGE", "MASK", "INT", "INT",) + RETURN_NAMES = ("IMAGE", "MASK", "WIDTH", "HEIGHT",) + FUNCTION = "resize" + CATEGORY = "🧪AILab/🖼️IMAGE" + + def resize(self, image, width, height, scale_by, upscale_method, resize_mode, pad_color, crop_position, divisible_by, device="cpu", mask=None): + B, H, W, C = image.shape + + if device == "gpu": + if upscale_method == "lanczos": + raise Exception("Lanczos is not supported on the GPU") + device = model_management.get_torch_device() + else: + device = torch.device("cpu") + + if width == 0 and height == 0: + if scale_by != 1.0: + width = int(W * scale_by) + height = int(H * scale_by) + else: + width = W + height = H + elif width == 0: + width = W + elif height == 0: + height = H + + new_width = width + new_height = height + + if resize_mode == "resize" or resize_mode.startswith("pad"): + if width != W or height != H: + if width == W and height != H: + ratio = height / H + new_width = round(W * ratio) + new_height = height + elif height == H and width != W: + ratio = width / W + new_height = round(H * ratio) + new_width = width + else: + ratio = min(width / W, height / H) + new_width = round(W * ratio) + new_height = round(H * ratio) + + if resize_mode.startswith("pad"): + pad_left = (width - new_width) // 2 + pad_right = width - new_width - pad_left + pad_top = (height - new_height) // 2 + pad_bottom = height - new_height - pad_top + + width = new_width + height = new_height + + width = max(1, width) + height = max(1, height) + + if divisible_by > 1: + width = width - (width % divisible_by) if width >= divisible_by else divisible_by + height = height - (height % divisible_by) if height >= divisible_by else divisible_by + + out_image = image.clone().to(device) + if mask is not None: + out_mask = mask.clone().to(device) + + if resize_mode == "crop": + old_width = W + old_height = H + old_aspect = old_width / old_height + new_aspect = width / height + + if old_aspect > new_aspect: + crop_w = round(old_height * new_aspect) + crop_h = old_height + else: + crop_w = old_width + crop_h = round(old_width / new_aspect) + + if crop_position == "center": + x = (old_width - crop_w) // 2 + y = (old_height - crop_h) // 2 + elif crop_position == "top": + x = (old_width - crop_w) // 2 + y = 0 + elif crop_position == "bottom": + x = (old_width - crop_w) // 2 + y = old_height - crop_h + elif crop_position == "left": + x = 0 + y = (old_height - crop_h) // 2 + elif crop_position == "right": + x = old_width - crop_w + y = (old_height - crop_h) // 2 + + out_image = out_image.narrow(-2, x, crop_w).narrow(-3, y, crop_h) + if mask is not None: + out_mask = out_mask.narrow(-1, x, crop_w).narrow(-2, y, crop_h) + + if (width != W or height != H) or (width != out_image.shape[2] or height != out_image.shape[1]): + out_image = common_upscale(out_image.movedim(-1,1), width, height, upscale_method, crop="disabled").movedim(1,-1) + + if mask is not None: + if upscale_method == "lanczos": + out_mask = common_upscale(out_mask.unsqueeze(1).repeat(1, 3, 1, 1), width, height, upscale_method, crop="disabled").movedim(1,-1)[:, :, :, 0] + else: + out_mask = common_upscale(out_mask.unsqueeze(1), width, height, upscale_method, crop="disabled").squeeze(1) + + if resize_mode.startswith("pad"): + if pad_left > 0 or pad_right > 0 or pad_top > 0 or pad_bottom > 0: + padded_width = width + pad_left + pad_right + padded_height = height + pad_top + pad_bottom + if divisible_by > 1: + width_remainder = padded_width % divisible_by + height_remainder = padded_height % divisible_by + if width_remainder > 0: + extra_width = divisible_by - width_remainder + pad_right += extra_width + if height_remainder > 0: + extra_height = divisible_by - height_remainder + pad_bottom += extra_height + + hex_color = fix_color_format(pad_color) + r, g, b = tuple(int(hex_color[i:i+2], 16) for i in (1, 3, 5)) + color = f"{r}, {g}, {b}" + + B, H, W, C = out_image.shape + padded_width = W + pad_left + pad_right + padded_height = H + pad_top + pad_bottom + + bg_color = [int(x.strip())/255.0 for x in color.split(",")] + if len(bg_color) == 1: + bg_color = bg_color * 3 + bg_color = torch.tensor(bg_color, dtype=out_image.dtype, device=out_image.device) + + padded_image = torch.zeros((B, padded_height, padded_width, C), dtype=out_image.dtype, device=out_image.device) + + for b in range(B): + if resize_mode == "pad_edge": + top_edge = out_image[b, 0, :, :] + bottom_edge = out_image[b, H-1, :, :] + left_edge = out_image[b, :, 0, :] + right_edge = out_image[b, :, W-1, :] + + padded_image[b, :pad_top, :, :] = top_edge.mean(dim=0) + padded_image[b, pad_top+H:, :, :] = bottom_edge.mean(dim=0) + padded_image[b, :, :pad_left, :] = left_edge.mean(dim=0) + padded_image[b, :, pad_left+W:, :] = right_edge.mean(dim=0) + else: + padded_image[b, :, :, :] = bg_color.unsqueeze(0).unsqueeze(0) + + padded_image[b, pad_top:pad_top+H, pad_left:pad_left+W, :] = out_image[b] + + if mask is not None: + padded_mask = F.pad( + out_mask, + (pad_left, pad_right, pad_top, pad_bottom), + mode='constant', + value=0 + ) + out_mask = padded_mask + + out_image = padded_image + + final_width = out_image.shape[2] + final_height = out_image.shape[1] + + # 创建默认掩码(如果没有提供) + if mask is None: + out_mask = torch.zeros((B, final_height, final_width), device=torch.device("cpu"), dtype=torch.float32) + else: + out_mask = out_mask.cpu() + + return (out_image.cpu(), out_mask, final_width, final_height) + # Node class mappings NODE_CLASS_MAPPINGS = { "AILab_LoadImage": AILab_LoadImage, "AILab_Preview": AILab_Preview, + "AILab_MaskOverlay": AILab_MaskOverlay, "AILab_ImagePreview": AILab_ImagePreview, "AILab_MaskPreview": AILab_MaskPreview, "AILab_ImageMaskConvert": AILab_ImageMaskConvert, @@ -1350,13 +1653,15 @@ NODE_CLASS_MAPPINGS = { "AILab_ICLoRAConcat": AILab_ICLoRAConcat, "AILab_CropObject": AILab_CropObject, "AILab_ImageCompare": AILab_ImageCompare, - "AILab_ColorInput": AILab_ColorInput + "AILab_ColorInput": AILab_ColorInput, + "AILab_ImageMaskResize": AILab_ImageMaskResize } # Node display name mappings NODE_DISPLAY_NAME_MAPPINGS = { "AILab_LoadImage": "Load Image (RMBG) 🖼️", "AILab_Preview": "Image / Mask Preview (RMBG) 🖼️🎭", + "AILab_MaskOverlay": "Mask Overlay (RMBG) 🖼️🎭", "AILab_ImagePreview": "Image Preview (RMBG) 🖼️", "AILab_MaskPreview": "Mask Preview (RMBG) 🎭", "AILab_ImageMaskConvert": "Image/Mask Converter (RMBG) 🖼️🎭", @@ -1369,5 +1674,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "AILab_ICLoRAConcat": "IC LoRA Concat (RMBG) 🖼️🎭", "AILab_CropObject": "Crop To Object (RMBG) 🖼️🎭", "AILab_ImageCompare": "Image Compare (RMBG) 🖼️🖼️", - "AILab_ColorInput": "Color Input (RMBG) 🎨" + "AILab_ColorInput": "Color Input (RMBG) 🎨", + "AILab_ImageMaskResize": "Image Mask Resize (RMBG) 🖼️🎭" } \ No newline at end of file diff --git a/AILab_LamaRemover.py b/AILab_LamaRemover.py new file mode 100644 index 0000000..24c9d1f --- /dev/null +++ b/AILab_LamaRemover.py @@ -0,0 +1,191 @@ +# ComfyUI-RMBG +# This custom node for ComfyUI provides functionality for Object removal using Big-Lama model. +# +# reference from https://github.com/advimman/lama +# +# This integration script follows GPL-3.0 License. +# When using or modifying this code, please respect both the original model licenses +# and this integration's license terms. +# +# Source: https://github.com/AILab-AI/ComfyUI-RMBG + + +import os +import torch +import numpy as np +from PIL import Image, ImageOps, ImageFilter +import folder_paths +from comfy.model_management import get_torch_device +from torchvision import transforms +from huggingface_hub import hf_hub_download +import shutil +import gc + +def tensor2pil(image): + return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8)) + +def pil2tensor(image): + return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0) + +def pil2comfy(image): + img_tensor = torch.from_numpy(np.array(image).astype(np.float32) / 255.0) + if len(img_tensor.shape) == 3: + img_tensor = img_tensor.unsqueeze(0) + return img_tensor + +def pad_image(image, is_mask=False): + w, h = image.size + if w % 8 != 0: + w = w + (8 - w % 8) + if h % 8 != 0: + h = h + (8 - h % 8) + + fill_color = 0 if is_mask else None + padded = Image.new(image.mode, (w, h), color=fill_color) + padded.paste(image, (0, 0)) + return padded + +def cropimage(image, w, h): + return image.crop((0, 0, w, h)) + +DEVICE = get_torch_device() +folder_paths.add_model_folder_path("rmbg", os.path.join(folder_paths.models_dir, "RMBG")) + +class AILab_LamaRemover: + @classmethod + def INPUT_TYPES(s): + tooltips = { + "images": "Input images to be processed", + "masks": "Masks defining areas to be removed (white=remove)", + "removal_strength": "Strength of the removal effect (higher values increase the effect area)", + "edge_smoothness": "Controls edge smoothness (higher values create smoother transitions)" + } + + return { + "required": { + "images": ("IMAGE", {"tooltip": tooltips["images"]}), + "masks": ("MASK", {"tooltip": tooltips["masks"]}), + "removal_strength": ("INT", {"default": 230, "min": 0, "max": 255, "step": 1, "display": "slider", "tooltip": tooltips["removal_strength"]}), + "edge_smoothness": ("INT", {"default": 8, "min": 0, "max": 20, "step": 1, "display": "slider", "tooltip": tooltips["edge_smoothness"]}), + }, + } + + CATEGORY = "🧪AILab/🧽RMBG" + RETURN_NAMES = ("images",) + RETURN_TYPES = ("IMAGE",) + FUNCTION = "remove_object" + + def __init__(self): + self.model = None + self.device = DEVICE + self.cache_dir = os.path.join(folder_paths.models_dir, "RMBG", "Lama") + self.model_path = os.path.join(self.cache_dir, "big-lama.pt") + self.to_pil = transforms.ToPILImage() + + def load_model(self): + if self.model is not None: + return + + if not os.path.exists(self.model_path): + self.download_model() + + try: + self.model = torch.jit.load(self.model_path, map_location=self.device) + except Exception as e: + print(f"Can't use comfy device: {str(e)}") + self.device = "cuda" if torch.cuda.is_available() else "cpu" + self.model = torch.jit.load(self.model_path, map_location=self.device) + + self.model.eval() + self.model.to(self.device) + + def download_model(self): + print("Downloading Big-Lama model...") + os.makedirs(self.cache_dir, exist_ok=True) + + try: + downloaded_path = hf_hub_download( + repo_id="1038lab/Lama", + filename="big-lama.pt", + local_dir=self.cache_dir, + local_dir_use_symlinks=False + ) + + if os.path.dirname(downloaded_path) != self.cache_dir: + shutil.move(downloaded_path, self.model_path) + + print("Big-Lama model downloaded successfully") + except Exception as e: + raise RuntimeError(f"Error downloading Big-Lama model: {str(e)}") + + def process_with_model(self, img_tensor, mask_tensor): + with torch.inference_mode(): + img_tensor = img_tensor.to(self.device) + mask_tensor = mask_tensor.to(self.device) + + result = self.model(img_tensor, mask_tensor) + result_cpu = result[0].cpu() + + del img_tensor + del mask_tensor + if torch.cuda.is_available(): + torch.cuda.empty_cache() + + return result_cpu + + def remove_object(self, images, masks, removal_strength, edge_smoothness): + try: + self.load_model() + results = [] + + for image, mask in zip(images, masks): + ori_image = tensor2pil(image) + w, h = ori_image.size + p_image = pad_image(ori_image) + + mask_np = mask.cpu().numpy() + mask_pil = Image.fromarray((mask_np * 255).astype(np.uint8)) + p_mask = pad_image(mask_pil, is_mask=True) + + if p_mask.size != p_image.size: + try: + p_mask = p_mask.resize(p_image.size, Image.LANCZOS) + except AttributeError: + p_mask = p_mask.resize(p_image.size, Image.ANTIALIAS) + + p_mask = ImageOps.invert(p_mask) + p_mask = p_mask.filter(ImageFilter.GaussianBlur(radius=edge_smoothness)) + gray = p_mask.point(lambda x: 0 if x > removal_strength else 255) + + img_tensor = torch.FloatTensor(np.array(p_image)).permute(2, 0, 1).unsqueeze(0) / 255.0 + mask_tensor = torch.FloatTensor(np.array(gray)).unsqueeze(0).unsqueeze(0) / 255.0 + + result = self.process_with_model(img_tensor, mask_tensor) + result_img = self.to_pil(result.squeeze()) + + if result_img.width > w or result_img.height > h: + result_img = cropimage(result_img, w, h) + + result_tensor = pil2comfy(result_img) + results.append(result_tensor) + + del result + gc.collect() + + return (torch.cat(results, dim=0),) + + except Exception as e: + import traceback + print(traceback.format_exc()) + raise RuntimeError(f"Error in object removal: {str(e)}") + finally: + if torch.cuda.is_available(): + torch.cuda.empty_cache() + +NODE_CLASS_MAPPINGS = { + "AILab_LamaRemover": AILab_LamaRemover, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "AILab_LamaRemover": "Lama Remover (RMBG)", +} \ No newline at end of file diff --git a/AILab_RMBG.py b/AILab_RMBG.py index 0858526..88a5b4e 100644 --- a/AILab_RMBG.py +++ b/AILab_RMBG.py @@ -33,7 +33,6 @@ import types device = "cuda" if torch.cuda.is_available() else "cpu" -# Add model path folder_paths.add_model_folder_path("rmbg", os.path.join(folder_paths.models_dir, "RMBG")) # Model configuration @@ -156,88 +155,86 @@ class RMBGModel(BaseModelLoader): def load_model(self, model_name): if self.current_model_version != model_name: self.clear_model() - + cache_dir = self.get_cache_dir(model_name) try: - # Try standard loading first + # Primary path: Modern transformers compatibility mode (optimized for newer versions) try: - self.model = AutoModelForImageSegmentation.from_pretrained( - cache_dir, - trust_remote_code=True, - local_files_only=True + from transformers import PreTrainedModel + import json + + config_path = os.path.join(cache_dir, "config.json") + with open(config_path, 'r') as f: + config = json.load(f) + + birefnet_path = os.path.join(cache_dir, "birefnet.py") + BiRefNetConfig_path = os.path.join(cache_dir, "BiRefNet_config.py") + + # Load the BiRefNetConfig + config_spec = importlib.util.spec_from_file_location("BiRefNetConfig", BiRefNetConfig_path) + config_module = importlib.util.module_from_spec(config_spec) + sys.modules["BiRefNetConfig"] = config_module + config_spec.loader.exec_module(config_module) + + # Fix and load birefnet module + with open(birefnet_path, 'r') as f: + birefnet_content = f.read() + + birefnet_content = birefnet_content.replace( + "from .BiRefNet_config import BiRefNetConfig", + "from BiRefNetConfig import BiRefNetConfig" ) - except AttributeError as ae: - if "'Config' object has no attribute 'get_text_config'" in str(ae): - print("[RMBG WARNING] Detected newer transformers version, using compatibility mode...") - try: - from transformers import PreTrainedModel - import json - - config_path = os.path.join(cache_dir, "config.json") - with open(config_path, 'r') as f: - config = json.load(f) - - birefnet_path = os.path.join(cache_dir, "birefnet.py") - BiRefNetConfig_path = os.path.join(cache_dir, "BiRefNet_config.py") - - # Load the BiRefNetConfig - config_spec = importlib.util.spec_from_file_location("BiRefNetConfig", BiRefNetConfig_path) - config_module = importlib.util.module_from_spec(config_spec) - sys.modules["BiRefNetConfig"] = config_module - config_spec.loader.exec_module(config_module) - - # Fix and load birefnet module - with open(birefnet_path, 'r') as f: - birefnet_content = f.read() - - birefnet_content = birefnet_content.replace( - "from .BiRefNet_config import BiRefNetConfig", - "from BiRefNetConfig import BiRefNetConfig" - ) - - module_name = f"custom_birefnet_model_{hash(birefnet_path)}" - module = types.ModuleType(module_name) - sys.modules[module_name] = module - exec(birefnet_content, module.__dict__) - - for attr_name in dir(module): - attr = getattr(module, attr_name) - if isinstance(attr, type) and issubclass(attr, PreTrainedModel) and attr != PreTrainedModel: - BiRefNetConfig = getattr(config_module, "BiRefNetConfig") - model_config = BiRefNetConfig() - self.model = attr(model_config) - - weights_path = os.path.join(cache_dir, "model.safetensors") - try: - try: - import safetensors.torch - self.model.load_state_dict(safetensors.torch.load_file(weights_path)) - except ImportError: - from transformers.modeling_utils import load_state_dict - state_dict = load_state_dict(weights_path) - self.model.load_state_dict(state_dict) - except Exception as load_error: - pytorch_weights = os.path.join(cache_dir, "pytorch_model.bin") - if os.path.exists(pytorch_weights): - self.model.load_state_dict(torch.load(pytorch_weights, map_location="cpu")) - else: - raise RuntimeError(f"Failed to load weights: {str(load_error)}") - break - - if self.model is None: - raise RuntimeError("Could not find suitable model class") - - except Exception as custom_e: - handle_model_error(f"Failed to load model in compatibility mode: {str(custom_e)}") - else: - raise ae + + module_name = f"custom_birefnet_model_{hash(birefnet_path)}" + module = types.ModuleType(module_name) + sys.modules[module_name] = module + exec(birefnet_content, module.__dict__) + + for attr_name in dir(module): + attr = getattr(module, attr_name) + if isinstance(attr, type) and issubclass(attr, PreTrainedModel) and attr != PreTrainedModel: + BiRefNetConfig = getattr(config_module, "BiRefNetConfig") + model_config = BiRefNetConfig() + self.model = attr(model_config) + + weights_path = os.path.join(cache_dir, "model.safetensors") + try: + try: + import safetensors.torch + self.model.load_state_dict(safetensors.torch.load_file(weights_path)) + except ImportError: + from transformers.modeling_utils import load_state_dict + state_dict = load_state_dict(weights_path) + self.model.load_state_dict(state_dict) + except Exception as load_error: + pytorch_weights = os.path.join(cache_dir, "pytorch_model.bin") + if os.path.exists(pytorch_weights): + self.model.load_state_dict(torch.load(pytorch_weights, map_location="cpu")) + else: + raise RuntimeError(f"Failed to load weights: {str(load_error)}") + break + + if self.model is None: + raise RuntimeError("Could not find suitable model class") + + except Exception as modern_e: + print(f"[RMBG INFO] Using standard transformers loading (fallback mode)...") + try: + self.model = AutoModelForImageSegmentation.from_pretrained( + cache_dir, + trust_remote_code=True, + local_files_only=True + ) + except Exception as standard_e: + handle_model_error(f"Failed to load model with both modern and standard methods. Modern error: {str(modern_e)}. Standard error: {str(standard_e)}") + except Exception as e: handle_model_error(f"Error loading model: {str(e)}") - + self.model.eval() for param in self.model.parameters(): param.requires_grad = False - + torch.set_float32_matmul_precision('high') self.model.to(device) self.current_model_version = model_name @@ -253,17 +250,14 @@ class RMBGModel(BaseModelLoader): transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) - # Ensure input is in list format if isinstance(images, torch.Tensor): if len(images.shape) == 3: images = [images] else: images = [img for img in images] - # Store original image sizes original_sizes = [tensor2pil(img).size for img in images] - # Batch process transformations input_tensors = [transform_image(tensor2pil(img)).unsqueeze(0) for img in images] input_batch = torch.cat(input_tensors, dim=0).to(device) @@ -290,13 +284,11 @@ class RMBGModel(BaseModelLoader): masks = [] - # Process each result and resize back to original dimensions for i, (result, (orig_w, orig_h)) in enumerate(zip(results, original_sizes)): result = result.squeeze() result = result * (1 + (1 - params["sensitivity"])) result = torch.clamp(result, 0, 1) - - # Resize back to original dimensions + result = F.interpolate(result.unsqueeze(0).unsqueeze(0), size=(orig_h, orig_w), mode='bilinear').squeeze() @@ -337,13 +329,11 @@ class InspyrenetModel(BaseModelLoader): orig_image = tensor2pil(image) w, h = orig_image.size - # Resize for processing aspect_ratio = h / w new_w = params["process_res"] new_h = int(params["process_res"] * aspect_ratio) resized_image = orig_image.resize((new_w, new_h), Image.LANCZOS) - # Process image foreground = self.model.process(resized_image, type='rgba') foreground = foreground.resize((w, h), Image.LANCZOS) mask = foreground.split()[-1] @@ -580,7 +570,6 @@ class RMBG: model_instance = self.models[model] - # Check and download model if needed cache_status, message = model_instance.check_model_cache(model) if not cache_status: print(f"Cache check: {message}") @@ -591,17 +580,14 @@ class RMBG: print("Model files downloaded successfully") for img in image: - # Get mask from specific model mask = model_instance.process_image(img, model, params) - # Ensure mask is in the correct format if isinstance(mask, list): masks = [m.convert("L") for m in mask if isinstance(m, Image.Image)] mask = masks[0] if masks else None elif isinstance(mask, Image.Image): mask = mask.convert("L") - # Post-process mask mask_tensor = pil2tensor(mask) mask_tensor = mask_tensor * (1 + (1 - params["sensitivity"])) mask_tensor = torch.clamp(mask_tensor, 0, 1) @@ -621,11 +607,9 @@ class RMBG: if params["invert_output"]: mask = Image.fromarray(255 - np.array(mask)) - # Convert to tensors for refine_foreground img_tensor = torch.from_numpy(np.array(tensor2pil(img))).permute(2, 0, 1).unsqueeze(0) / 255.0 mask_tensor = torch.from_numpy(np.array(mask)).unsqueeze(0).unsqueeze(0) / 255.0 - # Create final image orig_image = tensor2pil(img) if params.get("refine_foreground", False): @@ -658,10 +642,8 @@ class RMBG: processed_masks.append(pil2tensor(mask)) - # Create mask image for visualization mask_images = [] for mask_tensor in processed_masks: - # Convert mask to RGB image format for visualization mask_image = mask_tensor.reshape((-1, 1, mask_tensor.shape[-2], mask_tensor.shape[-1])).movedim(1, -1).expand(-1, -1, -1, 3) mask_images.append(mask_image) @@ -671,12 +653,10 @@ class RMBG: except Exception as e: handle_model_error(f"Error in image processing: {str(e)}") - # Return original image and empty mask on error empty_mask = torch.zeros((image.shape[0], image.shape[2], image.shape[3])) empty_mask_image = empty_mask.reshape((-1, 1, empty_mask.shape[-2], empty_mask.shape[-1])).movedim(1, -1).expand(-1, -1, -1, 3) return (image, empty_mask, empty_mask_image) -# Node Mapping NODE_CLASS_MAPPINGS = { "RMBG": RMBG } diff --git a/__init__.py b/__init__.py index 594b6aa..b869f1f 100644 --- a/__init__.py +++ b/__init__.py @@ -3,7 +3,7 @@ import sys import os import importlib.util -__version__ = "2.4.0" +__version__ = "2.5.0" # Add module directory to Python path current_dir = Path(__file__).parent