Add files via upload
This commit is contained in:
+312
-6
@@ -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) 🖼️🎭"
|
||||
}
|
||||
@@ -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)",
|
||||
}
|
||||
+72
-92
@@ -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
|
||||
}
|
||||
|
||||
+1
-1
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user