fix item_id field in repository datasets

This commit is contained in:
Dr. Kashif Rasul
2020-01-21 12:23:00 +01:00
parent f4c2406934
commit c32bc0977a
4 changed files with 22 additions and 6 deletions
@@ -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)
],
+4 -1
View File
@@ -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,
)
)
+7 -3
View File
@@ -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)
],
)
+10 -2
View File
@@ -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