diff --git a/__init__.py b/__init__.py index 3f5a9b7..0edb451 100644 --- a/__init__.py +++ b/__init__.py @@ -20,6 +20,7 @@ from utility_nodes import TRI3D_extract_facer_mask from .AEMatter import (load_AEMatter_Model, run_AEMatter_inference) from .light_layer import main_light_layer +from .remove_small_mask_islands import TRI3D_RemoveSmallMaskIslands from .image_stack import ( @@ -3696,6 +3697,7 @@ class TRI3D_BGREMOVE_MEGA(): from photoroom import TRI3D_photoroom_bgremove_api from smart_box import TRI3D_SmartBox, TRI3D_Skip_HeadMask, TRI3D_Skip_HeadMask_AddNeck, TRI3D_Image_extend, TRI3D_Smart_Depth, TRI3D_NarrowfyImage, TRI3D_Skip_LipMask from nsfw import TRI3DNSFWFilter +from cut_by_mask_aspect_ratio import TRI3D_CutByMaskAspectRatio # A dictionary that contains all nodes you want to export with their names # NOTE: names should be globally unique @@ -3765,6 +3767,8 @@ NODE_CLASS_MAPPINGS = { "tri3d_NSFWFilter": TRI3DNSFWFilter, "tri3d_NarrowfyImage": TRI3D_NarrowfyImage, "tri3d_Skip_LipMask": TRI3D_Skip_LipMask, + "tri3d_Remove_Small_Mask_Islands": TRI3D_RemoveSmallMaskIslands, + "tri3d_CutByMaskAspectRatio": TRI3D_CutByMaskAspectRatio, } @@ -3837,4 +3841,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "tri3d_Image_extend": "Image extend" + " v" + VERSION, "tri3d_Smart_Depth": "Smart Depth" + " v" + VERSION, "tri3d_NarrowfyImage": "Narrowfy Image" + " v" + VERSION, + "tri3d_Remove_Small_Mask_Islands": "Remove Small Mask Islands" + " v" + VERSION, + "tri3d_CutByMaskAspectRatio": "Cut by mask aspect ratio" + " v" + VERSION, } diff --git a/cut_by_mask_aspect_ratio.py b/cut_by_mask_aspect_ratio.py new file mode 100644 index 0000000..9690b99 --- /dev/null +++ b/cut_by_mask_aspect_ratio.py @@ -0,0 +1,162 @@ +import os +import cv2 +import numpy as np +import torch + +class TRI3D_CutByMaskAspectRatio: + """ + ComfyUI node that crops an image based on a mask's bounding box, + adjusts the aspect ratio, and resizes to specified dimensions. + """ + + def from_torch_image(self, image): + """Convert a torch tensor image to numpy array for OpenCV processing""" + image = image.cpu().numpy() * 255.0 + image = np.clip(image, 0, 255).astype(np.uint8) + return image + + def to_torch_image(self, image): + """Convert numpy array back to torch tensor format""" + image = image.astype(dtype=np.float32) + image /= 255.0 + image = torch.from_numpy(image) + return image + + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "mask": ("IMAGE",), + "margin": ("INT", {"default": 10, "min": 0, "max": 100, "step": 1}), + "target_width": ("INT", {"default": 512, "min": 64, "max": 4096, "step": 8}), + "target_height": ("INT", {"default": 512, "min": 64, "max": 4096, "step": 8}), + }, + } + + FUNCTION = "run" + RETURN_TYPES = ("IMAGE",) + CATEGORY = "TRI3D" + + def run(self, image, mask, margin, target_width, target_height): + # Convert Torch images to OpenCV format + cv_image = self.from_torch_image(image) + cv_mask = self.from_torch_image(mask) + + # Remove batch dimension if present + if len(cv_image.shape) == 4: + cv_image = cv_image[0] + if len(cv_mask.shape) == 4: + cv_mask = cv_mask[0] + + # Convert mask to grayscale if it's not already + if len(cv_mask.shape) == 3 and cv_mask.shape[2] > 1: + mask_gray = cv2.cvtColor(cv_mask, cv2.COLOR_RGB2GRAY) + else: + mask_gray = cv_mask[:, :, 0] + + # Create binary mask + _, binary_mask = cv2.threshold(mask_gray, 127, 255, cv2.THRESH_BINARY) + + # Find contours in the binary mask + contours, _ = cv2.findContours(binary_mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) + + if not contours: + # If no contours found, return the original image + print("No contours found in mask. Returning original image.") + return (image,) + + # Find bounding box around all contours + x_min, y_min = float('inf'), float('inf') + x_max, y_max = 0, 0 + + for contour in contours: + x, y, w, h = cv2.boundingRect(contour) + x_min = min(x_min, x) + y_min = min(y_min, y) + x_max = max(x_max, x + w) + y_max = max(y_max, y + h) + + # Add margin to bounding box + x_min = max(0, x_min - margin) + y_min = max(0, y_min - margin) + x_max = min(cv_image.shape[1], x_max + margin) + y_max = min(cv_image.shape[0], y_max + margin) + + # Current dimensions of the bounding box + height = y_max - y_min + width = x_max - x_min + + # Calculate the target aspect ratio (width/height) + target_aspect_ratio = target_width / target_height + + # Calculate current aspect ratio + current_aspect_ratio = width / height + + # Adjust width to match the target aspect ratio while keeping height constant + if current_aspect_ratio < target_aspect_ratio: + # Current width is too narrow, extend it + new_width = int(height * target_aspect_ratio) + width_difference = new_width - width + + # Add equal padding on both sides if possible + left_extend = width_difference // 2 + right_extend = width_difference - left_extend + + # Ensure we don't go out of bounds + if x_min - left_extend < 0: + # Not enough space on the left + left_extend = x_min + right_extend = width_difference - left_extend + + if x_max + right_extend > cv_image.shape[1]: + # Not enough space on the right + right_extend = cv_image.shape[1] - x_max + left_extend = width_difference - right_extend + + # Double-check left boundary again + if x_min - left_extend < 0: + left_extend = x_min + + # Apply the extension + x_min -= left_extend + x_max += right_extend + + elif current_aspect_ratio > target_aspect_ratio: + # Current width is too wide, crop it + new_width = int(height * target_aspect_ratio) + width_difference = width - new_width + + # Crop equally from both sides if possible + left_crop = width_difference // 2 + right_crop = width_difference - left_crop + + # Apply the crop + x_min += left_crop + x_max -= right_crop + + # Crop the image to the adjusted bounding box + cropped_image = cv_image[y_min:y_max, x_min:x_max] + + # Resize the cropped image to the target dimensions using Lanczos interpolation + resized_image = cv2.resize(cropped_image, (target_width, target_height), interpolation=cv2.INTER_LANCZOS4) + + # Convert back to torch format + torch_image = self.to_torch_image(resized_image) + + # Add batch dimension back + torch_image = torch_image.unsqueeze(0) + + return (torch_image,) + +# # Node registration for ComfyUI +# NODE_CLASS_MAPPINGS = { +# "TRI3D_CutByMaskAspectRatio": TRI3D_CutByMaskAspectRatio +# } + +# NODE_DISPLAY_NAME_MAPPINGS = { +# "TRI3D_CutByMaskAspectRatio": "TRI3D Cut By Mask Aspect Ratio" +# } diff --git a/remove_small_mask_islands.py b/remove_small_mask_islands.py new file mode 100644 index 0000000..cddeef0 --- /dev/null +++ b/remove_small_mask_islands.py @@ -0,0 +1,117 @@ +import os +import cv2 +import numpy as np +import torch + +class TRI3D_RemoveSmallMaskIslands: + """ + ComfyUI node that removes small islands of white pixels from a mask image + based on a specified area threshold. + """ + + def from_torch_image(self, image): + """Convert a torch tensor image to numpy array for OpenCV processing""" + image = image.cpu().numpy() * 255.0 + image = np.clip(image, 0, 255).astype(np.uint8) + return image + + def to_torch_image(self, image): + """Convert numpy array back to torch tensor format""" + image = image.astype(dtype=np.float32) + image /= 255.0 + image = torch.from_numpy(image) + return image + + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE", ), + "min_island_area": ("INT", {"default": 100, "min": 1, "max": 10000, "step": 10}), + "invert": ("BOOLEAN", {"default": False}), + }, + } + + FUNCTION = "run" + RETURN_TYPES = ("IMAGE", ) + CATEGORY = "TRI3D" + + def run(self, image, min_island_area, invert): + # Convert Torch image to OpenCV format + cv_image = self.from_torch_image(image) + + # Remove batch dimension if present + if len(cv_image.shape) == 4: + cv_image = cv_image[0] + + # Make a copy to work with + result_image = cv_image.copy() + + # Process each channel (if grayscale, it will just be one iteration) + height, width = cv_image.shape[:2] + + # If the image has 3 channels (RGB), convert to grayscale for contour detection + if len(cv_image.shape) == 3 and cv_image.shape[2] == 3: + # Convert to grayscale for processing + gray = cv2.cvtColor(cv_image, cv2.COLOR_RGB2GRAY) + else: + # Use the first channel if it's already grayscale or has alpha + gray = cv_image[:, :, 0] + + # Invert if needed (to work with black islands instead of white) + if invert: + gray = 255 - gray + + # Create binary image + _, binary = cv2.threshold(gray, 127, 255, cv2.THRESH_BINARY) + + # Find contours in the binary image + contours, _ = cv2.findContours(binary, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) + + # Create a blank mask for the cleaned image + clean_mask = np.zeros((height, width), dtype=np.uint8) + + # Draw only contours with area greater than the threshold + for contour in contours: + area = cv2.contourArea(contour) + if area >= min_island_area: + cv2.drawContours(clean_mask, [contour], 0, 255, -1) + + # Invert back if needed + if invert: + clean_mask = 255 - clean_mask + + # Apply the clean mask to each channel of the original image + if len(cv_image.shape) == 3 and cv_image.shape[2] == 3: + # RGB image + for i in range(3): + result_image[:, :, i] = cv2.bitwise_and(cv_image[:, :, i], clean_mask) + elif len(cv_image.shape) == 3 and cv_image.shape[2] == 4: + # RGBA image + for i in range(4): + result_image[:, :, i] = cv2.bitwise_and(cv_image[:, :, i], clean_mask) + else: + # Single channel image + result_image = cv2.bitwise_and(cv_image, clean_mask) + # Reshape to match expected dimensions + result_image = result_image.reshape(height, width, 1) + + # Convert back to torch format + torch_image = self.to_torch_image(result_image) + + # Add batch dimension back + torch_image = torch_image.unsqueeze(0) + + return (torch_image,) + +# # Node registration for ComfyUI +# NODE_CLASS_MAPPINGS = { +# "TRI3D_RemoveSmallMaskIslands": TRI3D_RemoveSmallMaskIslands +# } + +# NODE_DISPLAY_NAME_MAPPINGS = { +# "TRI3D_RemoveSmallMaskIslands": "TRI3D Remove Small Mask Islands" +# } \ No newline at end of file