Source code for MRF.Online.Performances

import torch
import numpy as np


[docs]class Performances(): """ Class designed to handle the computations and the definition of the validation loss and errors. """ def __init__(self, validation_settings): self.losses = [] self.training_relative_errors = [] self.training_absolute_errors = [] for key, value in validation_settings.items(): setattr(self, key, value) self.gradients = []
[docs] def dico_save(self): """ Save the parameters of the instance of the class Performances.""" dic = (self.__dict__).copy() return dic
[docs] def loss_function(self,outputs,params,size): raise NotImplementedError("Must override loss_function.")
[docs] def compute_relative_errors(self, estimations_validation, parameters, size): raise NotImplementedError("Must override compute_relative_errors.")
[docs] def validation_step(self, estimations_validation): """ Compute the loss and the relative errors on the parameters on the validation dataset. The parameter 'estimation_validation' represents the estimation of the network for the parameters on the validation dataset. """ self.losses_validation.append((self.loss_function(estimations_validation,self.params_validation, self.validation_size, self.CRBs_validation)).cpu().detach().numpy()) self.validation_relative_errors.append((self.compute_relative_errors(estimations_validation, self.params_validation, self.validation_size)).cpu().detach().numpy()) self.validation_absolute_errors.append((self.compute_absolute_errors(estimations_validation, self.params_validation, self.validation_size)).cpu().detach().numpy()) if self.data_class.CRBrequired: self.validation_absolute_errors_over_CRBs.append((self.compute_absolute_errors_over_CRBs(estimations_validation, self.params_validation, self.validation_size, self.CRBs_validation)).cpu().detach().numpy())
[docs] def init_validation(self): """ Define the validation dataset. """ if self.validation: self.validation_absolute_errors = [] self.CRBs_validation = None self.losses_validation = [] self.validation_relative_errors = [] self.dico_validation = np.zeros((self.validation_size,666)) self.params_validation = np.zeros((self.validation_size,5)) if self.CRBrequired: self.validation_absolute_errors_over_CRBs = [] self.CRBs_validation = np.zeros((self.validation_size,len(self.traparas.params))) for i in range(self.validation_size): if self.CRBrequired: self.dico_validation[i,:], self.params_validation[i,:], self.CRBs_validation[i,:] = self.data_class.__getitem__(0) self.dico_validation = torch.tensor(self.dico_validation, dtype=torch.float, device='cpu') self.params_validation = torch.tensor(self.params_validation, dtype=torch.float, device='cpu') if self.CRBrequired: self.CRBs_validation = torch.tensor(self.CRBs_validation, dtype=torch.float, device='cpu')