使用OpenCV快速分割蒙版并处理图像
(self, mask, image)
| 39 | return min_x, min_y |
| 40 | |
| 41 | def segment_mask(self, mask, image): |
| 42 | """使用OpenCV快速分割蒙版并处理图像""" |
| 43 | # 保存原始设备信息 |
| 44 | device = mask.device if isinstance(mask, torch.Tensor) else torch.device('cpu') |
| 45 | |
| 46 | # 确保mask是正确的形状并转换为numpy数组 |
| 47 | if isinstance(mask, torch.Tensor): |
| 48 | if len(mask.shape) == 2: |
| 49 | mask = mask.unsqueeze(0) |
| 50 | mask_np = (mask[0] * 255).cpu().numpy().astype(np.uint8) |
| 51 | else: |
| 52 | mask_np = (mask * 255).astype(np.uint8) |
| 53 | |
| 54 | # 使用OpenCV找到轮廓 |
| 55 | contours, hierarchy = cv2.findContours( |
| 56 | mask_np, |
| 57 | cv2.RETR_TREE, |
| 58 | cv2.CHAIN_APPROX_SIMPLE |
| 59 | ) |
| 60 | |
| 61 | mask_info = [] # 用于排序的信息列表 |
| 62 | |
| 63 | if hierarchy is not None and len(contours) > 0: |
| 64 | hierarchy = hierarchy[0] |
| 65 | contour_masks = {} |
| 66 | |
| 67 | # 创建每个轮廓的mask |
| 68 | for i, contour in enumerate(contours): |
| 69 | mask = np.zeros_like(mask_np) |
| 70 | cv2.drawContours(mask, [contour], -1, 255, -1) |
| 71 | contour_masks[i] = mask |
| 72 | |
| 73 | # 处理每个轮廓 |
| 74 | processed_indices = set() |
| 75 | |
| 76 | for i, (contour, h) in enumerate(zip(contours, hierarchy)): |
| 77 | if i in processed_indices: |
| 78 | continue |
| 79 | |
| 80 | current_mask = contour_masks[i].copy() |
| 81 | child_idx = h[2] |
| 82 | |
| 83 | if child_idx != -1: |
| 84 | while child_idx != -1: |
| 85 | current_mask = cv2.subtract(current_mask, contour_masks[child_idx]) |
| 86 | processed_indices.add(child_idx) |
| 87 | child_idx = hierarchy[child_idx][0] |
| 88 | |
| 89 | # 找到最左上角的点 |
| 90 | min_x, min_y = self.find_top_left_point(current_mask) |
| 91 | |
| 92 | # 转换为tensor |
| 93 | mask_tensor = torch.from_numpy(current_mask).float() / 255.0 |
| 94 | mask_tensor = mask_tensor.unsqueeze(0) |
| 95 | mask_tensor = mask_tensor.to(device) |
| 96 | |
| 97 | # 保存mask和排序信息 |
| 98 | mask_info.append((mask_tensor, min_x, min_y)) |
nothing calls this directly
no test coverage detected