Files
rui40000-RUI-Nodes/mask_preview_node.py
T
2026-03-27 10:54:04 +08:00

173 lines
4.9 KiB
Python

import torch
import torch.nn.functional as F
import numpy as np
from PIL import Image
import folder_paths
import os
import json
class RuiMaskPreview:
"""
遮罩预览节点:将遮罩以半透明彩色形式叠加到图像上进行可视化预览
"""
COLOR_MAP = {
"red": (1.0, 0.0, 0.0),
"green": (0.0, 1.0, 0.0),
"blue": (0.0, 0.0, 1.0),
"yellow": (1.0, 1.0, 0.0),
"cyan": (0.0, 1.0, 1.0),
"magenta": (1.0, 0.0, 1.0),
"white": (1.0, 1.0, 1.0),
}
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"mask": ("MASK",),
"mask_color": (
["红色 / Red", "绿色 / Green", "蓝色 / Blue", "黄色 / Yellow",
"青色 / Cyan", "品红 / Magenta", "白色 / White"],
{"default": "红色 / Red"}
),
},
"optional": {
"opacity": ("FLOAT", {
"default": 0.5,
"min": 0.0,
"max": 1.0,
"step": 0.05
}),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("图像 / Image",)
FUNCTION = "preview_mask"
CATEGORY = "Rui-Node🐶/遮罩处理🎭"
OUTPUT_NODE = True
def preview_mask(self, image, mask, mask_color, opacity=0.5):
"""
将遮罩以半透明彩色形式叠加到图像上
参数:
image: 输入图像张量 (N, H, W, C)
mask: 输入遮罩张量 (N, H, W) 或 (H, W)
mask_color: 遮罩显示颜色(中英双语字符串)
opacity: 不透明度 (0.0-1.0)
返回:
合成后的图像张量和预览信息
"""
color_mapping = {
"红色 / Red": "red",
"绿色 / Green": "green",
"蓝色 / Blue": "blue",
"黄色 / Yellow": "yellow",
"青色 / Cyan": "cyan",
"品红 / Magenta": "magenta",
"白色 / White": "white",
}
color_key = color_mapping.get(mask_color, "red")
color_rgb = self.COLOR_MAP[color_key]
batch, height, width, channels = image.shape
if mask.dim() == 2:
mask = mask.unsqueeze(0)
mask_batch, mask_height, mask_width = mask.shape
if mask_height != height or mask_width != width:
mask = mask.unsqueeze(1)
mask = F.interpolate(
mask,
size=(height, width),
mode='bilinear',
align_corners=False
)
mask = mask.squeeze(1)
if mask_batch == 1 and batch > 1:
mask = mask.repeat(batch, 1, 1)
elif mask_batch != batch:
min_batch = min(mask_batch, batch)
mask = mask[:min_batch]
image = image[:min_batch]
batch = min_batch
print(f"警告: 遮罩批次数({mask_batch})与图像批次数({batch})不匹配,已截取为{min_batch}")
mask = torch.clamp(mask, 0.0, 1.0)
if channels > 3:
image = image[:, :, :, :3]
mask_expanded = mask.unsqueeze(-1)
color_tensor = torch.tensor(
color_rgb,
dtype=image.dtype,
device=image.device
).view(1, 1, 1, 3)
color_layer = mask_expanded * color_tensor
alpha = mask_expanded * opacity
output = image * (1 - alpha) + color_layer * opacity
output = torch.clamp(output, 0.0, 1.0)
results = self.save_images(output)
return {
"ui": {"images": results},
"result": (output,)
}
def save_images(self, images):
"""
保存图像供预览使用
参数:
images: 图像张量 (N, H, W, C)
返回:
包含图像信息的列表
"""
results = []
output_dir = folder_paths.get_temp_directory()
for i, image_tensor in enumerate(images):
img_np = image_tensor.cpu().numpy()
img_np = np.clip(img_np * 255, 0, 255).astype(np.uint8)
img_pil = Image.fromarray(img_np, 'RGB')
filename = f"mask_preview_{i:05d}.png"
filepath = os.path.join(output_dir, filename)
img_pil.save(filepath, compress_level=4)
results.append({
"filename": filename,
"subfolder": "",
"type": "temp"
})
return results
NODE_CLASS_MAPPINGS = {
"RuiMaskPreview": RuiMaskPreview
}
NODE_DISPLAY_NAME_MAPPINGS = {
"RuiMaskPreview": "遮罩预览 / Mask Preview"
}