Modify training script with option to save predictions on test set

This commit is contained in:
Filip Stefaniuk
2024-09-14 14:10:29 +02:00
parent 063ea18d00
commit a3c5c76000
3 changed files with 129 additions and 75 deletions
+19 -13
View File
@@ -11,18 +11,22 @@ def get_dataset_from_wandb(run, window=None):
base_path = artifact.download()
name = artifact.metadata['name']
in_sample_name = f"in-sample-{window or run.config['data']['sliding_window']}"
in_sample_name =\
f"in-sample-{window or run.config['data']['sliding_window']}"
in_sample_data = pd.read_csv(os.path.join(
base_path, name + '-' + in_sample_name + '.csv'))
out_of_sample_name = f"out-of-sample-{window or run.config['data']['sliding_window']}"
out_of_sample_name =\
f"out-of-sample-{window or run.config['data']['sliding_window']}"
out_of_sample_data = pd.read_csv(os.path.join(
base_path, name + '-' + out_of_sample_name + '.csv'))
return in_sample_data, out_of_sample_data
def get_train_validation_split(config, in_sample_data):
validation_part = config['data']['validation']
train_data = in_sample_data.iloc[:int(len(in_sample_data) * (1 - validation_part))]
train_data = in_sample_data.iloc[:int(
len(in_sample_data) * (1 - validation_part))]
val_data = in_sample_data.iloc[len(train_data) - config['past_window']:]
return train_data, val_data
@@ -30,25 +34,27 @@ def get_train_validation_split(config, in_sample_data):
def build_time_series_dataset(config, data):
data = data.copy()
# TODO: Fix in dataset
data['weekday'] = data['weekday'].astype('str')
data['hour'] = data['hour'].astype('str')
time_series_dataset = TimeSeriesDataSet(
data,
time_idx=config['data']['fields']['time_index'],
target=config['data']['fields']['target'],
group_ids=config['data']['fields']['group_ids'],
time_idx=config['fields']['time_index'],
target=config['fields']['target'],
group_ids=config['fields']['group_ids'],
min_encoder_length=config['past_window'],
max_encoder_length=config['past_window'],
min_prediction_length=config['future_window'],
max_prediction_length=config['future_window'],
static_reals=config['data']['fields']['static_real'],
static_categoricals=config['data']['fields']['static_cat'],
time_varying_known_reals=config['data']['fields']['dynamic_known_real'],
time_varying_known_categoricals=config['data']['fields']['dynamic_known_cat'],
time_varying_unknown_reals=config['data']['fields']['dynamic_unknown_real'],
time_varying_unknown_categoricals=config['data']['fields']['dynamic_unknown_cat'],
static_reals=config['fields']['static_real'],
static_categoricals=config['fields']['static_cat'],
time_varying_known_reals=config['fields']['dynamic_known_real'],
time_varying_known_categoricals=config['fields']['dynamic_known_cat'],
time_varying_unknown_reals=config['fields']['dynamic_unknown_real'],
time_varying_unknown_categoricals=config['fields'][
'dynamic_unknown_cat'],
randomize_length=False,
)
return time_series_dataset
return time_series_dataset
+16
View File
@@ -1,8 +1,24 @@
import torch
from pytorch_forecasting import QuantileLoss
from pytorch_forecasting.metrics.base_metrics import MultiHorizonMetric
def get_loss(config):
loss_name = config['loss']['name']
if loss_name == 'Quantile':
return QuantileLoss(config['loss']['quantiles'])
if loss_name == 'GMADL':
return GMADL(
a=config['loss']['a'],
b=config['loss']['b']
)
raise ValueError("Unknown loss")
class GMADL(MultiHorizonMetric):
"""GMADL loss function."""