Get the connected components (8-connectivity) of binary masks of shape (N, 1, H, W). Inputs: - mask: A binary mask tensor of shape (N, 1, H, W), where 1 is foreground and 0 is background. Outputs: - labels: A tensor of shape (N, 1, H, W) containing the connected co
(mask)
| 45 | |
| 46 | |
| 47 | def get_connected_components(mask): |
| 48 | """ |
| 49 | Get the connected components (8-connectivity) of binary masks of shape (N, 1, H, W). |
| 50 | |
| 51 | Inputs: |
| 52 | - mask: A binary mask tensor of shape (N, 1, H, W), where 1 is foreground and 0 is |
| 53 | background. |
| 54 | |
| 55 | Outputs: |
| 56 | - labels: A tensor of shape (N, 1, H, W) containing the connected component labels |
| 57 | for foreground pixels and 0 for background pixels. |
| 58 | - counts: A tensor of shape (N, 1, H, W) containing the area of the connected |
| 59 | components for foreground pixels and 0 for background pixels. |
| 60 | """ |
| 61 | from sam2_train import _C |
| 62 | |
| 63 | return _C.get_connected_componnets(mask.to(torch.uint8).contiguous()) |
| 64 | |
| 65 | |
| 66 | def mask_to_box(masks: torch.Tensor): |
no outgoing calls
no test coverage detected