(self, phase, gt_dict, pred_dict, latent_dict, hyperparams)
| 12 | self.eps = eps |
| 13 | |
| 14 | def forward(self, phase, gt_dict, pred_dict, latent_dict, hyperparams): |
| 15 | gt_strands= gt_dict["strand_positions"] |
| 16 | pred_hair_strands= pred_dict["strand_positions"] |
| 17 | |
| 18 | # w_kl=map_range_val(phase.epoch_nr, hyperparams.loss_kl_anneal_epoch_nr_start, hyperparams.loss_kl_anneal_epoch_nr_finish, 0.0, hyperparams.loss_kl_weight) |
| 19 | |
| 20 | |
| 21 | loss_dict = {} |
| 22 | |
| 23 | #we always predict the xyz loss just because it's fast |
| 24 | # loss_l2 = compute_loss_l2(gt_strands, pred_hair_strands) |
| 25 | loss_pos = compute_loss_l1(gt_strands, pred_hair_strands) |
| 26 | loss_dir = compute_loss_dir_l1(gt_strands, pred_hair_strands) |
| 27 | loss_curv = compute_loss_curv_l1(gt_strands, pred_hair_strands) |
| 28 | loss_kl = 0.0 |
| 29 | if "z_logstd" in latent_dict: |
| 30 | loss_kl = compute_loss_kl(latent_dict["z_mean"], latent_dict["z_logstd"]) |
| 31 | loss = loss_pos*hyperparams.loss_pos_weight + loss_dir*hyperparams.loss_dir_weight + loss_curv*hyperparams.loss_curv_weight + loss_kl*hyperparams.loss_kl_weight |
| 32 | # loss = loss_pos*hyperparams.loss_pos_weight + loss_dir*hyperparams.loss_dir_weight + loss_curv*hyperparams.loss_curv_weight + loss_kl*w_kl |
| 33 | loss_dict['loss'] = loss |
| 34 | loss_dict['loss_pos'] = loss_pos |
| 35 | loss_dict['loss_dir'] = loss_dir |
| 36 | loss_dict['loss_curv'] = loss_curv |
| 37 | loss_dict['loss_kl'] = loss_kl |
| 38 | |
| 39 | |
| 40 | return loss_dict |
| 41 | |
| 42 | |
| 43 |
nothing calls this directly
no test coverage detected