Files
2026-01-28 19:33:35 +13:00

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)