123 lines
4.0 KiB
Python
123 lines
4.0 KiB
Python
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
|