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

Class SMPL_LOSS

ADHMR/lib/models/loss.py:66–98  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

64
65
66class 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
100class DPO_SMPL_LOSS(nn.Module):
101 def __init__(self,cfg):

Callers 1

get_modelFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected