File size: 2,891 Bytes
bfbe8cd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 | import torch
import numpy as np
from PIL import Image, ImageOps
class MaskCropMaster:
"""Mask Crop Region replacement that ALWAYS outputs a square crop.
Based on WAS_Mask_Crop_Region from was-ns, with the key fix:
when the square extends beyond the image bounds, it REPOSITIONS
the crop instead of clamping dimensions (which produces a rectangle).
Compatible with Mask Paste Region (same crop_data format).
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"mask": ("MASK",),
"padding": ("INT", {"default": 24, "min": 0, "max": 4096, "step": 1}),
"region_type": (["dominant", "minority"],),
}
}
RETURN_TYPES = ("MASK", "CROP_DATA", "INT", "INT", "INT", "INT", "INT", "INT")
RETURN_NAMES = ("cropped_mask", "crop_data", "top_int", "left_int", "right_int", "bottom_int", "width_int", "height_int")
FUNCTION = "mask_crop_master"
CATEGORY = "WAS Suite/Image/Masking"
DESCRIPTION = "Mask Crop Region that always outputs a square, even at image edges."
def mask_crop_master(self, mask, padding=24, region_type="dominant"):
mask_np = mask.cpu().squeeze().numpy()
mask_pil = Image.fromarray(np.clip(255.0 * mask_np, 0, 255).astype(np.uint8))
img_w, img_h = mask_pil.size
bbox = mask_pil.getbbox()
if bbox is None:
empty = Image.new("L", (img_w, img_h), 0)
empty_tensor = torch.from_numpy(np.array(empty).astype(np.float32) / 255.0).unsqueeze(0).unsqueeze(1)
crop_data = ((img_w, img_h), (0, 0, 0, 0))
return (empty_tensor, crop_data, 0, 0, 0, 0, img_w, img_h)
bbox_x1, bbox_y1, bbox_x2, bbox_y2 = bbox
bbox_w = bbox_x2 - bbox_x1
bbox_h = bbox_y2 - bbox_y1
side = max(bbox_w, bbox_h) + 2 * padding
side = min(side, img_w, img_h)
cx = (bbox_x1 + bbox_x2) / 2.0
cy = (bbox_y1 + bbox_y2) / 2.0
crop_x = round(cx - side / 2.0)
crop_y = round(cy - side / 2.0)
if crop_x < 0:
crop_x = 0
if crop_y < 0:
crop_y = 0
if crop_x + side > img_w:
crop_x = img_w - side
if crop_y + side > img_h:
crop_y = img_h - side
crop_x2 = crop_x + side
crop_y2 = crop_y + side
cropped_mask = mask_pil.crop((crop_x, crop_y, crop_x2, crop_y2))
region_tensor = torch.from_numpy(
np.array(cropped_mask).astype(np.float32) / 255.0
).unsqueeze(0).unsqueeze(1)
crop_data = (cropped_mask.size, (crop_x, crop_y, crop_x2, crop_y2))
return (region_tensor, crop_data, crop_y, crop_x, crop_y2, crop_x2, side, side)
NODE_CLASS_MAPPINGS = {
"MaskCropMaster": MaskCropMaster,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"MaskCropMaster": "MASK CROP MASTER",
}
|