fixes
This commit is contained in:
@@ -0,0 +1,132 @@
|
||||
from typing import Callable
|
||||
|
||||
import configs
|
||||
from configs.feature_config import feature_config
|
||||
from configs.cnn_run_config import get_cnn_run_config
|
||||
from configs.transformer_run_config import get_transformer_run_config
|
||||
|
||||
take_every_nth = int(288 / 12)
|
||||
shift_in_hours = 12
|
||||
input_window_length = (288 // take_every_nth) * 80
|
||||
# output_window_length = (288 // take_every_nth) * 1
|
||||
output_window_length = 1
|
||||
output_window_offset = input_window_length + (288 // take_every_nth) * 0
|
||||
window_shift = int((288 // take_every_nth) / 24 * shift_in_hours)
|
||||
min_input_length_fraction_for_padding = ((288 // take_every_nth) * 4) / input_window_length
|
||||
|
||||
|
||||
def config_generator(config_gen_fn: Callable,
|
||||
fixed_params: dict,
|
||||
variable_param_configs: list,
|
||||
) -> list:
|
||||
configs = list()
|
||||
for variable_config in variable_param_configs:
|
||||
config = fixed_params.copy()
|
||||
for param_name, param_value in variable_config.items():
|
||||
# update the config with the variable parameter, values can be None, if default should be used
|
||||
if param_value is not None:
|
||||
config[param_name] = param_value
|
||||
configs.append(config_gen_fn(**config))
|
||||
|
||||
return configs
|
||||
|
||||
|
||||
run_configuration = {
|
||||
"name": "ovulation_regression",
|
||||
"runs": config_generator(
|
||||
get_cnn_run_config,
|
||||
fixed_params={
|
||||
"run_name": "cnn_ovulation_regression",
|
||||
"run_description": "CNN model for ovulation regression",
|
||||
"input_window_length": input_window_length,
|
||||
"output_window_length": output_window_length,
|
||||
"take_every_nth": take_every_nth,
|
||||
"shift_in_hours": shift_in_hours,
|
||||
"output_window_offset": output_window_offset,
|
||||
"batch_size": 256,
|
||||
"max_lr": 1e-5,
|
||||
"num_epochs": 10,
|
||||
"patience": 3,
|
||||
"feature_config": feature_config,
|
||||
},
|
||||
variable_param_configs=[
|
||||
{
|
||||
"model_parameters": {
|
||||
"cnn_channels": 32,
|
||||
"kernel_size": 3,
|
||||
"embed_dim": 32,
|
||||
"num_enc_layers": 2,
|
||||
"num_heads": 2,
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_parameters": {
|
||||
"cnn_channels": 64,
|
||||
"kernel_size": 3,
|
||||
"embed_dim": 64,
|
||||
"num_enc_layers": 2,
|
||||
"num_heads": 2,
|
||||
},
|
||||
},
|
||||
{
|
||||
"batch_size": 128,
|
||||
"model_parameters": {
|
||||
"cnn_channels": 128,
|
||||
"kernel_size": 5,
|
||||
"embed_dim": 128,
|
||||
"num_enc_layers": 4,
|
||||
"num_heads": 4,
|
||||
},
|
||||
},
|
||||
]
|
||||
) + config_generator(
|
||||
get_transformer_run_config,
|
||||
fixed_params={
|
||||
"run_name": "transformer_ovulation_regression",
|
||||
"run_description": "Transformer model for ovulation regression",
|
||||
"input_window_length": input_window_length,
|
||||
"output_window_length": output_window_length,
|
||||
"take_every_nth": take_every_nth,
|
||||
"shift_in_hours": shift_in_hours,
|
||||
"output_window_offset": output_window_offset,
|
||||
"batch_size": 128,
|
||||
"max_lr": 1e-5,
|
||||
"num_epochs": 10,
|
||||
"patience": 3,
|
||||
"feature_config": feature_config,
|
||||
},
|
||||
variable_param_configs=[
|
||||
{
|
||||
"model_parameters": {
|
||||
"embed_dim": 64,
|
||||
"num_enc_layers": 2,
|
||||
"num_heads": 2,
|
||||
},
|
||||
},
|
||||
{
|
||||
"batch_size": 64,
|
||||
"model_parameters": {
|
||||
"embed_dim": 128,
|
||||
"num_enc_layers": 4,
|
||||
"num_heads": 4,
|
||||
},
|
||||
},
|
||||
{
|
||||
"batch_size": 64,
|
||||
"model_parameters": {
|
||||
"embed_dim": 256,
|
||||
"num_enc_layers": 4,
|
||||
"num_heads": 4,
|
||||
},
|
||||
},
|
||||
{
|
||||
"batch_size": 32,
|
||||
"model_parameters": {
|
||||
"embed_dim": 512,
|
||||
"num_enc_layers": 4,
|
||||
"num_heads": 4,
|
||||
},
|
||||
},
|
||||
]
|
||||
)
|
||||
}
|
||||
Reference in New Issue
Block a user