Add files via upload

This commit is contained in:
AI Lab
2025-07-27 13:46:20 -07:00
committed by GitHub
parent fe109d4a23
commit 75dc006dfc
9 changed files with 751 additions and 165 deletions
+2 -1
View File
@@ -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")))
+1
View File
@@ -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)
+1
View File
@@ -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)
+1
View File
@@ -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)
+1
View File
@@ -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)
+595 -71
View File
@@ -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) 🖼️",
+2 -1
View File
@@ -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")))
+8
View File
@@ -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()
+140 -92
View File
@@ -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,