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')