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

Method loss_smpl_pose

models/aios/criterion_smplx.py:1141–1193  ·  view source on GitHub ↗
(self, outputs, targets, indices, idx, num_boxes,
                       data_batch, face_hand_kpt=False)

Source from the content-addressed store, hash-verified

1139 return losses
1140
1141 def loss_smpl_pose(self, outputs, targets, indices, idx, num_boxes,
1142 data_batch, face_hand_kpt=False):
1143 indices = indices[0]
1144 device = outputs['pred_logits'].device
1145
1146 pred_smpl_body_pose = outputs['pred_smpl_pose'][idx] # 22
1147 pred_smpl_lhand_pose = outputs['pred_smpl_lhand_pose'][idx] # 15
1148 pred_smpl_rhand_pose = outputs['pred_smpl_rhand_pose'][idx] # 15
1149 pred_smpl_jaw_pose = outputs['pred_smpl_jaw_pose'][idx]
1150
1151 pred_smplx_pose = torch.cat((pred_smpl_body_pose, pred_smpl_lhand_pose,
1152 pred_smpl_rhand_pose, pred_smpl_jaw_pose),
1153 dim=1)
1154
1155 targets_smpl_pose = torch.cat(
1156 [t[i] for t, (_, i) in zip(data_batch['smplx_pose'], indices)],
1157 dim=0)
1158 targets_smpl_pose = batch_rodrigues(targets_smpl_pose.view(
1159 -1, 3)).view(-1, 53, 3, 3)
1160 conf = torch.cat([
1161 t[i] for t, (_, i) in zip(data_batch['smplx_pose_valid'], indices)
1162 ], dim=0)
1163 body_pose_valid = conf[:, :22].sum(-1) > 0
1164 lhand_pose_valid = conf[:, 22:37].sum(-1) > 0
1165 rhand_pose_valid = conf[:, 37:52].sum(-1) > 0
1166 face_pose_valid = conf[:, 52].sum(-1) > 0
1167
1168 losses = {}
1169 loss_smpl_pose = \
1170 F.l1_loss(
1171 pred_smplx_pose,
1172 targets_smpl_pose,
1173 reduction='none'
1174 )
1175 loss_smpl_pose = loss_smpl_pose.sum([-1,-2]) * conf
1176
1177 if face_hand_kpt:
1178 losses = {
1179 'loss_smpl_pose_root': loss_smpl_pose[:, 0].sum() / (body_pose_valid.sum() + 1e-6),
1180 'loss_smpl_pose_body': loss_smpl_pose[:, 1:22].sum() / (body_pose_valid.sum() + 1e-6),
1181 'loss_smpl_pose_lhand': loss_smpl_pose[:, 22:37].sum() / (lhand_pose_valid.sum() + 1e-6),
1182 'loss_smpl_pose_rhand': loss_smpl_pose[:, 37:52].sum() / (rhand_pose_valid.sum() + 1e-6),
1183 'loss_smpl_pose_jaw': loss_smpl_pose[:, 52].sum() / (face_pose_valid.sum() + 1e-6),
1184 }
1185 else:
1186 losses = {
1187 'loss_smpl_pose_root': loss_smpl_pose[:, 0].sum() / (body_pose_valid.sum() + 1e-6),
1188 'loss_smpl_pose_body': loss_smpl_pose[:, 1:22].sum() / (body_pose_valid.sum() + 1e-6),
1189 'loss_smpl_pose_lhand': 0 * loss_smpl_pose[:, 22:37].sum()/(lhand_pose_valid.sum() + 1e-6),
1190 'loss_smpl_pose_rhand': 0 * loss_smpl_pose[:, 37:52].sum() / (rhand_pose_valid.sum() + 1e-6),
1191 'loss_smpl_pose_jaw': 0*loss_smpl_pose[:, 52].sum() / (face_pose_valid.sum() + 1e-6),
1192 }
1193 return losses
1194
1195 def loss_smpl_beta(self, outputs, targets, indices, idx, num_boxes,
1196 data_batch, face_hand_kpt=False):

Callers

nothing calls this directly

Calls 1

batch_rodriguesFunction · 0.90

Tested by

no test coverage detected