Sample `num_pt` random points (along with their labels) independently from the error regions. Inputs: - gt_masks: [B, 1, H_im, W_im] masks, dtype=torch.bool - pred_masks: [B, 1, H_im, W_im] masks, dtype=torch.bool or None - num_pt: int, number of points to sample independently
(gt_masks, pred_masks, num_pt=1, positive_only=False)
| 193 | |
| 194 | |
| 195 | def sample_random_points_from_errors(gt_masks, pred_masks, num_pt=1, positive_only=False): |
| 196 | """ |
| 197 | Sample `num_pt` random points (along with their labels) independently from the error regions. |
| 198 | |
| 199 | Inputs: |
| 200 | - gt_masks: [B, 1, H_im, W_im] masks, dtype=torch.bool |
| 201 | - pred_masks: [B, 1, H_im, W_im] masks, dtype=torch.bool or None |
| 202 | - num_pt: int, number of points to sample independently for each of the B error maps |
| 203 | |
| 204 | Outputs: |
| 205 | - points: [B, num_pt, 2], dtype=torch.float, contains (x, y) coordinates of each sampled point |
| 206 | - labels: [B, num_pt], dtype=torch.int32, where 1 means positive clicks and 0 means |
| 207 | negative clicks |
| 208 | """ |
| 209 | if pred_masks is None: # if pred_masks is not provided, treat it as empty |
| 210 | pred_masks = torch.zeros_like(gt_masks) |
| 211 | assert gt_masks.dtype == torch.bool and gt_masks.size(1) == 1 |
| 212 | assert pred_masks.dtype == torch.bool and pred_masks.shape == gt_masks.shape |
| 213 | assert num_pt >= 0 |
| 214 | |
| 215 | B, _, H_im, W_im = gt_masks.shape |
| 216 | device = gt_masks.device |
| 217 | |
| 218 | # false positive region, a new point sampled in this region should have |
| 219 | # negative label to correct the FP error |
| 220 | fp_masks = ~gt_masks & pred_masks |
| 221 | # false negative region, a new point sampled in this region should have |
| 222 | # positive label to correct the FN error |
| 223 | fn_masks = gt_masks & ~pred_masks |
| 224 | # whether the prediction completely match the ground-truth on each mask |
| 225 | all_correct = torch.all((gt_masks == pred_masks).flatten(2), dim=2) |
| 226 | all_correct = all_correct[..., None, None] |
| 227 | |
| 228 | # channel 0 is FP map, while channel 1 is FN map |
| 229 | pts_noise = torch.rand(B, num_pt, H_im, W_im, 2, device=device) |
| 230 | # sample a negative new click from FP region or a positive new click |
| 231 | # from FN region, depend on where the maximum falls, |
| 232 | # and in case the predictions are all correct (no FP or FN), we just |
| 233 | # sample a negative click from the background region |
| 234 | pts_noise[..., 0] *= fp_masks | (all_correct & ~gt_masks) |
| 235 | if positive_only: |
| 236 | pts_noise[..., 0] = -1 |
| 237 | pts_noise[..., 1] *= fn_masks |
| 238 | pts_idx = pts_noise.flatten(2).argmax(dim=2) |
| 239 | labels = (pts_idx % 2).to(torch.int32) |
| 240 | pts_idx = pts_idx // 2 |
| 241 | pts_x = pts_idx % W_im |
| 242 | pts_y = pts_idx // W_im |
| 243 | points = torch.stack([pts_x, pts_y], dim=2).to(torch.float) |
| 244 | return points, labels |
| 245 | |
| 246 | |
| 247 | def sample_one_point_from_error_center(gt_masks, pred_masks, padding=True, positive_only=False): |