(
self,
cls_scores,
box_preds,
instance_id=None,
quality=None,
output_idx=-1,
)
| 34 | self.sorted = sorted |
| 35 | |
| 36 | def decode( |
| 37 | self, |
| 38 | cls_scores, |
| 39 | box_preds, |
| 40 | instance_id=None, |
| 41 | quality=None, |
| 42 | output_idx=-1, |
| 43 | ): |
| 44 | squeeze_cls = instance_id is not None |
| 45 | |
| 46 | cls_scores = cls_scores[output_idx].sigmoid() |
| 47 | |
| 48 | if squeeze_cls: |
| 49 | cls_scores, cls_ids = cls_scores.max(dim=-1) |
| 50 | cls_scores = cls_scores.unsqueeze(dim=-1) |
| 51 | |
| 52 | box_preds = box_preds[output_idx] |
| 53 | bs, num_pred, num_cls = cls_scores.shape |
| 54 | cls_scores, indices = cls_scores.flatten(start_dim=1).topk( |
| 55 | self.num_output, dim=1, sorted=self.sorted |
| 56 | ) |
| 57 | if not squeeze_cls: |
| 58 | cls_ids = indices % num_cls |
| 59 | if self.score_threshold is not None: |
| 60 | mask = cls_scores >= self.score_threshold |
| 61 | |
| 62 | if quality[output_idx] is None: |
| 63 | quality = None |
| 64 | if quality is not None: |
| 65 | centerness = quality[output_idx][..., CNS] |
| 66 | centerness = torch.gather(centerness, 1, indices // num_cls) |
| 67 | cls_scores_origin = cls_scores.clone() |
| 68 | cls_scores *= centerness.sigmoid() |
| 69 | cls_scores, idx = torch.sort(cls_scores, dim=1, descending=True) |
| 70 | if not squeeze_cls: |
| 71 | cls_ids = torch.gather(cls_ids, 1, idx) |
| 72 | if self.score_threshold is not None: |
| 73 | mask = torch.gather(mask, 1, idx) |
| 74 | indices = torch.gather(indices, 1, idx) |
| 75 | |
| 76 | output = [] |
| 77 | for i in range(bs): |
| 78 | category_ids = cls_ids[i] |
| 79 | if squeeze_cls: |
| 80 | category_ids = category_ids[indices[i]] |
| 81 | scores = cls_scores[i] |
| 82 | box = box_preds[i, indices[i] // num_cls] |
| 83 | if self.score_threshold is not None: |
| 84 | category_ids = category_ids[mask[i]] |
| 85 | scores = scores[mask[i]] |
| 86 | box = box[mask[i]] |
| 87 | if quality is not None: |
| 88 | scores_origin = cls_scores_origin[i] |
| 89 | if self.score_threshold is not None: |
| 90 | scores_origin = scores_origin[mask[i]] |
| 91 | |
| 92 | box = decode_box(box) |
| 93 | output.append( |
no test coverage detected