MCPcopy Create free account
hub / github.com/MotrixLab/ADHMR / SCORE_LOSS

Class SCORE_LOSS

ADHMR/lib/models/loss.py:25–63  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

23 return out.sum()
24
25class 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
66class SMPL_LOSS(nn.Module):

Callers 1

get_model_scoreFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected