(self, outputs, targets, indices, idx, num_boxes,
data_batch, face_hand_kpt=False)
| 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): |
nothing calls this directly
no test coverage detected