MCPcopy Create free account
hub / github.com/Meshcapade/difflocks / forward

Method forward

losses/loss.py:14–40  ·  view source on GitHub ↗
(self, phase, gt_dict, pred_dict, latent_dict, hyperparams)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 4

compute_loss_l1Function · 0.90
compute_loss_dir_l1Function · 0.90
compute_loss_curv_l1Function · 0.90
compute_loss_klFunction · 0.90

Tested by

no test coverage detected