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",
}