MCPcopy Create free account
hub / github.com/MotrixLab/AiOS / get_loss

Method get_loss

models/aios/criterion_smplx.py:814–836  ·  view source on GitHub ↗
(self, loss, outputs, targets, data_batch, indices, num_boxes,
                 **kwargs)

Source from the content-addressed store, hash-verified

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']

Callers 1

forwardMethod · 0.95

Calls 1

Tested by

no test coverage detected