| 1409 | ) |
| 1410 | |
| 1411 | def custom_mse_loss(noise_pred, target, weighting=None, threshold=50): |
| 1412 | noise_pred = noise_pred.float() |
| 1413 | target = target.float() |
| 1414 | diff = noise_pred - target |
| 1415 | mse_loss = F.mse_loss(noise_pred, target, reduction='none') |
| 1416 | |
| 1417 | mask_loss_flag = torch.rand(1).item() |
| 1418 | if mask_loss_flag >= 0.5 and mask_loss_flag < 0.7: |
| 1419 | mse_loss = mse_loss * tgt_face_masks |
| 1420 | elif mask_loss_flag >= 0.7: |
| 1421 | mse_loss = mse_loss * tgt_lip_masks |
| 1422 | else: |
| 1423 | mse_loss = mse_loss * (1 + tgt_face_masks + tgt_lip_masks) |
| 1424 | |
| 1425 | if weighting is not None: |
| 1426 | mse_loss = mse_loss * weighting |
| 1427 | final_loss = mse_loss.mean() |
| 1428 | return final_loss |
| 1429 | |
| 1430 | tgt_face_masks = F.interpolate(tgt_face_masks, size=(target.size()[-3], target.size()[-2], target.size()[-1]), mode='trilinear', align_corners=False) |
| 1431 | tgt_lip_masks = F.interpolate(tgt_lip_masks, size=(target.size()[-3], target.size()[-2], target.size()[-1]), mode='trilinear', align_corners=False) |