188 lines
7.4 KiB
Python
188 lines
7.4 KiB
Python
"""
|
|
Roto Mask Node - Backend implementation
|
|
Converts curve data to rasterized masks
|
|
"""
|
|
|
|
import json
|
|
import numpy as np
|
|
import torch
|
|
from typing import Dict, List, Tuple, Optional, Any
|
|
|
|
from .core.curve_validator import validate_schema, validate_curve
|
|
from .mask.generator import generate_single_mask
|
|
from .image.loader import load_image_from_disk
|
|
from .image.processor import resize_mask_to_image, apply_mask_overlay
|
|
from .image.utils import create_empty_mask
|
|
|
|
|
|
class RotoMaskNode:
|
|
"""
|
|
ComfyUI node for generating masks from curve data.
|
|
Accepts image input, curve JSON data, and frame index.
|
|
Outputs a single-channel mask image.
|
|
"""
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls) -> Dict[str, Dict[str, Any]]:
|
|
"""
|
|
Define input parameters for the node.
|
|
|
|
Returns:
|
|
Dictionary specifying required input types and their constraints
|
|
"""
|
|
return {
|
|
"required": {
|
|
"curve_json": (
|
|
"STRING",
|
|
{
|
|
"multiline": True,
|
|
"default": '{"meta":{"version":1,"width":512,"height":512},"shapes":{}}',
|
|
},
|
|
),
|
|
"frame_index": ("INT", {"default": 1, "min": 1, "max": 9999, "step": 1}),
|
|
},
|
|
"optional": {
|
|
"frame_count": ("INT", {"default": 1, "min": 1, "max": 9999, "step": 1}),
|
|
},
|
|
}
|
|
|
|
RETURN_TYPES = ("IMAGE", "MASK", "IMAGE", "IMAGE", "STRING")
|
|
RETURN_NAMES = ("Mask as image", "Mask", "Masked Image", "Loaded Image", "Mask Data")
|
|
FUNCTION = "generate_mask"
|
|
CATEGORY = "mask"
|
|
|
|
def generate_mask(
|
|
self, curve_json: str, frame_index: int, frame_count: Optional[int] = None
|
|
) -> Tuple[torch.Tensor]:
|
|
"""
|
|
Main execution function for the node.
|
|
|
|
Args:
|
|
curve_json: JSON string containing curve data
|
|
frame_index: Starting frame number to render
|
|
frame_count: Number of frames to render (1 = single frame, >1 = batch sequence)
|
|
|
|
Returns:
|
|
Tuple containing mask tensor as RGB image (B, H, W, 3)
|
|
"""
|
|
# Default frame_count to 1 if not provided
|
|
if frame_count is None:
|
|
frame_count = 1
|
|
|
|
# For batch rendering (frame_count > 1), always start from frame 1
|
|
# frame_index is only used for single-frame preview mode
|
|
if frame_count > 1:
|
|
frame_index = 1
|
|
|
|
# Parse curve data
|
|
try:
|
|
curve_data = json.loads(curve_json)
|
|
except json.JSONDecodeError as e:
|
|
print(f"[RotoMaskNode] Invalid JSON: {e}")
|
|
# Return empty masks on error (default 512x512)
|
|
empty_mask = create_empty_mask(frame_count, 512, 512)
|
|
empty_mask_single = torch.zeros((frame_count, 512, 512), dtype=torch.float32)
|
|
return (empty_mask, empty_mask_single, empty_mask, empty_mask, curve_json)
|
|
|
|
# Validate schema
|
|
if not validate_schema(curve_data):
|
|
print("[RotoMaskNode] Invalid curve data schema")
|
|
# Return empty masks on error (default 512x512)
|
|
empty_mask = create_empty_mask(frame_count, 512, 512)
|
|
empty_mask_single = torch.zeros((frame_count, 512, 512), dtype=torch.float32)
|
|
return (empty_mask, empty_mask_single, empty_mask, empty_mask, curve_json)
|
|
|
|
# Get dimensions from curve_json metadata, default to 512x512
|
|
meta = curve_data.get("meta", {})
|
|
width = meta.get("width", 512)
|
|
height = meta.get("height", 512)
|
|
|
|
# Get feather_amount from metadata, validate and clamp
|
|
feather_amount = meta.get("feather_amount", 0.0)
|
|
if not isinstance(feather_amount, (int, float)) or not np.isfinite(feather_amount):
|
|
feather_amount = 0.0
|
|
else:
|
|
feather_amount = float(np.clip(feather_amount, 0.0, 100.0))
|
|
|
|
# Get background_color from metadata, default to black
|
|
background_color = meta.get("background_color", [0.0, 0.0, 0.0])
|
|
if not isinstance(background_color, list) or len(background_color) != 3:
|
|
background_color = [0.0, 0.0, 0.0]
|
|
else:
|
|
background_color = [
|
|
float(np.clip(background_color[0], 0.0, 1.0)),
|
|
float(np.clip(background_color[1], 0.0, 1.0)),
|
|
float(np.clip(background_color[2], 0.0, 1.0)),
|
|
]
|
|
|
|
# Get frame_paths from metadata
|
|
frame_paths = meta.get("frame_paths", [])
|
|
|
|
# Generate masks for all frames in the sequence
|
|
masks = []
|
|
for i in range(frame_count):
|
|
current_frame = frame_index + i
|
|
mask = generate_single_mask(curve_data, current_frame, height, width, feather_amount)
|
|
masks.append(mask)
|
|
|
|
# Stack all masks into batch: list of (H, W) -> (B, H, W)
|
|
mask_batch = np.stack(masks, axis=0)
|
|
|
|
# Output 1: Convert to 3-channel RGB for preview compatibility: (B, H, W) -> (B, H, W, 3)
|
|
mask_rgb = np.stack([mask_batch, mask_batch, mask_batch], axis=-1)
|
|
mask_tensor = torch.from_numpy(mask_rgb).float()
|
|
|
|
# Output 2: MASK format (single channel, B, H, W)
|
|
mask_single_channel = torch.from_numpy(mask_batch).float()
|
|
|
|
# Outputs 3 and 4: Load images and apply mask overlay
|
|
masked_images = []
|
|
raw_images = []
|
|
|
|
for i in range(frame_count):
|
|
current_frame = frame_index + i
|
|
# Map 1-based frame_index to 0-based array index
|
|
frame_array_index = current_frame - 1
|
|
|
|
# Load image if available
|
|
if frame_paths and 0 <= frame_array_index < len(frame_paths):
|
|
frame_info = frame_paths[frame_array_index]
|
|
image = load_image_from_disk(
|
|
frame_info.get("filename", ""), frame_info.get("subfolder", "")
|
|
)
|
|
|
|
if image is not None:
|
|
# Resize mask to match image dimensions
|
|
img_h, img_w = image.shape[:2]
|
|
mask_resized = resize_mask_to_image(mask_batch[i], img_h, img_w)
|
|
|
|
# Apply mask overlay
|
|
masked_image = apply_mask_overlay(image, mask_resized, background_color)
|
|
masked_images.append(masked_image)
|
|
raw_images.append(image)
|
|
else:
|
|
# Image load failed, create empty tensors matching mask dimensions
|
|
empty_img = np.zeros((height, width, 3), dtype=np.float32)
|
|
masked_images.append(empty_img)
|
|
raw_images.append(empty_img)
|
|
else:
|
|
# No frame_paths or out of bounds, create empty tensors matching mask dimensions
|
|
empty_img = np.zeros((height, width, 3), dtype=np.float32)
|
|
masked_images.append(empty_img)
|
|
raw_images.append(empty_img)
|
|
|
|
# Stack images into batches
|
|
if masked_images:
|
|
masked_batch = np.stack(masked_images, axis=0)
|
|
masked_tensor = torch.from_numpy(masked_batch).float()
|
|
else:
|
|
masked_tensor = torch.zeros((frame_count, height, width, 3), dtype=torch.float32)
|
|
|
|
if raw_images:
|
|
raw_batch = np.stack(raw_images, axis=0)
|
|
raw_tensor = torch.from_numpy(raw_batch).float()
|
|
else:
|
|
raw_tensor = torch.zeros((frame_count, height, width, 3), dtype=torch.float32)
|
|
|
|
return (mask_tensor, mask_single_channel, masked_tensor, raw_tensor, curve_json)
|