(self, loss, outputs, targets, data_batch, indices, num_boxes,
**kwargs)
| 1667 | return batch_idx, tgt_idx |
| 1668 | |
| 1669 | def get_loss(self, loss, outputs, targets, data_batch, indices, num_boxes, |
| 1670 | **kwargs): |
| 1671 | loss_map = { |
| 1672 | 'smpl_pose': self.loss_smpl_pose, |
| 1673 | 'smpl_beta': self.loss_smpl_beta, |
| 1674 | 'smpl_expr': self.loss_smpl_expr, |
| 1675 | 'smpl_kp2d': self.loss_smpl_kp2d, |
| 1676 | 'smpl_kp3d_ra': self.loss_smpl_kp3d_ra, |
| 1677 | 'smpl_kp3d': self.loss_smpl_kp3d, |
| 1678 | 'labels': self.loss_labels, |
| 1679 | 'cardinality': self.loss_cardinality, |
| 1680 | 'boxes': self.loss_boxes, |
| 1681 | 'dn_label': self.loss_dn_labels, |
| 1682 | 'dn_bbox': self.loss_dn_boxes, |
| 1683 | 'matching': self.loss_matching_cost, |
| 1684 | } |
| 1685 | |
| 1686 | idx = self._get_src_permutation_idx(indices[0]) |
| 1687 | # pdb.set_trace() |
| 1688 | assert loss in loss_map, f'do you really want to compute {loss} loss?' |
| 1689 | return loss_map[loss](outputs, targets, indices, idx, num_boxes, |
| 1690 | data_batch, **kwargs) |
| 1691 | |
| 1692 | def prep_for_dn2(self, mask_dict): |
| 1693 | known_bboxs = mask_dict['known_bboxs'] |
no test coverage detected