MCPcopy Create free account
hub / github.com/PolyU-ChenLab/UniPixel / sample_random_points_from_errors

Function sample_random_points_from_errors

sam2/modeling/sam2_utils.py:195–244  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

193
194
195def 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
247def sample_one_point_from_error_center(gt_masks, pred_masks, padding=True, positive_only=False):

Callers 1

get_next_pointFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected