Modify training script with option to save predictions on test set
This commit is contained in:
+19
-13
@@ -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
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user