diff --git a/grabcut_nodes.py b/grabcut_nodes.py index a95e662..e03bbef 100644 --- a/grabcut_nodes.py +++ b/grabcut_nodes.py @@ -1,9 +1,15 @@ -import numpy as np -import torch -from typing import Tuple, Optional +from __future__ import annotations + import time +from typing import Optional, Tuple + +import numpy as np +import structlog +import torch from PIL import Image +log = structlog.get_logger(__name__) + # Try to import ComfyUI modules try: import folder_paths @@ -11,13 +17,26 @@ try: COMFY_AVAILABLE = True except ImportError: COMFY_AVAILABLE = False - print("ComfyUI modules not available. Running in standalone mode.") + log.info("comfyui_modules_unavailable", mode="standalone") # Import our GrabCut processor try: - from .grabcut_remover import GrabCutProcessor, create_fallback_processor + from .grabcut_remover import GrabCutProcessor, create_fallback_processor, _log_gpu_memory except ImportError: - from grabcut_remover import GrabCutProcessor, create_fallback_processor + from grabcut_remover import GrabCutProcessor, create_fallback_processor, _log_gpu_memory + +# Pydantic validation — security-hardened parameter sanitisation before GPU execution +try: + from src.validation import validate_node_params, GrabCutParams, MaskParams +except ImportError: + import os.path as _path + import sys as _sys + _sys.path.insert(0, _path.dirname(__file__)) + try: + from src.validation import validate_node_params, GrabCutParams, MaskParams + except ImportError: + # Graceful degradation — log and continue without validation + validate_node_params = GrabCutParams = MaskParams = None class ScalingMixin: @@ -303,8 +322,8 @@ class AutoGrabCutRemover(ScalingMixin): try: self.processor = GrabCutProcessor() except Exception as e: - print(f"Warning: Could not initialize YOLO-based processor: {e}") - print("Using fallback processor without YOLO") + log.warning("grabcut_node.yolo_init_failed", error=str(e), + fallback="FallbackGrabCutProcessor") self.processor = create_fallback_processor()() def _map_object_class(self, object_class: str) -> Optional[str]: @@ -320,7 +339,8 @@ class AutoGrabCutRemover(ScalingMixin): } return mapping.get(object_class, None) - def remove_background(self, image: torch.Tensor, + @torch.no_grad() + def remove_background(self, image: torch.Tensor, initial_mask: Optional[torch.Tensor] = None, object_class: str = "auto", confidence_threshold: float = 0.5, @@ -356,10 +376,42 @@ class AutoGrabCutRemover(ScalingMixin): Returns: Tuple of (processed_image, mask, bbox_string, confidence, metrics) """ + # --- Pydantic validation: sanitise ALL user params before GPU execution --- + if validate_node_params is not None: + try: + validate_node_params( + grabcut_iterations=grabcut_iterations, + margin=margin_pixels, + edge_threshold=0.5, + confidence_threshold=confidence_threshold, + target_long_edge=4096, + maintain_aspect=True, + scaling_method="auto", + edge_blur_amount=int(edge_blur_amount), + invert_mask=False, + edge_refinement_strength=edge_refinement, + bbox_safety_margin=bbox_safety_margin, + min_bbox_size=min_bbox_size, + fallback_margin_percent=fallback_margin_percent, + binary_threshold=binary_threshold, + output_format=output_format, + auto_adjust=auto_adjust, + ) + except Exception as exc: + log.error("grabcut_node.validation_failed", node="AutoGrabCutRemover", error=str(exc)) + raise ValueError(f"[AutoGrabCutRemover] Invalid parameters: {exc}") from exc + + log.info("grabcut_node.remove_background.start", + batch_size=image.shape[0] if len(image.shape) == 4 else 1, + output_format=output_format) + if torch.cuda.is_available(): + log.debug("gpu_memory.remove_background.start", + allocated_gb=round(torch.cuda.memory_allocated() / 1e9, 3)) + # Ensure processor is initialized if self.processor is None: self._initialize_processor() - + # Update processor parameters self.processor.confidence_threshold = confidence_threshold self.processor.iterations = grabcut_iterations @@ -449,15 +501,16 @@ class AutoGrabCutRemover(ScalingMixin): all_metrics = [] for i in range(batch_size): + _log_gpu_memory(f"batch_item_{i}.start") try: # Convert from ComfyUI tensor format to numpy img_tensor = image[i] - + # Validate input tensor if len(img_tensor.shape) != 3: - print(f"Warning: Unexpected tensor shape {img_tensor.shape}, expected 3D tensor") + log.warning("grabcut_node.unexpected_shape", shape=list(img_tensor.shape)) continue - + img_np = (img_tensor.cpu().numpy() * 255).astype(np.uint8) # Ensure channels last format (H, W, C) @@ -466,7 +519,7 @@ class AutoGrabCutRemover(ScalingMixin): # Final validation if len(img_np.shape) != 3 or img_np.shape[2] not in [3, 4]: - print(f"Warning: Invalid image shape {img_np.shape} after conversion") + log.warning("grabcut_node.invalid_shape", shape=list(img_np.shape)) continue # Convert to RGB if needed @@ -500,17 +553,17 @@ class AutoGrabCutRemover(ScalingMixin): if output_format == "RGBA": # For RGBA output: preserve transparency, don't premultiply alpha # Create 4-channel RGBA tensor - rgba_tensor = torch.from_numpy(rgba.astype(np.float32) / 255.0).float() + rgba_tensor = torch.from_numpy(rgba.astype(np.float32) / 255.0).to(dtype=torch.float32) processed_images.append(rgba_tensor) else: # output_format == "MASK" # For MASK output: return binary mask as primary output # Convert alpha to binary mask (0 or 255) binary_mask = (alpha > 0.5).astype(np.float32) - mask_tensor = torch.from_numpy(binary_mask).float() + mask_tensor = torch.from_numpy(binary_mask).to(dtype=torch.float32) processed_images.append(mask_tensor.unsqueeze(-1)) # Add channel dimension # Alpha tensor for mask output (always provided) - alpha_tensor = torch.from_numpy(alpha).float() + alpha_tensor = torch.from_numpy(alpha).to(dtype=torch.float32) masks.append(alpha_tensor) # Format bbox and metrics @@ -536,11 +589,11 @@ class AutoGrabCutRemover(ScalingMixin): rgba_fallback = np.zeros((h, w, 4), dtype=np.float32) rgba_fallback[:, :, :3] = img_np.astype(np.float32) / 255.0 rgba_fallback[:, :, 3] = 1.0 # Full opacity - rgb_tensor = torch.from_numpy(rgba_fallback).float() + rgb_tensor = torch.from_numpy(rgba_fallback).to(dtype=torch.float32) else: # output_format == "MASK" # Return full foreground mask alpha_fallback = np.ones((img_np.shape[0], img_np.shape[1], 1), dtype=np.float32) - rgb_tensor = torch.from_numpy(alpha_fallback).float() + rgb_tensor = torch.from_numpy(alpha_fallback).to(dtype=torch.float32) alpha_tensor = torch.ones((img_np.shape[0], img_np.shape[1]), dtype=torch.float32) @@ -551,37 +604,39 @@ class AutoGrabCutRemover(ScalingMixin): all_metrics.append(f"Batch {i+1}/{batch_size}: Processing failed") except Exception as e: - print(f"Error processing batch item {i+1}: {e}") - # Add fallback empty tensors to maintain batch consistency - h, w = 512, 512 # Default dimensions + log.error("grabcut_node.batch_error", item=i, error=str(e)) + h, w = 512, 512 if len(processed_images) > 0: - # Use dimensions from previous successful processing h, w = processed_images[0].shape[:2] - + if output_format == "RGBA": - # Empty RGBA tensor rgb_tensor = torch.zeros((h, w, 4), dtype=torch.float32) - else: # output_format == "MASK" - # Empty mask tensor (single channel) + else: rgb_tensor = torch.zeros((h, w, 1), dtype=torch.float32) - + alpha_tensor = torch.zeros((h, w), dtype=torch.float32) - + processed_images.append(rgb_tensor) masks.append(alpha_tensor) all_bboxes.append("(0,0,0,0)") all_confidences.append(0.0) all_metrics.append(f"Batch {i+1}/{batch_size}: Error occurred") - - # Stack results + finally: + if torch.cuda.is_available(): + torch.cuda.empty_cache() + _log_gpu_memory(f"batch_item_{i}.end") + output_image = torch.stack(processed_images) output_mask = torch.stack(masks) - - # Format outputs + bbox_output = " | ".join(all_bboxes) confidence_output = float(np.mean(all_confidences)) metrics_output = "\n".join(all_metrics) - + + _log_gpu_memory("remove_background.end") + log.info("grabcut_node.remove_background.done", + batch_size=batch_size, mean_confidence=round(confidence_output, 3)) + return (output_image, output_mask, bbox_output, confidence_output, metrics_output) @@ -685,6 +740,7 @@ class GrabCutRefinement(ScalingMixin): except Exception: self.processor = create_fallback_processor()(iterations=3) + @torch.no_grad() def refine_mask(self, image: torch.Tensor, mask: torch.Tensor, grabcut_iterations: int = 3, edge_refinement: float = 0.5, @@ -709,9 +765,32 @@ class GrabCutRefinement(ScalingMixin): Returns: Tuple of (image_with_refined_alpha, refined_mask) """ + # --- Pydantic validation: sanitise params before GPU execution --- + if GrabCutParams is not None and MaskParams is not None: + try: + GrabCutParams( + iterations=grabcut_iterations, + margin=expand_margin, + edge_threshold=0.5, + confidence_threshold=0.5, + ) + MaskParams( + edge_blur_amount=int(edge_blur_amount), + invert_mask=False, + edge_refinement_strength=edge_refinement, + ) + except Exception as exc: + log.error("grabcut_node.validation_failed", + node="GrabCutRefinement", error=str(exc)) + raise ValueError(f"[GrabCutRefinement] Invalid parameters: {exc}") from exc + + log.info("grabcut_node.refine_mask.start", + batch_size=image.shape[0] if len(image.shape) == 4 else 1) + _log_gpu_memory("refine_mask.start") + if self.processor is None: self._initialize_processor() - + # Update parameters self.processor.iterations = grabcut_iterations self.processor.edge_refinement_strength = edge_refinement @@ -719,9 +798,8 @@ class GrabCutRefinement(ScalingMixin): self.processor.margin_pixels = expand_margin self.processor.bbox_safety_margin = bbox_safety_margin self.processor.min_bbox_size = min_bbox_size - # Use default fallback margin for refinement self.processor.fallback_margin_percent = 0.15 - + # Handle batch if len(image.shape) == 4: batch_size = image.shape[0] @@ -729,49 +807,47 @@ class GrabCutRefinement(ScalingMixin): batch_size = 1 image = image.unsqueeze(0) mask = mask.unsqueeze(0) - + refined_images = [] refined_masks = [] - + for i in range(batch_size): - # Convert to numpy - img_np = (image[i].cpu().numpy() * 255).astype(np.uint8) - if img_np.shape[0] == 3: - img_np = np.transpose(img_np, (1, 2, 0)) - - mask_np = (mask[i].cpu().numpy() * 255).astype(np.uint8) - if len(mask_np.shape) == 3: - mask_np = mask_np.squeeze() - - # Refine with GrabCut - result = self.processor.process_with_initial_mask(img_np, mask_np, None) - - if result['success']: - rgba = result['rgba_image'] - - # Apply resize if requested - rgba = self._apply_resize(rgba, output_size, scaling_method, custom_width, custom_height) - - rgb = rgba[:, :, :3].astype(np.float32) / 255.0 - alpha = rgba[:, :, 3].astype(np.float32) / 255.0 - - # Apply refined alpha - for c in range(3): - rgb[:, :, c] *= alpha - - rgb_tensor = torch.from_numpy(rgb).float() - alpha_tensor = torch.from_numpy(alpha).float() - - refined_images.append(rgb_tensor) - refined_masks.append(alpha_tensor) - else: - # Return original if refinement fails + _log_gpu_memory(f"refine_batch_{i}.start") + try: + img_np = (image[i].cpu().numpy() * 255).astype(np.uint8) + if img_np.shape[0] == 3: + img_np = np.transpose(img_np, (1, 2, 0)) + mask_np = (mask[i].cpu().numpy() * 255).astype(np.uint8) + if len(mask_np.shape) == 3: + mask_np = mask_np.squeeze() + result = self.processor.process_with_initial_mask(img_np, mask_np, None) + if result['success']: + rgba = result['rgba_image'] + rgba = self._apply_resize(rgba, output_size, scaling_method, custom_width, custom_height) + rgb = rgba[:, :, :3].astype(np.float32) / 255.0 + alpha = rgba[:, :, 3].astype(np.float32) / 255.0 + for c in range(3): + rgb[:, :, c] *= alpha + rgb_tensor = torch.from_numpy(rgb).to(dtype=torch.float32) + alpha_tensor = torch.from_numpy(alpha).to(dtype=torch.float32) + refined_images.append(rgb_tensor) + refined_masks.append(alpha_tensor) + else: + refined_images.append(image[i]) + refined_masks.append(mask[i]) + except Exception as e: + log.error("grabcut_node.refine_error", item=i, error=str(e)) refined_images.append(image[i]) refined_masks.append(mask[i]) - + finally: + if torch.cuda.is_available(): + torch.cuda.empty_cache() + _log_gpu_memory(f"refine_batch_{i}.end") + output_image = torch.stack(refined_images) output_mask = torch.stack(refined_masks) - + _log_gpu_memory("refine_mask.end") + log.info("grabcut_node.refine_mask.done", batch_size=batch_size) return (output_image, output_mask) diff --git a/grabcut_remover.py b/grabcut_remover.py index 5e7f4f3..8973773 100644 --- a/grabcut_remover.py +++ b/grabcut_remover.py @@ -1,108 +1,105 @@ -import numpy as np -import cv2 +from __future__ import annotations + +import os import time -from typing import Tuple, Dict, Optional, List +from typing import Any, Dict, List, Optional, Tuple + +import cv2 +import numpy as np +import structlog import torch from ultralytics import YOLO -import os +log = structlog.get_logger(__name__) +# --------------------------------------------------------------------------- # Parameter Adjustment Thresholds and Constants -# These constants define the thresholds used in auto_adjust_parameters() -# for intelligent parameter tuning based on image characteristics +# --------------------------------------------------------------------------- # Contrast Analysis Thresholds -_CONTRAST_HIGH = 40 # High contrast threshold - allows lower confidence -_CONTRAST_LOW = 25 # Low contrast threshold - requires higher confidence +_CONTRAST_HIGH = 40 +_CONTRAST_LOW = 25 -# Edge Density Thresholds -_EDGE_DENSITY_HIGH = 0.08 # High edge density - clear edges detected -_EDGE_DENSITY_LOW = 0.04 # Low edge density - few edges detected -_EDGE_DENSITY_SHARP = 0.1 # Very sharp edges - can use smaller margin -_EDGE_DENSITY_SOFT = 0.05 # Soft edges - need larger margin +# Edge Density Thresholds +_EDGE_DENSITY_HIGH = 0.08 +_EDGE_DENSITY_LOW = 0.04 +_EDGE_DENSITY_SHARP = 0.1 +_EDGE_DENSITY_SOFT = 0.05 # Complexity Score Calculation Constants -_EDGE_DENSITY_MULTIPLIER = 10 # Weight for edge density in complexity score -_COLOR_VARIANCE_DIVISOR = 1000 # Divisor for color variance normalization +_EDGE_DENSITY_MULTIPLIER = 10 +_COLOR_VARIANCE_DIVISOR = 1000 # Complexity Score Thresholds -_COMPLEXITY_HIGH = 1.5 # High complexity - more iterations needed -_COMPLEXITY_LOW = 0.5 # Low complexity - fewer iterations sufficient +_COMPLEXITY_HIGH = 1.5 +_COMPLEXITY_LOW = 0.5 -# Laplacian Variance Thresholds (noise/sharpness detection) -_LAPLACIAN_HIGH_NOISE = 1000 # High noise level - reduce edge refinement -_LAPLACIAN_SHARP_EDGES = 500 # Sharp, well-defined edges -_LAPLACIAN_SOFT_EDGES = 100 # Soft or unclear edges -_LAPLACIAN_LOW_NOISE = 200 # Low noise level - can use stronger refinement +# Laplacian Variance Thresholds +_LAPLACIAN_HIGH_NOISE = 1000 +_LAPLACIAN_SHARP_EDGES = 500 +_LAPLACIAN_SOFT_EDGES = 100 +_LAPLACIAN_LOW_NOISE = 200 # Brightness Thresholds -_BRIGHTNESS_DARK = 80 # Dark image threshold - lower binary threshold -_BRIGHTNESS_BRIGHT = 180 # Bright image threshold - higher binary threshold +_BRIGHTNESS_DARK = 80 +_BRIGHTNESS_BRIGHT = 180 # Parameter Adjustment Values -_CONFIDENCE_ADJUSTMENT_DOWN = 0.1 # Amount to decrease confidence for clear images -_CONFIDENCE_ADJUSTMENT_UP = 0.15 # Amount to increase confidence for unclear images -_ITERATIONS_ADJUSTMENT = 2 # Amount to adjust iterations -_MARGIN_ADJUSTMENT = 5 # Amount to adjust margin pixels -_REFINEMENT_ADJUSTMENT = 0.2 # Amount to adjust edge refinement strength -_BINARY_THRESHOLD_ADJUSTMENT = 30 # Amount to adjust binary threshold for dark images -_BINARY_THRESHOLD_BRIGHT_ADJUSTMENT = 20 # Amount to adjust for bright images -_EDGE_BLUR_ADJUSTMENT = 0.5 # Amount to adjust edge blur for different conditions +_CONFIDENCE_ADJUSTMENT_DOWN = 0.1 +_CONFIDENCE_ADJUSTMENT_UP = 0.15 +_ITERATIONS_ADJUSTMENT = 2 +_MARGIN_ADJUSTMENT = 5 +_REFINEMENT_ADJUSTMENT = 0.2 +_BINARY_THRESHOLD_ADJUSTMENT = 30 +_BINARY_THRESHOLD_BRIGHT_ADJUSTMENT = 20 +_EDGE_BLUR_ADJUSTMENT = 0.5 # Edge Blur Processing Constants -_SHARP_EDGE_BLUR_THRESHOLD = 0.5 # Threshold below which binary threshold is applied -_KERNEL_SIZE_SCALAR = 4 # Multiplier for blur amount to kernel size conversion -_MAX_KERNEL_SIZE = 31 # Maximum kernel size for performance -_SIGMA_SCALAR = 0.5 # Multiplier for blur amount to sigma conversion +_SHARP_EDGE_BLUR_THRESHOLD = 0.5 +_KERNEL_SIZE_SCALAR = 4 +_MAX_KERNEL_SIZE = 31 +_SIGMA_SCALAR = 0.5 # Parameter Limits -_CONFIDENCE_MIN = 0.3 # Minimum confidence threshold -_CONFIDENCE_MAX = 0.8 # Maximum confidence threshold -_ITERATIONS_MIN = 3 # Minimum GrabCut iterations -_ITERATIONS_MAX = 8 # Maximum GrabCut iterations -_MARGIN_MIN = 10 # Minimum margin pixels -_MARGIN_MAX = 35 # Maximum margin pixels -_REFINEMENT_MIN = 0.4 # Minimum edge refinement strength -_REFINEMENT_MAX = 0.9 # Maximum edge refinement strength -_BINARY_THRESHOLD_MIN = 150 # Minimum binary threshold -_BINARY_THRESHOLD_MAX = 240 # Maximum binary threshold -_EDGE_BLUR_MIN = 0.0 # Minimum edge blur amount -_EDGE_BLUR_MAX = 3.0 # Maximum edge blur amount +_CONFIDENCE_MIN = 0.3 +_CONFIDENCE_MAX = 0.8 +_ITERATIONS_MIN = 3 +_ITERATIONS_MAX = 8 +_MARGIN_MIN = 10 +_MARGIN_MAX = 35 +_REFINEMENT_MIN = 0.4 +_REFINEMENT_MAX = 0.9 +_BINARY_THRESHOLD_MIN = 150 +_BINARY_THRESHOLD_MAX = 240 +_EDGE_BLUR_MIN = 0.0 +_EDGE_BLUR_MAX = 3.0 class GrabCutProcessor: """ Advanced GrabCut background removal with automated object detection. + Combines YOLOv8 object detection with OpenCV's GrabCut algorithm for precise foreground extraction with zero manual intervention. """ - - def __init__(self, - confidence_threshold: float = 0.5, - iterations: int = 5, - margin_pixels: int = 20, - edge_refinement_strength: float = 0.7, - edge_blur_amount: float = 0.0, - binary_threshold: int = 200, - model_path: Optional[str] = None, - bbox_safety_margin: int = 30, - min_bbox_size: int = 64, - fallback_margin_percent: float = 0.2): - """ - Initialize GrabCut processor with configuration. - - Args: - confidence_threshold: Minimum confidence for object detection (0.0-1.0) - iterations: Number of GrabCut iterations - margin_pixels: Pixel margin around detected object - edge_refinement_strength: Strength of edge refinement (0.0-1.0) - edge_blur_amount: Amount of Gaussian blur to apply to mask edges (0.0-10.0) - binary_threshold: Threshold for binary mask conversion - model_path: Optional custom YOLO model path - bbox_safety_margin: Extra pixels beyond detected bbox for safety - min_bbox_size: Minimum bbox dimensions to prevent over-cropping - fallback_margin_percent: Margin percentage for fallback bbox (0.0-0.5) - """ + + # Class-level YOLO model cache — shared across all instances + _yolo_cache: Dict[str, YOLO] = {} + + def __init__( + self, + confidence_threshold: float = 0.5, + iterations: int = 5, + margin_pixels: int = 20, + edge_refinement_strength: float = 0.7, + edge_blur_amount: float = 0.0, + binary_threshold: int = 200, + model_path: Optional[str] = None, + bbox_safety_margin: int = 30, + min_bbox_size: int = 64, + fallback_margin_percent: float = 0.2, + ) -> None: + """Initialize GrabCut processor with configuration.""" self.confidence_threshold = confidence_threshold self.iterations = iterations self.margin_pixels = margin_pixels @@ -112,98 +109,80 @@ class GrabCutProcessor: self.bbox_safety_margin = bbox_safety_margin self.min_bbox_size = min_bbox_size self.fallback_margin_percent = max(0.1, min(0.5, fallback_margin_percent)) - - # Initialize YOLO model - self.yolo_model = None + + # Initialize YOLO model (cached at class level) + self.yolo_model: Optional[YOLO] = None self._initialize_yolo(model_path) - + # Object class mapping - self.target_classes = { - 'person': 0, - 'bicycle': 1, - 'car': 2, - 'motorcycle': 3, - 'airplane': 4, - 'bus': 5, - 'train': 6, - 'truck': 7, - 'boat': 8, - 'bird': 14, - 'cat': 15, - 'dog': 16, - 'horse': 17, - 'sheep': 18, - 'cow': 19, - 'elephant': 20, - 'bear': 21, - 'zebra': 22, - 'giraffe': 23, - 'backpack': 24, - 'umbrella': 25, - 'handbag': 26, - 'tie': 27, - 'suitcase': 28, - 'bottle': 39, - 'wine glass': 40, - 'cup': 41, - 'fork': 42, - 'knife': 43, - 'spoon': 44, - 'bowl': 45, - 'chair': 56, - 'couch': 57, - 'potted plant': 58, - 'bed': 59, - 'dining table': 60, - 'toilet': 61, - 'tv': 62, - 'laptop': 63, - 'mouse': 64, - 'remote': 65, - 'keyboard': 66, - 'cell phone': 67, - 'book': 73, - 'clock': 74, - 'vase': 75, - 'teddy bear': 77, + self.target_classes: Dict[str, int] = { + 'person': 0, 'bicycle': 1, 'car': 2, 'motorcycle': 3, + 'airplane': 4, 'bus': 5, 'train': 6, 'truck': 7, 'boat': 8, + 'bird': 14, 'cat': 15, 'dog': 16, 'horse': 17, 'sheep': 18, + 'cow': 19, 'elephant': 20, 'bear': 21, 'zebra': 22, 'giraffe': 23, + 'backpack': 24, 'umbrella': 25, 'handbag': 26, 'tie': 27, + 'suitcase': 28, 'bottle': 39, 'wine glass': 40, 'cup': 41, + 'fork': 42, 'knife': 43, 'spoon': 44, 'bowl': 45, + 'chair': 56, 'couch': 57, 'potted plant': 58, 'bed': 59, + 'dining table': 60, 'toilet': 61, 'tv': 62, 'laptop': 63, + 'mouse': 64, 'remote': 65, 'keyboard': 66, 'cell phone': 67, + 'book': 73, 'clock': 74, 'vase': 75, 'teddy bear': 77, } - - def _initialize_yolo(self, model_path: Optional[str] = None): - """Initialize YOLO model for object detection.""" + + log.info("grabcut_processor.initialized", + confidence=self.confidence_threshold, + iterations=self.iterations, + yolo_available=self.yolo_model is not None) + + # ------------------------------------------------------------------ + # YOLO initialization + # ------------------------------------------------------------------ + + def _initialize_yolo(self, model_path: Optional[str] = None) -> None: + """Initialize YOLO model, using class-level cache for efficiency.""" + cache_key = model_path if model_path and os.path.exists(model_path) else "yolov8n.pt" + try: - if model_path and os.path.exists(model_path): - self.yolo_model = YOLO(model_path) - else: - # Use YOLOv8 nano model for speed - self.yolo_model = YOLO('yolov8n.pt') + if cache_key not in GrabCutProcessor._yolo_cache: + log.info("grabcut_processor.yolo_loading", model=cache_key) + GrabCutProcessor._yolo_cache[cache_key] = YOLO(cache_key) + # Attempt torch.compile for H100 acceleration + try: + GrabCutProcessor._yolo_cache[cache_key].model = torch.compile( + GrabCutProcessor._yolo_cache[cache_key].model, + ) + log.info("grabcut_processor.torch_compile.success", model=cache_key) + except Exception as compile_err: + log.info("grabcut_processor.torch_compile.skipped", + reason=str(compile_err)) + log.info("grabcut_processor.yolo_loaded", model=cache_key) + self.yolo_model = GrabCutProcessor._yolo_cache[cache_key] except Exception as e: - print(f"Warning: Could not initialize YOLO model: {e}") - print("Falling back to manual rectangle mode") + log.warning("grabcut_processor.yolo_init_failed", + error=str(e), fallback="manual_rectangle") self.yolo_model = None - - def _validate_and_fix_bbox(self, bbox: Tuple[int, int, int, int], image_shape: Tuple[int, int]) -> Tuple[int, int, int, int]: - """ - Validate and fix bounding box to prevent cropping and ensure minimum size. - - Args: - bbox: Original bounding box (x1, y1, x2, y2) - image_shape: Image shape (height, width) - - Returns: - Corrected bounding box (x1, y1, x2, y2) - """ + + # ------------------------------------------------------------------ + # Bounding box validation + # ------------------------------------------------------------------ + + def _validate_and_fix_bbox( + self, + bbox: tuple[int, int, int, int], + image_shape: tuple[int, int], + ) -> tuple[int, int, int, int]: + """Validate and fix bounding box to prevent cropping and ensure minimum size.""" h, w = image_shape x1, y1, x2, y2 = bbox - + # Ensure coordinates are in correct order x1, x2 = min(x1, x2), max(x1, x2) y1, y2 = min(y1, y2), max(y1, y2) - - # Calculate current dimensions + + # Ensure minimum size bbox_w = x2 - x1 bbox_h = y2 - y1 - - # Ensure minimum size + if bbox_w < self.min_bbox_size: center_x = (x1 + x2) // 2 x1 = center_x - self.min_bbox_size // 2 @@ -214,7 +193,7 @@ class GrabCutProcessor: elif x1 < 0: x1 = 0 x2 = min(w, self.min_bbox_size) - + if bbox_h < self.min_bbox_size: center_y = (y1 + y2) // 2 y1 = center_y - self.min_bbox_size // 2 @@ -225,600 +204,520 @@ class GrabCutProcessor: elif y1 < 0: y1 = 0 y2 = min(h, self.min_bbox_size) - + # Add safety margin x1 = max(0, x1 - self.bbox_safety_margin) y1 = max(0, y1 - self.bbox_safety_margin) x2 = min(w, x2 + self.bbox_safety_margin) y2 = min(h, y2 + self.bbox_safety_margin) - + # Final bounds check x1, x2 = max(0, x1), min(w, x2) y1, y2 = max(0, y1), min(h, y2) - + # Ensure we still have a valid bbox if x2 <= x1 or y2 <= y1: - # Fall back to center area if bbox is invalid margin = int(min(h, w) * self.fallback_margin_percent) x1, y1 = margin, margin x2, y2 = w - margin, h - margin - + return (x1, y1, x2, y2) - - def detect_object(self, image: np.ndarray, target_class: Optional[str] = None) -> Optional[Tuple[int, int, int, int, float]]: - """ - Detect primary object in image using YOLO. - - Args: - image: Input image (RGB) - target_class: Specific object class to detect, None for auto-detect - - Returns: - Tuple of (x1, y1, x2, y2, confidence) or None if no object detected - """ + + # ------------------------------------------------------------------ + # Object detection + # ------------------------------------------------------------------ + + @torch.no_grad() + def detect_object( + self, + image: np.ndarray, + target_class: Optional[str] = None, + ) -> Optional[tuple[int, int, int, int, float]]: + """Detect primary object in image using YOLO.""" if self.yolo_model is None: return None - + + _log_gpu_memory("detect_object.start") + try: - # Run YOLO detection results = self.yolo_model(image, conf=self.confidence_threshold, verbose=False) - + if len(results) == 0 or len(results[0].boxes) == 0: + log.debug("grabcut_processor.no_detection") return None - + boxes = results[0].boxes confidences = boxes.conf.cpu().numpy() classes = boxes.cls.cpu().numpy() xyxy = boxes.xyxy.cpu().numpy() - + # Filter by target class if specified if target_class and target_class != 'auto': if target_class in self.target_classes: target_class_id = self.target_classes[target_class] valid_indices = np.where(classes == target_class_id)[0] - + if len(valid_indices) == 0: return None - - # Get highest confidence match for target class + best_idx = valid_indices[np.argmax(confidences[valid_indices])] else: - # Unknown class, use largest detection areas = (xyxy[:, 2] - xyxy[:, 0]) * (xyxy[:, 3] - xyxy[:, 1]) best_idx = np.argmax(areas) else: - # Auto-detect: use highest confidence detection best_idx = np.argmax(confidences) - - # Get bounding box + x1, y1, x2, y2 = xyxy[best_idx].astype(int) confidence = float(confidences[best_idx]) - + + log.info("grabcut_processor.detected", + bbox=(int(x1), int(y1), int(x2), int(y2)), + confidence=round(confidence, 3)) return (x1, y1, x2, y2, confidence) - + except Exception as e: - print(f"Error in object detection: {e}") + log.error("grabcut_processor.detection_error", error=str(e)) return None - - def apply_grabcut(self, image: np.ndarray, bbox: Tuple[int, int, int, int]) -> np.ndarray: - """ - Apply GrabCut algorithm with given bounding box. - - Args: - image: Input image (RGB) - bbox: Bounding box (x1, y1, x2, y2) - - Returns: - Binary mask (0=background, 255=foreground) - """ + finally: + if torch.cuda.is_available(): + torch.cuda.empty_cache() + _log_gpu_memory("detect_object.end") + + # ------------------------------------------------------------------ + # GrabCut core + # ------------------------------------------------------------------ + + def apply_grabcut( + self, + image: np.ndarray, + bbox: tuple[int, int, int, int], + ) -> np.ndarray: + """Apply GrabCut algorithm with given bounding box.""" h, w = image.shape[:2] - + log.debug("grabcut_processor.apply_grabcut", image_size=(w, h)) + # Validate and fix bounding box x1, y1, x2, y2 = self._validate_and_fix_bbox(bbox, (h, w)) - - # Add additional margin for GrabCut processing + + # Add additional margin x1 = max(0, x1 - self.margin_pixels) y1 = max(0, y1 - self.margin_pixels) x2 = min(w, x2 + self.margin_pixels) y2 = min(h, y2 + self.margin_pixels) - - # Convert to GrabCut rectangle format (x, y, width, height) + rect = (x1, y1, x2 - x1, y2 - y1) - - # Initialize mask mask = np.zeros((h, w), np.uint8) - - # Initialize foreground and background models bgd_model = np.zeros((1, 65), np.float64) fgd_model = np.zeros((1, 65), np.float64) - + try: - # Apply GrabCut - cv2.grabCut(image, mask, rect, bgd_model, fgd_model, - self.iterations, cv2.GC_INIT_WITH_RECT) - - # Convert mask to binary - # GrabCut mask values: 0=BG, 1=FG, 2=PR_BG, 3=PR_FG + cv2.grabCut(image, mask, rect, bgd_model, fgd_model, + self.iterations, cv2.GC_INIT_WITH_RECT) output_mask = np.where((mask == 2) | (mask == 0), 0, 255).astype('uint8') - return output_mask - except Exception as e: - print(f"Error in GrabCut: {e}") - # Return rectangle mask as fallback + log.error("grabcut_processor.grabcut_error", error=str(e)) fallback_mask = np.zeros((h, w), dtype=np.uint8) fallback_mask[y1:y2, x1:x2] = 255 return fallback_mask - + + # ------------------------------------------------------------------ + # Edge refinement + # ------------------------------------------------------------------ + def refine_edges(self, mask: np.ndarray, image: np.ndarray) -> np.ndarray: - """ - Apply edge refinement to improve mask quality with artifact-free edge blur. - - Args: - mask: Binary mask - image: Original image for guided filtering - - Returns: - Refined mask - """ + """Apply edge refinement to improve mask quality with artifact-free edge blur.""" if self.edge_refinement_strength <= 0 and self.edge_blur_amount <= 0: return mask - - # Convert to float for processing + mask_float = mask.astype(np.float32) / 255.0 - - # Step 1: Apply bilateral filter for edge-preserving smoothing + + # Step 1: Bilateral filter for edge-preserving smoothing if self.edge_refinement_strength > 0: - refined = cv2.bilateralFilter( - mask_float, - d=9, - sigmaColor=0.1, - sigmaSpace=7 - ) - - # Step 2: Apply guided filter using original image + refined = cv2.bilateralFilter(mask_float, d=9, sigmaColor=0.1, sigmaSpace=7) + + # Step 2: Guided filter try: - gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY) - # Ensure proper data types for guided filter - gray = gray.astype(np.float32) + gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY).astype(np.float32) refined_input = refined.astype(np.float32) - refined = cv2.ximgproc.guidedFilter( - guide=gray, - src=refined_input, - radius=4, - eps=self.edge_refinement_strength * 0.01 + guide=gray, src=refined_input, radius=4, + eps=self.edge_refinement_strength * 0.01, ) except Exception as e: - print(f"Guided filter failed, using bilateral filter only: {e}") - # Continue with just the bilateral filter result - - # Step 3: Morphological operations to clean up + log.warning("grabcut_processor.guided_filter_failed", error=str(e)) + + # Step 3: Morphological cleanup kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (5, 5)) refined = cv2.morphologyEx(refined, cv2.MORPH_CLOSE, kernel) refined = cv2.morphologyEx(refined, cv2.MORPH_OPEN, kernel) else: refined = mask_float - - # Step 4: Apply edge blur BEFORE binary thresholding to prevent artifacts + + # Step 4: Edge blur BEFORE binary thresholding if self.edge_blur_amount > 0: - # Convert to 0-255 range for edge blur processing refined_255 = (refined * 255).astype(np.uint8) blurred = self._apply_edge_blur(refined_255) - # Convert back to float for final processing refined = blurred.astype(np.float32) / 255.0 - - # Step 5: Apply binary threshold as final step (only if no blur or minimal blur) + + # Step 5: Binary threshold if self.edge_blur_amount <= _SHARP_EDGE_BLUR_THRESHOLD: - # Apply binary threshold to eliminate semi-transparency for sharp edges - _, binary = cv2.threshold(refined, self.binary_threshold / 255.0, 1.0, cv2.THRESH_BINARY) + _, binary = cv2.threshold( + refined, self.binary_threshold / 255.0, 1.0, cv2.THRESH_BINARY, + ) return (binary * 255).astype(np.uint8) else: - # Keep soft edges when significant blur is applied return (refined * 255).astype(np.uint8) - + def _apply_edge_blur(self, mask: np.ndarray) -> np.ndarray: - """ - Apply Gaussian blur to mask edges for softer transitions. - - Args: - mask: Binary mask to blur - - Returns: - Blurred mask with soft edges - """ + """Apply Gaussian blur to mask edges for softer transitions.""" if self.edge_blur_amount <= 0: return mask - - # Calculate dynamic kernel size based on blur amount - # Ensure kernel size is odd and reasonable - kernel_size = max(3, int(self.edge_blur_amount * _KERNEL_SIZE_SCALAR) | 1) # Force odd using bitwise OR - # With edge_blur_amount max 10.0, kernel_size max is 41, cap for performance + + kernel_size = max(3, int(self.edge_blur_amount * _KERNEL_SIZE_SCALAR) | 1) kernel_size = min(kernel_size, _MAX_KERNEL_SIZE) - - # Calculate sigma based on blur amount for consistent results - # Using standard relationship: sigma = 0.3 * ((ksize-1) * 0.5 - 1) + 0.8 - # But simplified for direct control: sigma = blur_amount * scalar sigma = self.edge_blur_amount * _SIGMA_SCALAR - - # Apply Gaussian blur with calculated sigma - blurred_mask = cv2.GaussianBlur(mask, (kernel_size, kernel_size), sigma) - - return blurred_mask - - def auto_adjust_parameters(self, image: np.ndarray) -> dict: - """ - Automatically adjust parameters based on image analysis. - Analyzes image characteristics and returns optimal parameter adjustments. - - Args: - image: Input image for analysis (RGB format) - - Returns: - Dictionary of adjusted parameters - """ + + return cv2.GaussianBlur(mask, (kernel_size, kernel_size), sigma) + + # ------------------------------------------------------------------ + # Auto parameter adjustment + # ------------------------------------------------------------------ + + def auto_adjust_parameters(self, image: np.ndarray) -> dict[str, Any]: + """Automatically adjust parameters based on image analysis.""" + log.debug("grabcut_processor.auto_adjust.start") + gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY) h, w = gray.shape total_pixels = h * w - - # Analyze image characteristics - # 1. Edge density analysis + edges = cv2.Canny(gray, 50, 150) edge_density = np.sum(edges > 0) / total_pixels - - # 2. Contrast analysis contrast = gray.std() - - # 3. Brightness distribution brightness = gray.mean() - - # 4. Color variance analysis color_variance = np.var(image.reshape(-1, 3), axis=0).mean() - - # 5. Noise level estimation (using Laplacian variance) laplacian_var = cv2.Laplacian(gray, cv2.CV_64F).var() - - # Initialize adjustments dictionary - adjustments = {} - - # Adjust confidence_threshold based on contrast and edge clarity + + adjustments: dict[str, Any] = {} + + # Confidence threshold base_confidence = self.confidence_threshold if contrast > _CONTRAST_HIGH and edge_density > _EDGE_DENSITY_HIGH: - # High contrast, clear edges -> can lower confidence threshold - adjustments['confidence_threshold'] = max(_CONFIDENCE_MIN, base_confidence - _CONFIDENCE_ADJUSTMENT_DOWN) + adjustments['confidence_threshold'] = max( + _CONFIDENCE_MIN, base_confidence - _CONFIDENCE_ADJUSTMENT_DOWN) elif contrast < _CONTRAST_LOW or edge_density < _EDGE_DENSITY_LOW: - # Low contrast or few edges -> need higher confidence threshold - adjustments['confidence_threshold'] = min(_CONFIDENCE_MAX, base_confidence + _CONFIDENCE_ADJUSTMENT_UP) - - # Adjust iterations based on image complexity + adjustments['confidence_threshold'] = min( + _CONFIDENCE_MAX, base_confidence + _CONFIDENCE_ADJUSTMENT_UP) + + # Iterations base_iterations = self.iterations - complexity_score = (edge_density * _EDGE_DENSITY_MULTIPLIER) + (color_variance / _COLOR_VARIANCE_DIVISOR) + complexity_score = (edge_density * _EDGE_DENSITY_MULTIPLIER) + ( + color_variance / _COLOR_VARIANCE_DIVISOR) if complexity_score > _COMPLEXITY_HIGH: - # High complexity -> more iterations needed - adjustments['iterations'] = min(_ITERATIONS_MAX, base_iterations + _ITERATIONS_ADJUSTMENT) + adjustments['iterations'] = min( + _ITERATIONS_MAX, base_iterations + _ITERATIONS_ADJUSTMENT) elif complexity_score < _COMPLEXITY_LOW: - # Low complexity -> fewer iterations sufficient - adjustments['iterations'] = max(_ITERATIONS_MIN, base_iterations - _ITERATIONS_ADJUSTMENT) - - # Adjust margin_pixels based on edge sharpness and object size estimation + adjustments['iterations'] = max( + _ITERATIONS_MIN, base_iterations - _ITERATIONS_ADJUSTMENT) + + # Margin pixels base_margin = self.margin_pixels if edge_density > _EDGE_DENSITY_SHARP and laplacian_var > _LAPLACIAN_SHARP_EDGES: - # Sharp, well-defined edges -> can use smaller margin - adjustments['margin_pixels'] = max(_MARGIN_MIN, base_margin - _MARGIN_ADJUSTMENT) + adjustments['margin_pixels'] = max( + _MARGIN_MIN, base_margin - _MARGIN_ADJUSTMENT) elif edge_density < _EDGE_DENSITY_SOFT or laplacian_var < _LAPLACIAN_SOFT_EDGES: - # Soft or unclear edges -> need larger margin - adjustments['margin_pixels'] = min(_MARGIN_MAX, base_margin + _MARGIN_ADJUSTMENT) - - # Adjust edge_refinement_strength based on noise level + adjustments['margin_pixels'] = min( + _MARGIN_MAX, base_margin + _MARGIN_ADJUSTMENT) + + # Edge refinement strength base_refinement = self.edge_refinement_strength if laplacian_var > _LAPLACIAN_HIGH_NOISE: - # High noise -> reduce edge refinement to avoid artifacts - adjustments['edge_refinement_strength'] = max(_REFINEMENT_MIN, base_refinement - _REFINEMENT_ADJUSTMENT) + adjustments['edge_refinement_strength'] = max( + _REFINEMENT_MIN, base_refinement - _REFINEMENT_ADJUSTMENT) elif laplacian_var < _LAPLACIAN_LOW_NOISE: - # Low noise -> can use stronger edge refinement - adjustments['edge_refinement_strength'] = min(_REFINEMENT_MAX, base_refinement + _REFINEMENT_ADJUSTMENT) - - # Adjust binary_threshold based on brightness distribution + adjustments['edge_refinement_strength'] = min( + _REFINEMENT_MAX, base_refinement + _REFINEMENT_ADJUSTMENT) + + # Binary threshold base_threshold = self.binary_threshold if brightness < _BRIGHTNESS_DARK: - # Dark image -> lower threshold - adjustments['binary_threshold'] = max(_BINARY_THRESHOLD_MIN, base_threshold - _BINARY_THRESHOLD_ADJUSTMENT) + adjustments['binary_threshold'] = max( + _BINARY_THRESHOLD_MIN, base_threshold - _BINARY_THRESHOLD_ADJUSTMENT) elif brightness > _BRIGHTNESS_BRIGHT: - # Bright image -> higher threshold - adjustments['binary_threshold'] = min(_BINARY_THRESHOLD_MAX, base_threshold + _BINARY_THRESHOLD_BRIGHT_ADJUSTMENT) - - # Adjust edge_blur_amount based on edge characteristics and noise level + adjustments['binary_threshold'] = min( + _BINARY_THRESHOLD_MAX, base_threshold + _BINARY_THRESHOLD_BRIGHT_ADJUSTMENT) + + # Edge blur amount base_blur = self.edge_blur_amount if laplacian_var > _LAPLACIAN_HIGH_NOISE or edge_density < _EDGE_DENSITY_SOFT: - # Noisy or soft edges -> apply blur to smooth transitions - adjustments['edge_blur_amount'] = min(_EDGE_BLUR_MAX, base_blur + _EDGE_BLUR_ADJUSTMENT) + adjustments['edge_blur_amount'] = min( + _EDGE_BLUR_MAX, base_blur + _EDGE_BLUR_ADJUSTMENT) elif edge_density > _EDGE_DENSITY_SHARP and laplacian_var > _LAPLACIAN_SHARP_EDGES: - # Sharp, clean edges -> minimal blur to preserve detail - adjustments['edge_blur_amount'] = max(_EDGE_BLUR_MIN, base_blur - _EDGE_BLUR_ADJUSTMENT) + adjustments['edge_blur_amount'] = max( + _EDGE_BLUR_MIN, base_blur - _EDGE_BLUR_ADJUSTMENT) elif contrast < _CONTRAST_LOW: - # Low contrast images benefit from edge blur for smoother results - adjustments['edge_blur_amount'] = min(_EDGE_BLUR_MAX, base_blur + (_EDGE_BLUR_ADJUSTMENT * 0.5)) - + adjustments['edge_blur_amount'] = min( + _EDGE_BLUR_MAX, base_blur + (_EDGE_BLUR_ADJUSTMENT * 0.5)) + + log.info("grabcut_processor.auto_adjust.done", adjustments=adjustments) return adjustments - + + # ------------------------------------------------------------------ + # Pixel art detection + # ------------------------------------------------------------------ + def _detect_pixel_art_characteristics(self, image: np.ndarray) -> bool: - """ - Detect if image has pixel art characteristics. - Analyzes color count, edge sharpness, and dithering patterns. - - Args: - image: Input image (RGB format) - - Returns: - Boolean indicating if image appears to be pixel art - """ - # Convert to grayscale for analysis + """Detect if image has pixel art characteristics.""" gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY) h, w = gray.shape total_pixels = h * w - - # 1. Color count analysis - pixel art typically has limited colors + unique_colors = len(np.unique(image.reshape(-1, 3), axis=0)) color_density = unique_colors / total_pixels - - # 2. Edge sharpness analysis - pixel art has very sharp edges - # Calculate edge sharpness using Laplacian variance laplacian_var = cv2.Laplacian(gray, cv2.CV_64F).var() - - # 3. Dither pattern detection - look for regular patterns - # Use FFT to detect repeating patterns + fft = np.fft.fft2(gray) fft_magnitude = np.abs(fft) - - # Look for peaks in frequency domain that suggest regular patterns - # Remove DC component and low frequencies fft_magnitude[0:5, 0:5] = 0 peak_ratio = np.max(fft_magnitude) / np.mean(fft_magnitude) - - # 4. Small dimensions often indicate pixel art + is_small = w <= 512 or h <= 512 - - # Decision logic based on multiple factors + pixel_art_score = 0 - - # Low color density suggests pixel art - if color_density < 0.01: # Very limited colors + if color_density < 0.01: pixel_art_score += 3 - elif color_density < 0.05: # Limited colors + elif color_density < 0.05: pixel_art_score += 2 - elif color_density < 0.1: # Somewhat limited colors + elif color_density < 0.1: pixel_art_score += 1 - - # High edge sharpness suggests pixel art - if laplacian_var > 1000: # Very sharp edges + + if laplacian_var > 1000: pixel_art_score += 2 - elif laplacian_var > 500: # Sharp edges + elif laplacian_var > 500: pixel_art_score += 1 - - # Regular patterns suggest dithering/pixel art - if peak_ratio > 50: # Strong regular patterns + + if peak_ratio > 50: pixel_art_score += 2 - elif peak_ratio > 20: # Moderate patterns + elif peak_ratio > 20: pixel_art_score += 1 - - # Small dimensions bonus + if is_small: pixel_art_score += 1 - - # Threshold for pixel art detection + return pixel_art_score >= 3 - - def process_with_grabcut(self, image: np.ndarray, target_class: Optional[str] = None) -> Dict: - """ - Complete GrabCut processing pipeline with automated object detection. - - Args: - image: Input image (RGB) - target_class: Target object class or 'auto' for automatic - - Returns: - Dictionary containing: - - rgba_image: RGBA image with transparent background - - mask: Alpha mask - - bbox: Detected bounding box (x1, y1, x2, y2) - - confidence: Detection confidence - - processing_time_ms: Processing time in milliseconds - - success: Whether processing succeeded - """ + + # ------------------------------------------------------------------ + # Main processing pipelines + # ------------------------------------------------------------------ + + @torch.no_grad() + def process_with_grabcut( + self, + image: np.ndarray, + target_class: Optional[str] = None, + ) -> dict[str, Any]: + """Complete GrabCut processing pipeline with automated object detection.""" start_time = time.time() h, w = image.shape[:2] - - # Initialize result - result = { + + log.info("grabcut_processor.process.start", size=(w, h), target_class=target_class) + _log_gpu_memory("process_with_grabcut.start") + + result: dict[str, Any] = { 'rgba_image': None, 'mask': None, 'bbox': None, 'confidence': 0.0, 'processing_time_ms': 0, - 'success': False + 'success': False, } - + # Ensure RGB format if len(image.shape) == 2: image = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB) elif image.shape[2] == 4: image = image[:, :, :3] - + # Step 1: Detect object detection = self.detect_object(image, target_class) - + if detection is None: - # Fallback: use entire image with adaptive margin margin = int(min(h, w) * self.fallback_margin_percent) detection = (margin, margin, w - margin, h - margin, 0.5) - print(f"No object detected, using fallback rectangle with {self.fallback_margin_percent*100}% margin") - + log.info("grabcut_processor.using_fallback", + margin_pct=self.fallback_margin_percent) + x1, y1, x2, y2, confidence = detection result['bbox'] = (x1, y1, x2, y2) result['confidence'] = confidence - + # Step 2: Apply GrabCut mask = self.apply_grabcut(image, (x1, y1, x2, y2)) - + # Step 3: Refine edges try: mask = self.refine_edges(mask, image) except Exception as e: - print(f"Edge refinement skipped: {e}") - + log.warning("grabcut_processor.edge_refinement_skipped", error=str(e)) + # Step 4: Create RGBA output rgba = np.zeros((h, w, 4), dtype=np.uint8) rgba[:, :, :3] = image rgba[:, :, 3] = mask - - # Update result + result['rgba_image'] = rgba result['mask'] = mask result['success'] = True result['processing_time_ms'] = int((time.time() - start_time) * 1000) - + + if torch.cuda.is_available(): + torch.cuda.empty_cache() + _log_gpu_memory("process_with_grabcut.end") + + log.info("grabcut_processor.process.done", + time_ms=result['processing_time_ms'], + confidence=round(confidence, 3)) return result - - def process_with_initial_mask(self, image: np.ndarray, initial_mask: np.ndarray, - target_class: Optional[str] = None) -> Dict: - """ - Process image with an initial mask from previous processing. - Useful for refining results from other background removal methods. - - Args: - image: Input image (RGB) - initial_mask: Initial mask from previous processing - target_class: Target object class for detection - - Returns: - Dictionary with processing results - """ + + @torch.no_grad() + def process_with_initial_mask( + self, + image: np.ndarray, + initial_mask: np.ndarray, + target_class: Optional[str] = None, + ) -> dict[str, Any]: + """Process image with an initial mask from previous processing.""" start_time = time.time() h, w = image.shape[:2] - - result = { + + log.info("grabcut_processor.refine.start", size=(w, h)) + _log_gpu_memory("process_with_initial_mask.start") + + result: dict[str, Any] = { 'rgba_image': None, 'mask': None, 'bbox': None, 'confidence': 0.0, 'processing_time_ms': 0, - 'success': False + 'success': False, } - + # Ensure proper formats if len(image.shape) == 2: image = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB) elif image.shape[2] == 4: image = image[:, :, :3] - - # Get bounding box from initial mask if no detection + + # Get bounding box detection = self.detect_object(image, target_class) - + if detection is None: - # Find bounding box from initial mask if initial_mask.max() > 0: coords = np.where(initial_mask > 127) - # Correct axis ordering: coords[0] is y (rows), coords[1] is x (columns) y1, y2 = coords[0].min(), coords[0].max() x1, x2 = coords[1].min(), coords[1].max() - # Validate the extracted bbox raw_bbox = (x1, y1, x2, y2) x1, y1, x2, y2 = self._validate_and_fix_bbox(raw_bbox, (h, w)) detection = (x1, y1, x2, y2, 0.8) else: - # Fallback to adaptive margin margin = int(min(h, w) * self.fallback_margin_percent) detection = (margin, margin, w - margin, h - margin, 0.5) - + x1, y1, x2, y2, confidence = detection result['bbox'] = (x1, y1, x2, y2) result['confidence'] = confidence - + # Initialize GrabCut mask from initial mask grabcut_mask = np.zeros((h, w), np.uint8) - - # Convert initial mask to GrabCut format - # 0=BG, 1=FG, 2=PR_BG, 3=PR_FG - grabcut_mask[initial_mask > 200] = cv2.GC_FGD # Definite foreground - grabcut_mask[(initial_mask > 50) & (initial_mask <= 200)] = cv2.GC_PR_FGD # Probable foreground - grabcut_mask[(initial_mask > 0) & (initial_mask <= 50)] = cv2.GC_PR_BGD # Probable background - # grabcut_mask[initial_mask == 0] remains cv2.GC_BGD (0) - - # Apply GrabCut with mask initialization + grabcut_mask[initial_mask > 200] = cv2.GC_FGD + grabcut_mask[(initial_mask > 50) & (initial_mask <= 200)] = cv2.GC_PR_FGD + grabcut_mask[(initial_mask > 0) & (initial_mask <= 50)] = cv2.GC_PR_BGD + bgd_model = np.zeros((1, 65), np.float64) fgd_model = np.zeros((1, 65), np.float64) - + try: cv2.grabCut(image, grabcut_mask, None, bgd_model, fgd_model, - self.iterations, cv2.GC_INIT_WITH_MASK) - - # Convert to binary mask - output_mask = np.where((grabcut_mask == 2) | (grabcut_mask == 0), 0, 255).astype('uint8') - - # Refine edges + self.iterations, cv2.GC_INIT_WITH_MASK) + output_mask = np.where( + (grabcut_mask == 2) | (grabcut_mask == 0), 0, 255, + ).astype('uint8') output_mask = self.refine_edges(output_mask, image) - except Exception as e: - print(f"Error in GrabCut with initial mask: {e}") + log.error("grabcut_processor.mask_refinement_error", error=str(e)) output_mask = initial_mask - + # Create RGBA output rgba = np.zeros((h, w, 4), dtype=np.uint8) rgba[:, :, :3] = image rgba[:, :, 3] = output_mask - + result['rgba_image'] = rgba result['mask'] = output_mask result['success'] = True result['processing_time_ms'] = int((time.time() - start_time) * 1000) - + + if torch.cuda.is_available(): + torch.cuda.empty_cache() + _log_gpu_memory("process_with_initial_mask.end") + + log.info("grabcut_processor.refine.done", + time_ms=result['processing_time_ms']) return result -def create_fallback_processor(): - """ - Create a fallback processor that works without YOLO. - Uses simple image analysis to find the main subject. - """ - +# --------------------------------------------------------------------------- +# Helper: GPU memory logging +# --------------------------------------------------------------------------- + +def _log_gpu_memory(tag: str) -> None: + """Log GPU memory statistics if CUDA is available.""" + if torch.cuda.is_available(): + allocated = torch.cuda.memory_allocated() / 1e9 + reserved = torch.cuda.memory_reserved() / 1e9 + log.debug("gpu_memory", tag=tag, + allocated_gb=round(allocated, 3), + reserved_gb=round(reserved, 3)) + + +# --------------------------------------------------------------------------- +# Fallback processor (no YOLO) +# --------------------------------------------------------------------------- + +def create_fallback_processor() -> type[GrabCutProcessor]: + """Create a fallback processor class that works without YOLO.""" + class FallbackGrabCutProcessor(GrabCutProcessor): - def __init__(self, **kwargs): - # Initialize without YOLO + def __init__(self, **kwargs: Any) -> None: super().__init__(**kwargs) self.yolo_model = None - - def detect_object(self, image: np.ndarray, target_class: Optional[str] = None) -> Optional[Tuple[int, int, int, int, float]]: - """ - Fallback object detection using image analysis. - Finds the largest connected component that's not background. - """ + + def detect_object( + self, + image: np.ndarray, + target_class: Optional[str] = None, + ) -> Optional[tuple[int, int, int, int, float]]: + """Fallback object detection using image analysis.""" h, w = image.shape[:2] - - # Convert to grayscale gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY) - - # Apply edge detection edges = cv2.Canny(gray, 50, 150) - - # Find contours - contours, _ = cv2.findContours(edges, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) - + contours, _ = cv2.findContours( + edges, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) + if not contours: - # Return center region with adaptive margin margin = int(min(h, w) * self.fallback_margin_percent) return (margin, margin, w - margin, h - margin, 0.5) - - # Find largest contour + largest_contour = max(contours, key=cv2.contourArea) x, y, cw, ch = cv2.boundingRect(largest_contour) - - # Add some margin + margin = 20 x1 = max(0, x - margin) y1 = max(0, y - margin) x2 = min(w, x + cw + margin) y2 = min(h, y + ch + margin) - + return (x1, y1, x2, y2, 0.7) - - return FallbackGrabCutProcessor \ No newline at end of file + + return FallbackGrabCutProcessor diff --git a/pyproject.toml b/pyproject.toml index 60592dd..3260a59 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -3,7 +3,7 @@ name = "ComfyUI-TransparencyBackgroundRemover" version = "1.1.2" description = "Automatic background removal and transparency generation for ComfyUI" license = { file = "LICENSE" } -requires-python = ">=3.8" +requires-python = ">=3.10" classifiers = [ "Operating System :: OS Independent" ] diff --git a/requirements.txt b/requirements.txt index c067947..e15a188 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,7 +1,9 @@ -torch -numpy -Pillow -opencv-python -opencv-contrib-python -scikit-learn -ultralytics +torch>=2.1.0 +numpy>=1.24.0 +Pillow>=10.4.0 +opencv-python>=4.10.0 +opencv-contrib-python>=4.10.0 +scikit-learn>=1.4.0 +ultralytics>=8.3.0 +structlog>=24.4.0 +pydantic>=2.7.0 diff --git a/src/__init__.py b/src/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/validation.py b/src/validation.py new file mode 100644 index 0000000..a37dbb3 --- /dev/null +++ b/src/validation.py @@ -0,0 +1,150 @@ +"""Pydantic validation models for all user-facing GrabCut node parameters. + +Security-hardened: sanitizes and validates all inputs before GPU execution. +""" +from __future__ import annotations + +from typing import Annotated, Any + +from pydantic import BaseModel, Field, model_validator + + +class GrabCutParams(BaseModel): + """Parameters for the core GrabCut algorithm.""" + + iterations: int = Field(default=5, ge=1, le=100, description="GrabCut iterations (1-100)") + margin: int = Field(default=20, ge=0, le=500, description="Margin around detected object in pixels") + edge_threshold: float = Field( + default=0.5, ge=0.0, le=1.0, + description="Edge detection threshold for auto-adjustment" + ) + confidence_threshold: float = Field( + default=0.5, ge=0.0, le=1.0, + description="YOLO confidence threshold" + ) + + model_config = {"str_strip_whitespace": True} + + +class ScalingParams(BaseModel): + """Parameters for image scaling / resize preprocessing.""" + + target_long_edge: int = Field( + default=1024, ge=256, le=4096, + description="Target size for longest image edge" + ) + maintain_aspect: bool = Field( + default=True, + description="Preserve aspect ratio during scaling" + ) + scaling_method: str = Field( + default="auto", + description="Scaling method: auto, nearest, bilinear, lanczos, power-of-8" + ) + + @model_validator(mode="after") + def check_scaling_method(self) -> "ScalingParams": + valid = {"auto", "nearest", "bilinear", "lanczos", "power-of-8"} + if self.scaling_method not in valid: + raise ValueError( + f"Invalid scaling_method '{self.scaling_method}'. " + f"Must be one of: {', '.join(sorted(valid))}" + ) + return self + + +class MaskParams(BaseModel): + """Parameters for mask post-processing.""" + + edge_blur_amount: int = Field( + default=0, ge=0, le=20, + description="Gaussian blur kernel size for mask edge softening (0=off)" + ) + invert_mask: bool = Field( + default=False, + description="Invert the output mask (foreground becomes background)" + ) + edge_refinement_strength: float = Field( + default=0.7, ge=0.0, le=1.0, + description="Strength of edge refinement pass" + ) + + model_config = {"str_strip_whitespace": True} + + +class BBoxParams(BaseModel): + """Bounding-box / detection parameters.""" + + bbox_safety_margin: int = Field( + default=30, ge=0, le=200, + description="Safety margin added to detected bounding boxes in pixels" + ) + min_bbox_size: int = Field( + default=64, ge=8, le=1024, + description="Minimum bounding box side length in pixels" + ) + fallback_margin_percent: float = Field( + default=0.2, ge=0.01, le=0.5, + description="When no object is detected, use this fraction of image size as margin" + ) + + +class GrabCutNodeParams(GrabCutParams, ScalingParams, MaskParams, BBoxParams): + """Combined validation model for AutoGrabCutRemover node parameters. + + Used to validate ALL user inputs in a single pass before any GPU work. + """ + + binary_threshold: int = Field( + default=200, ge=0, le=255, + description="Threshold for binary mask generation" + ) + output_format: str = Field( + default="RGBA", + description="Output format: RGBA or MASK" + ) + auto_adjust: bool = Field( + default=True, + description="Enable automatic parameter adjustment based on image analysis" + ) + + @model_validator(mode="after") + def check_output_format(self) -> "GrabCutNodeParams": + valid = {"RGBA", "MASK"} + if self.output_format not in valid: + raise ValueError( + f"Invalid output_format '{self.output_format}'. Must be one of: {', '.join(sorted(valid))}" + ) + return self + + +# --------------------------------------------------------------------------- +# Convenience validator functions (call these at the top of each execute()) +# --------------------------------------------------------------------------- + +def validate_grabcut_params(**kwargs: Any) -> GrabCutParams: + """Validate and return GrabCutParams. Raises ValueError on failure.""" + return GrabCutParams(**kwargs) + + +def validate_scaling_params(**kwargs: Any) -> ScalingParams: + """Validate and return ScalingParams. Raises ValueError on failure.""" + return ScalingParams(**kwargs) + + +def validate_mask_params(**kwargs: Any) -> MaskParams: + """Validate and return MaskParams. Raises ValueError on failure.""" + return MaskParams(**kwargs) + + +def validate_bbox_params(**kwargs: Any) -> BBoxParams: + """Validate and return BBoxParams. Raises ValueError on failure.""" + return BBoxParams(**kwargs) + + +def validate_node_params(**kwargs: Any) -> GrabCutNodeParams: + """Validate ALL parameters for a GrabCut node. + + Call this at the START of every execute() method before any processing. + """ + return GrabCutNodeParams(**kwargs)