| 23 | return out.sum() |
| 24 | |
| 25 | class SCORE_LOSS(nn.Module): |
| 26 | def __init__(self, cfg): |
| 27 | super(SCORE_LOSS, self).__init__() |
| 28 | self.eps = cfg.loss.scorenet.eps |
| 29 | self.weight_pve = cfg.loss.scorenet.weight_pve |
| 30 | self.weight_mcam = cfg.loss.scorenet.weight_mcam |
| 31 | self.weight_2d = cfg.loss.scorenet.weight_2d |
| 32 | self.num_joints =cfg.hyponet.num_joints |
| 33 | self.num_twists = cfg.hyponet.num_twists |
| 34 | self.sigma = cfg.loss.scorenet.sigma |
| 35 | self.loss_func = functools.partial(rankloss, sigma=self.sigma) |
| 36 | def forward(self, score, output, gt, mask): |
| 37 | bs = gt['mesh'].shape[0] |
| 38 | assert output['pred_vertices'].shape[0] % bs==0 |
| 39 | multi_n = output['pred_vertices'].shape[0] // bs |
| 40 | pred_mesh = output['pred_vertices'].view(bs,multi_n,-1,3) |
| 41 | gt_mesh = gt['mesh'].unsqueeze(1) |
| 42 | pve = torch.sqrt(torch.sum((pred_mesh - gt_mesh) ** 2, dim=3)) |
| 43 | pve = torch.mean(pve, dim=2) * 1000 |
| 44 | |
| 45 | joint_cam_pred = output['pred_xyz_jts_17'].reshape(bs, multi_n, -1, 3) |
| 46 | joint_cam_pred = joint_cam_pred - joint_cam_pred[:,:,0].reshape(bs,multi_n,1,3) |
| 47 | joint_cam_pred = joint_cam_pred * 2000 |
| 48 | joint_cam_gt = gt['joint_cam'] - gt['joint_cam'][:,0].unsqueeze(1) |
| 49 | joint_cam_gt = joint_cam_gt.unsqueeze(1) |
| 50 | mpjpe_cam = torch.sqrt(torch.sum((joint_cam_gt - joint_cam_pred) ** 2, dim=3)) |
| 51 | mpjpe_cam = torch.mean(mpjpe_cam,dim=2) |
| 52 | |
| 53 | score_t = score.view(bs, multi_n) |
| 54 | p_pve = self.loss_func(score_t.clone(), pve.clone(), mask.clone()) * self.weight_pve |
| 55 | p_mpjpe_cam = self.loss_func(score_t.clone(), mpjpe_cam.clone(), mask.clone()) * self.weight_mcam |
| 56 | |
| 57 | |
| 58 | pred_2d = output['pred_2d'].view(-1, self.num_joints, 2) |
| 59 | gt_2d = gt['pred_2d'].view(-1, self.num_joints, 2) |
| 60 | loss_2d = (pred_2d - gt_2d).square().view(-1, self.num_joints, 2) * gt['mask_2d'] |
| 61 | loss_2d = loss_2d.view(-1,self.num_joints*2).sum(dim=1).mean(dim=0) * self.weight_2d |
| 62 | |
| 63 | return p_pve, p_mpjpe_cam, loss_2d |
| 64 | |
| 65 | |
| 66 | class SMPL_LOSS(nn.Module): |