feat: GrabCut H100 optimization — torch.no_grad, Pydantic validation, structlog, type hints

- grabcut_nodes.py: torch.no_grad() on remove_background + refine_mask,
  dtype=torch.float32 on all tensor conversions, torch.cuda.empty_cache()
  in finally blocks, Pydantic validation at start of every execute()
- grabcut_remover.py: Full type-hinting on GrabCutProcessor (mypy-ready),
  structlog observability with GPU memory tracking (_log_gpu_memory),
  torch.no_grad() + empty_cache() on all inference paths
- requirements.txt: Pin torch>=2.1.0, Pillow>=10.4.0 (CVE fix),
  opencv>=4.10.0, structlog>=24.4.0, pydantic>=2.7.0
- pyproject.toml: python>=3.10 (from 3.8)
- src/validation.py: NEW — Pydantic models for all user-facing params
  (GrabCutParams, ScalingParams, MaskParams, BBoxParams, GrabCutNodeParams)
- src/__init__.py: NEW — package marker
- 6/6 pytest tests pass, node import test passes

Ref: CD-14 | notion:32c40bf6-db23-8001-92b8-dfe4d91dccb4
This commit is contained in:
limbicnation
2026-03-29 06:14:15 +02:00
parent e582d1d22d
commit 3094651eb8
6 changed files with 719 additions and 592 deletions
+149 -73
View File
@@ -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)
+410 -511
View File
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -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"
]
+9 -7
View File
@@ -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
View File
+150
View File
@@ -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)