(
self,
cls_scores,
box_preds,
instance_id=None,
quality=None,
motion_output=None,
output_idx=-1,
)
| 16 | super(SparseBox3DMotionDecoder, self).__init__() |
| 17 | |
| 18 | def decode( |
| 19 | self, |
| 20 | cls_scores, |
| 21 | box_preds, |
| 22 | instance_id=None, |
| 23 | quality=None, |
| 24 | motion_output=None, |
| 25 | output_idx=-1, |
| 26 | ): |
| 27 | squeeze_cls = instance_id is not None |
| 28 | |
| 29 | cls_scores = cls_scores[output_idx].sigmoid() |
| 30 | |
| 31 | if squeeze_cls: |
| 32 | cls_scores, cls_ids = cls_scores.max(dim=-1) |
| 33 | cls_scores = cls_scores.unsqueeze(dim=-1) |
| 34 | |
| 35 | box_preds = box_preds[output_idx] |
| 36 | bs, num_pred, num_cls = cls_scores.shape |
| 37 | cls_scores, indices = cls_scores.flatten(start_dim=1).topk( |
| 38 | self.num_output, dim=1, sorted=self.sorted |
| 39 | ) |
| 40 | if not squeeze_cls: |
| 41 | cls_ids = indices % num_cls |
| 42 | if self.score_threshold is not None: |
| 43 | mask = cls_scores >= self.score_threshold |
| 44 | |
| 45 | if quality[output_idx] is None: |
| 46 | quality = None |
| 47 | if quality is not None: |
| 48 | centerness = quality[output_idx][..., CNS] |
| 49 | centerness = torch.gather(centerness, 1, indices // num_cls) |
| 50 | cls_scores_origin = cls_scores.clone() |
| 51 | cls_scores *= centerness.sigmoid() |
| 52 | cls_scores, idx = torch.sort(cls_scores, dim=1, descending=True) |
| 53 | if not squeeze_cls: |
| 54 | cls_ids = torch.gather(cls_ids, 1, idx) |
| 55 | if self.score_threshold is not None: |
| 56 | mask = torch.gather(mask, 1, idx) |
| 57 | indices = torch.gather(indices, 1, idx) |
| 58 | |
| 59 | output = [] |
| 60 | anchor_queue = motion_output["anchor_queue"] |
| 61 | anchor_queue = torch.stack(anchor_queue, dim=2) |
| 62 | period = motion_output["period"] |
| 63 | |
| 64 | for i in range(bs): |
| 65 | category_ids = cls_ids[i] |
| 66 | if squeeze_cls: |
| 67 | category_ids = category_ids[indices[i]] |
| 68 | scores = cls_scores[i] |
| 69 | box = box_preds[i, indices[i] // num_cls] |
| 70 | if self.score_threshold is not None: |
| 71 | category_ids = category_ids[mask[i]] |
| 72 | scores = scores[mask[i]] |
| 73 | box = box[mask[i]] |
| 74 | if quality is not None: |
| 75 | scores_origin = cls_scores_origin[i] |
no test coverage detected