from functools import partial import numpy as np import sklearn from utils.evaluation import * def get_eval_functions(model_configuration: dict): """ Returns the evaluation functions for the model :return: list of evaluation functions """ eval_functions = [ { "name": "mean_absolute_error_overall", "eval_fn": sklearn.metrics.mean_absolute_error, "accumulation_fn": np.mean, }, { "name": "mean_absolute_error_pre_ov", "eval_fn": pre_ov_error, "accumulation_fn": np.mean, }, { "name": "mean_absolute_error_after_ov", "eval_fn": after_ov_error, "accumulation_fn": np.mean, }, { "name": "mean_absolute_error_ov_in_days", "eval_fn": partial(ov_error, model_configuration=model_configuration), "accumulation_fn": np.mean, }, { "name": "mean_absolute_error_five_days_before_ov", "eval_fn": partial(day_relative_to_ov_error, day_relative_to_ov=-5, model_configuration=model_configuration), "accumulation_fn": np.mean, } ] return eval_functions def get_fertility_based_eval_functions(model_configuration: dict): eval_functions = [ { "name": "mae_ov_over", "eval_fn": sklearn.metrics.mean_absolute_error, "accumulation_fn": np.mean, "input_index": 1 }, { "name": "mae_ov_over_before_ov", "eval_fn": partial(get_ov_over_pre_ov_error, error_fn=sklearn.metrics.mean_absolute_error), "accumulation_fn": np.mean, "input_index": 1 }, { "name": "mae_ov_over_after_ov", "eval_fn": partial(get_ov_over_post_ov_error, error_fn=sklearn.metrics.mean_absolute_error), "accumulation_fn": np.mean, "input_index": 1 }, { "name": "mse_ov_over", "eval_fn": sklearn.metrics.mean_squared_error, "accumulation_fn": np.mean, "input_index": 1 }, { "name": "mse_ov_over_before_ov", "eval_fn": partial(get_ov_over_pre_ov_error, error_fn=sklearn.metrics.mean_squared_error), "accumulation_fn": np.mean, "input_index": 1 }, { "name": "mse_ov_over_after_ov", "eval_fn": partial(get_ov_over_post_ov_error, error_fn=sklearn.metrics.mean_squared_error), "accumulation_fn": np.mean, "input_index": 1 }, { "name": "mae_fertility", "eval_fn": sklearn.metrics.mean_absolute_error, "accumulation_fn": np.mean, "input_index": 0 }, { "name": "mae_during_fertility", "eval_fn": partial(get_during_fertility_error, error_fn=sklearn.metrics.mean_absolute_error), "accumulation_fn": np.mean, "input_index": 0 }, { "name": "mae_non_fertility", "eval_fn": partial(get_non_fertility_error, error_fn=sklearn.metrics.mean_absolute_error), "accumulation_fn": np.mean, "input_index": 0 }, { "name": "mse_fertility", "eval_fn": sklearn.metrics.mean_squared_error, "accumulation_fn": np.mean, "input_index": 0 }, { "name": "mse_during_fertility", "eval_fn": partial(get_during_fertility_error, error_fn=sklearn.metrics.mean_squared_error), "accumulation_fn": np.mean, "input_index": 0 }, { "name": "mse_non_fertility", "eval_fn": partial(get_non_fertility_error, error_fn=sklearn.metrics.mean_squared_error), "accumulation_fn": np.mean, "input_index": 0 }, ] return eval_functions