MCPcopy Create free account
hub / github.com/UX-Decoder/Semantic-SAM / HungarianMatcher

Class HungarianMatcher

semantic_sam/modules/matcher.py:77–229  ·  view source on GitHub ↗

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,

Source from the content-addressed store, hash-verified

75
76
77class 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,

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected