MCPcopy Create free account
hub / github.com/Francis-Rings/FlashPortrait / custom_mse_loss

Function custom_mse_loss

train_portrait.py:1411–1428  ·  view source on GitHub ↗
(noise_pred, target, weighting=None, threshold=50)

Source from the content-addressed store, hash-verified

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)

Callers 1

mainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected