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

200 lines
6.6 KiB
Python

from functools import partial
from configs.config_utils import *
from configs.feature_config import feature_config
from experiment_setup import get_fertility_based_eval_functions
from utils.data_utils import *
from utils.training import *
from models.lstm import *
from models.utils import *
from models.collation import *
lstm_base_config = {
"model_name": "lstm",
"version": "1.0.0",
"model_class": LSTM,
"feature_config": feature_config | {"ignored_features":
[
"fertility_probability",
"ov_over_probability",
# "days_relative_to_ov",
# "is_biphasic",
"ov_day",
]},
"batch_fn": produce_window_batches,
"collate_fn": simple_x_y_collate,
"model_creation_fn": simple_model_creation,
"model_save_fn": simple_model_save,
"model_load_fn": partial(simple_model_load, model_creation_fn=simple_model_creation),
"batch_loss_fn": get_model_loss,
"actual_fn": simple_get_y,
"predict_fn": simple_x_y_predict,
}
base_lstm_training_config = {
"batch_size": 256,
"model_class": LSTM,
"learning_parameters": {
"max_lr": 1e-4,
# "learning_rate": base_lr * (batch_size / 4),
"epochs": 30,
"patience": 5,
},
"max_grad_norm": 1.0,
"train_size": 0.7,
"val_size": 0.15,
"test_size": 0.15,
}
model_input_run = get_run_config(
run_name="lstm_input_size_run",
base_model_configuration=lstm_base_config |
{
"model_parameters": {
"lstm_hidden_size": 64,
"num_layers": 2,
},
},
base_training_configuration=base_lstm_training_config,
model_config_kwargs_list=input_run_kwargs_list
)
fertility_complexity_run = get_run_config(
run_name="lstm_fertility_complexity_run",
base_model_configuration=lstm_base_config | get_data_config_set(
values_per_day=12,
shift_in_hours=12,
input_window_length_in_days=160,
output_window_length=1,
output_window_offset=0
) | {
"feature_config": feature_config | {"ignored_features":
[
# "fertility_probability",
# "ov_over_probability",
"days_relative_to_ov",
"is_biphasic",
"ov_day",
]},
# "model_parameters": {
# "lstm_hidden_size": 128,
# "num_layers": 4,
# },
},
base_training_configuration=base_lstm_training_config,
model_config_kwargs_list=[
{
"model_parameters": {
"lstm_hidden_size": 16,
"num_layers": 1,
},
},
{
"model_parameters": {
"lstm_hidden_size": 32,
"num_layers": 1,
},
},
{
"model_parameters": {
"lstm_hidden_size": 32,
"num_layers": 2,
},
},
{
"model_parameters": {
"lstm_hidden_size": 64,
"num_layers": 2,
},
},
{
"model_parameters": {
"lstm_hidden_size": 128,
"num_layers": 2,
},
},
{
"model_parameters": {
"lstm_hidden_size": 128,
"num_layers": 4,
},
},
{
"model_parameters": {
"lstm_hidden_size": 256,
"num_layers": 4,
},
},
{
"model_parameters": {
"lstm_hidden_size": 512,
"num_layers": 4,
},
},
],
evaluation_configuration={
"eval_functions_getter": get_fertility_based_eval_functions
},
)
fertility_input_run = get_run_config(
run_name="lstm_fertility_input_size_run",
base_model_configuration=lstm_base_config |
{
"feature_config": feature_config | {"ignored_features":
[
# "fertility_probability",
# "ov_over_probability",
"days_relative_to_ov",
"is_biphasic",
"ov_day",
]},
"model_parameters": {
"lstm_hidden_size": 128,
"num_layers": 4,
},
},
base_training_configuration=base_lstm_training_config,
model_config_kwargs_list=input_run_kwargs_list,
evaluation_configuration={
"eval_functions_getter": get_fertility_based_eval_functions
},
)
best_config_run = get_run_config(
run_name="lstm_best_config_run",
base_model_configuration=lstm_base_config | get_data_config_set(
values_per_day=12,
shift_in_hours=12,
input_window_length_in_days=160,
output_window_length=1,
output_window_offset=0
) | {
"feature_config": feature_config | {"ignored_features":
[
# "fertility_probability",
# "ov_over_probability",
"days_relative_to_ov",
"is_biphasic",
"ov_day",
]},
"model_parameters": {
"lstm_hidden_size": 512,
"num_layers": 4,
},
},
base_training_configuration=base_lstm_training_config | {
"learning_parameters": {
"max_lr": 1e-4,
# "learning_rate": base_lr * (batch_size / 4),
"epochs": 50,
"patience": 7,
},
},
model_config_kwargs_list=[],
evaluation_configuration={
"eval_functions_getter": get_fertility_based_eval_functions
},
)