(self, mask)
| 796 | CATEGORY = "Masquerade Nodes" |
| 797 | |
| 798 | def separate(self, mask): |
| 799 | mask = tensor2mask(mask) |
| 800 | |
| 801 | thresholded = torch.gt(mask,0).unsqueeze(1) |
| 802 | B, H, W = mask.shape |
| 803 | components = torch.arange(B * H * W, device=mask.device, dtype=mask.dtype).reshape(B, 1, H, W) + 1 |
| 804 | components[~thresholded] = 0 |
| 805 | |
| 806 | while True: |
| 807 | previous_components = components |
| 808 | components = torch.nn.functional.max_pool2d(components, kernel_size=3, stride=1, padding=1) |
| 809 | components[~thresholded] = 0 |
| 810 | if torch.equal(previous_components, components): |
| 811 | break |
| 812 | |
| 813 | components = components.reshape(B, H, W) |
| 814 | segments = torch.unique(components) |
| 815 | result = torch.zeros([len(segments) - 1, H, W]) |
| 816 | index = 0 |
| 817 | mapping = torch.zeros([len(segments) - 1], device=mask.device, dtype=torch.int) |
| 818 | for i in range(len(segments)): |
| 819 | segment = segments[i].item() |
| 820 | if segment == 0: |
| 821 | continue |
| 822 | image_index = int((segment - 1) // (H * W)) |
| 823 | segment_mask = (components[image_index,:,:] == segment) |
| 824 | result[index][segment_mask] = mask[image_index][segment_mask] |
| 825 | mapping[index] = image_index |
| 826 | index += 1 |
| 827 | |
| 828 | return (result,mapping,) |
| 829 | |
| 830 | |
| 831 | class PasteByMask: |
nothing calls this directly
no test coverage detected