This class computes an assignment between the targets and the predictions of the network For efficiency reasons, the targets don't include the no_object. Because of this, in general, there are more predictions than targets. In this case, we do a 1-to-1 matching of the best predictions,
| 75 | |
| 76 | |
| 77 | class HungarianMatcher(nn.Module): |
| 78 | """This class computes an assignment between the targets and the predictions of the network |
| 79 | |
| 80 | For efficiency reasons, the targets don't include the no_object. Because of this, in general, |
| 81 | there are more predictions than targets. In this case, we do a 1-to-1 matching of the best predictions, |
| 82 | while the others are un-matched (and thus treated as non-objects). |
| 83 | """ |
| 84 | |
| 85 | def __init__(self, cost_class: float = 1, cost_mask: float = 1, cost_dice: float = 1, num_points: int = 0, |
| 86 | cost_box: float = 0, cost_giou: float = 0, panoptic_on: bool = False): |
| 87 | """Creates the matcher |
| 88 | |
| 89 | Params: |
| 90 | cost_class: This is the relative weight of the classification error in the matching cost |
| 91 | cost_mask: This is the relative weight of the focal loss of the binary mask in the matching cost |
| 92 | cost_dice: This is the relative weight of the dice loss of the binary mask in the matching cost |
| 93 | """ |
| 94 | super().__init__() |
| 95 | self.cost_class = cost_class |
| 96 | self.cost_mask = cost_mask |
| 97 | self.cost_dice = cost_dice |
| 98 | self.cost_box = cost_box |
| 99 | self.cost_giou = cost_giou |
| 100 | |
| 101 | self.panoptic_on = panoptic_on |
| 102 | |
| 103 | assert cost_class != 0 or cost_mask != 0 or cost_dice != 0, "all costs cant be 0" |
| 104 | |
| 105 | self.num_points = num_points |
| 106 | |
| 107 | @torch.no_grad() |
| 108 | def memory_efficient_forward(self, outputs, targets, cost=["cls", "box", "mask"]): |
| 109 | """More memory-friendly matching. Change cost to compute only certain loss in matching""" |
| 110 | bs, num_queries = outputs["pred_logits"].shape[:2] |
| 111 | |
| 112 | indices = [] |
| 113 | |
| 114 | # Iterate through batch size |
| 115 | for b in range(bs): |
| 116 | out_bbox = outputs["pred_boxes"][b] |
| 117 | if 'box' in cost: |
| 118 | tgt_bbox=targets[b]["boxes"] |
| 119 | cost_bbox = torch.cdist(out_bbox, tgt_bbox, p=1) |
| 120 | cost_giou = -generalized_box_iou(box_cxcywh_to_xyxy(out_bbox), box_cxcywh_to_xyxy(tgt_bbox)) |
| 121 | else: |
| 122 | cost_bbox = torch.tensor(0).to(out_bbox) |
| 123 | cost_giou = torch.tensor(0).to(out_bbox) |
| 124 | |
| 125 | out_prob = outputs["pred_logits"][b].sigmoid() # [num_queries, num_classes] |
| 126 | tgt_ids = targets[b]["labels"] |
| 127 | # focal loss |
| 128 | alpha = 0.25 |
| 129 | gamma = 2.0 |
| 130 | neg_cost_class = (1 - alpha) * (out_prob ** gamma) * (-(1 - out_prob + 1e-6).log()) |
| 131 | pos_cost_class = alpha * ((1 - out_prob) ** gamma) * (-(out_prob + 1e-6).log()) |
| 132 | cost_class = pos_cost_class[:, tgt_ids] - neg_cost_class[:, tgt_ids] |
| 133 | |
| 134 | # Compute the classification cost. Contrary to the loss, we don't use the NLL, |
nothing calls this directly
no outgoing calls
no test coverage detected