(self)
| 126 | return all_output |
| 127 | |
| 128 | def backward_G(self): |
| 129 | loss_mot_rec = self.mse_criterion(self.fake_noise, self.real_noise).mean(dim=-1) |
| 130 | loss_mot_rec = (loss_mot_rec * self.src_mask).sum() / self.src_mask.sum() |
| 131 | self.loss_mot_rec = loss_mot_rec |
| 132 | loss_logs = OrderedDict({}) |
| 133 | loss_logs['loss_mot_rec'] = self.loss_mot_rec.item() |
| 134 | return loss_logs |
| 135 | |
| 136 | def update(self): |
| 137 | self.zero_grad([self.opt_encoder]) |