| 64 | |
| 65 | |
| 66 | class SMPL_LOSS(nn.Module): |
| 67 | def __init__(self,cfg): |
| 68 | super(SMPL_LOSS, self).__init__() |
| 69 | self.weight_shape = cfg.loss.hyponet.weight_shape |
| 70 | self.weight_diff = cfg.loss.hyponet.weight_diff |
| 71 | self.criterion_smpl = nn.MSELoss() |
| 72 | self.weight_2d = cfg.loss.hyponet.weight_2d |
| 73 | self.num_joints = cfg.hyponet.num_joints |
| 74 | |
| 75 | def forward(self,output, gt, labels): |
| 76 | loss_beta = (output['pred_shape'] - labels['target_beta']) * labels['target_smpl_weight'] |
| 77 | loss = loss_beta.square().sum(dim=1).mean(dim=0) * self.weight_shape |
| 78 | |
| 79 | loss_2d = (output['joint_2d'] - gt['joint_2d']).square().view(-1,self.num_joints,2) * labels['joints_vis_29'][:,:,:2] |
| 80 | loss_2d = loss_2d.sum(dim=-1).sum(dim=-1).mean() * self.weight_2d |
| 81 | loss += loss_2d |
| 82 | |
| 83 | loss_joint = (output['noise_j'] - gt['noise_j']).square().view(-1, self.num_joints, 3) |
| 84 | try: |
| 85 | loss_joint = loss_joint * labels['joints_vis_29'] |
| 86 | except: |
| 87 | loss_joint = loss_joint * labels['joints_vis_29'].repeat(2, 1, 1) # 2*bs, 29, 3 |
| 88 | loss_joint = loss_joint.view(-1, self.num_joints*3).sum(dim=1).mean(dim=0)*self.weight_diff |
| 89 | loss_twist = (output['noise_t'] - gt['noise_t']).square().view(-1, 23, 2) |
| 90 | try: |
| 91 | loss_twist = loss_twist * labels['target_twist_weight'] |
| 92 | except: |
| 93 | loss_twist = loss_twist * labels['target_twist_weight'].repeat(2, 1, 1) # 2*bs, 23, 3 |
| 94 | loss_twist = loss_twist.view(-1, 23*2).sum(dim=1).mean(dim=0)*self.weight_diff |
| 95 | loss += (loss_twist+loss_joint) |
| 96 | |
| 97 | assert torch.isnan(loss).sum()==0 |
| 98 | return loss, {'loss_2d':loss_2d, 'loss_twist':loss_twist, 'loss_joint':loss_joint, 'loss_beta':loss-loss_joint-loss_twist-loss_2d} |
| 99 | |
| 100 | class DPO_SMPL_LOSS(nn.Module): |
| 101 | def __init__(self,cfg): |