mirror of
https://github.com/wassname/pytorch-ts.git
synced 2026-08-05 13:21:07 +08:00
fix item_id field in repository datasets
This commit is contained in:
@@ -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)
|
||||
],
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -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)
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user