(self, loss, outputs, targets, data_batch, indices, num_boxes,
**kwargs)
| 812 | return batch_idx, tgt_idx |
| 813 | |
| 814 | def get_loss(self, loss, outputs, targets, data_batch, indices, num_boxes, |
| 815 | **kwargs): |
| 816 | loss_map = { |
| 817 | 'smpl_pose': self.loss_smpl_pose, |
| 818 | 'smpl_beta': self.loss_smpl_beta, |
| 819 | 'smpl_expr': self.loss_smpl_expr, |
| 820 | 'smpl_kp2d': self.loss_smpl_kp2d, |
| 821 | 'smpl_kp3d_ra': self.loss_smpl_kp3d_ra, |
| 822 | 'smpl_kp3d': self.loss_smpl_kp3d, |
| 823 | 'labels': self.loss_labels, |
| 824 | 'cardinality': self.loss_cardinality, |
| 825 | 'keypoints': self.loss_keypoints, |
| 826 | 'boxes': self.loss_boxes, |
| 827 | 'dn_label': self.loss_dn_labels, |
| 828 | 'dn_bbox': self.loss_dn_boxes, |
| 829 | 'matching': self.loss_matching_cost, |
| 830 | } |
| 831 | |
| 832 | idx = self._get_src_permutation_idx(indices[0]) |
| 833 | # pdb.set_trace() |
| 834 | assert loss in loss_map, f'do you really want to compute {loss} loss?' |
| 835 | return loss_map[loss](outputs, targets, indices, idx, num_boxes, |
| 836 | data_batch, **kwargs) |
| 837 | |
| 838 | def prep_for_dn2(self, mask_dict): |
| 839 | known_bboxs = mask_dict['known_bboxs'] |
no test coverage detected