From c32bc0977a6e1e6510fd28ab1575c1a39ae10a3f Mon Sep 17 00:00:00 2001 From: "Dr. Kashif Rasul" Date: Tue, 21 Jan 2020 12:23:00 +0100 Subject: [PATCH] fix item_id field in repository datasets --- pts/dataset/repository/_gp_copula_2019.py | 1 + pts/dataset/repository/_lstnet.py | 5 ++++- pts/dataset/repository/_m4.py | 10 +++++++--- pts/dataset/repository/_util.py | 12 ++++++++++-- 4 files changed, 22 insertions(+), 6 deletions(-) diff --git a/pts/dataset/repository/_gp_copula_2019.py b/pts/dataset/repository/_gp_copula_2019.py index 0428470..4293e07 100644 --- a/pts/dataset/repository/_gp_copula_2019.py +++ b/pts/dataset/repository/_gp_copula_2019.py @@ -149,6 +149,7 @@ def save_dataset(dataset_path: Path, ds_info: GPCopulaDataset): # Handles adding categorical features of rolling # evaluation dates cat=[cat - ds_info.num_series * (cat // ds_info.num_series)], + item_id=cat, ) for cat, data_entry in enumerate(dataset) ], diff --git a/pts/dataset/repository/_lstnet.py b/pts/dataset/repository/_lstnet.py index 638a321..311385c 100644 --- a/pts/dataset/repository/_lstnet.py +++ b/pts/dataset/repository/_lstnet.py @@ -162,7 +162,10 @@ def generate_lstnet_dataset(dataset_path: Path, dataset_name: str): if len(sliced_ts) > 0: train_ts.append( to_dict( - target_values=sliced_ts.values, start=sliced_ts.index[0], cat=[cat], + target_values=sliced_ts.values, + start=sliced_ts.index[0], + cat=[cat], + item_id=cat, ) ) diff --git a/pts/dataset/repository/_m4.py b/pts/dataset/repository/_m4.py index bda0250..ffb1641 100644 --- a/pts/dataset/repository/_m4.py +++ b/pts/dataset/repository/_m4.py @@ -66,7 +66,9 @@ def generate_m4_dataset( save_to_file( train_file, [ - to_dict(target_values=target, start=mock_start_dataset, cat=[cat]) + to_dict( + target_values=target, start=mock_start_dataset, cat=[cat], item_id=cat + ) for cat, target in enumerate(train_target_values) ], ) @@ -74,8 +76,10 @@ def generate_m4_dataset( save_to_file( test_file, [ - to_dict(target_values=target, start=mock_start_dataset, cat=[cat]) + to_dict( + target_values=target, start=mock_start_dataset, cat=[cat], item_id=cat + ) for cat, target in enumerate(test_target_values) ], ) - \ No newline at end of file + diff --git a/pts/dataset/repository/_util.py b/pts/dataset/repository/_util.py index 46f45a3..dd2e1bc 100644 --- a/pts/dataset/repository/_util.py +++ b/pts/dataset/repository/_util.py @@ -14,12 +14,17 @@ import json import os from pathlib import Path -from typing import Dict, List, Optional +from typing import Dict, List, Optional, Any import numpy as np -def to_dict(target_values: np.ndarray, start: str, cat: Optional[List[int]] = None): +def to_dict( + target_values: np.ndarray, + start: str, + cat: Optional[List[int]] = None, + item_id: Optional[Any] = None, +): def serialize(x): if np.isnan(x): return "NaN" @@ -34,6 +39,9 @@ def to_dict(target_values: np.ndarray, start: str, cat: Optional[List[int]] = No if cat is not None: res["feat_static_cat"] = cat + + if item_id is not None: + res["item_id"] = item_id return res