Files
temperature-based-fertility…/code/experiment_setup.py
T
2025-09-10 10:37:55 +02:00

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