Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -30,8 +30,8 @@ def __init__(
) -> None:
self.fill_value = fill_value

self.first_timestamp = pd.Timestamp(2200, 1, 1, 12)
self.last_timestamp = pd.Timestamp(1800, 1, 1, 12)
self.first_timestamp = None
self.last_timestamp = None
self.frequency = None
self.align_data = align_data
self.max_target_length = 0
Expand All @@ -55,7 +55,7 @@ def _group_all(self, dataset: Dataset) -> Dataset:
def to_ts(self, data: DataEntry):
return pd.Series(
data["target"],
index=pd.date_range(
index=pd.period_range(
start=data["start"],
periods=len(data["target"]),
freq=data["start"].freq,
Expand All @@ -67,7 +67,7 @@ def _align_data_entry(self, data: DataEntry) -> DataEntry:
# fill target invidually if we want to fill all of them, we should use a dataframe
ts = self.to_ts(data)
d["target"] = ts.reindex(
pd.date_range(
pd.period_range(
start=self.first_timestamp,
end=self.last_timestamp,
freq=d["start"].freq,
Expand All @@ -89,6 +89,11 @@ def _preprocess(self, dataset: Dataset) -> None:
"""
for data in dataset:
timestamp = data["start"]
if self.first_timestamp is None:
self.first_timestamp = timestamp
if self.last_timestamp is None:
self.last_timestamp = timestamp

self.first_timestamp = min(self.first_timestamp, timestamp)

self.frequency = (
Expand Down Expand Up @@ -136,7 +141,7 @@ def _prepare_test_data(self, dataset):
def left_pad_data(data: DataEntry):
ts = self.to_ts(data)
filled_ts = ts.reindex(
pd.date_range(
pd.period_range(
start=self.first_timestamp,
end=ts.index[-1],
freq=data["start"].freq,
Expand All @@ -146,9 +151,13 @@ def left_pad_data(data: DataEntry):
return filled_ts.values

grouped_entry = [left_pad_data(data) for data in dataset]
grouped_entry = np.array(grouped_entry)

split_dataset = np.split(grouped_entry, self.num_test_dates)
assert len(grouped_entry) % self.num_test_dates == 0
split_size = len(grouped_entry) // self.num_test_dates
split_dataset = [
grouped_entry[i : i + split_size]
for i in range(0, len(grouped_entry), split_size)
]

all_entries = list()
for dataset_at_test_date in split_dataset:
Expand Down
48 changes: 48 additions & 0 deletions test/nursery/robust_mts_attack/test_grouper.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,48 @@
import importlib.util
from pathlib import Path

import numpy as np

from gluonts.dataset.common import ListDataset


def load_grouper():
path = (
Path(__file__).parents[3]
/ "src"
/ "gluonts"
/ "nursery"
/ "robust-mts-attack"
/ "multivariate"
/ "datasets"
/ "grouper.py"
)
spec = importlib.util.spec_from_file_location(
"robust_mts_attack_grouper", path
)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module.Grouper


def test_grouper_splits_rolling_test_data_before_stacking():
dataset = ListDataset(
[
{"start": "2014-09-07", "target": [1, 2, 3, 4]},
{"start": "2014-09-07", "target": [5, 6, 7, 8]},
{"start": "2014-09-08", "target": [0, 1, 2, 3]},
{"start": "2014-09-08", "target": [4, 5, 6, 7]},
],
freq="1D",
)

Grouper = load_grouper()
grouped_data = list(Grouper(num_test_dates=2)(dataset))

np.testing.assert_array_equal(
grouped_data[0]["target"], np.array([[1, 2, 3, 4], [5, 6, 7, 8]])
)
np.testing.assert_array_equal(
grouped_data[1]["target"],
np.array([[0, 0, 1, 2, 3], [0, 4, 5, 6, 7]]),
)