MCPcopy Create free account
hub / github.com/adept-thu/DIVER / decode

Method decode

mmdet3d_plugin/models/detection3d/decoder.py:36–107  ·  view source on GitHub ↗
(
        self,
        cls_scores,
        box_preds,
        instance_id=None,
        quality=None,
        output_idx=-1,
    )

Source from the content-addressed store, hash-verified

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(

Callers 1

post_processMethod · 0.45

Calls 1

decode_boxFunction · 0.85

Tested by

no test coverage detected