diff --git a/pts/dataset/artificial.py b/pts/dataset/artificial.py index 7a9a602..de36bd9 100644 --- a/pts/dataset/artificial.py +++ b/pts/dataset/artificial.py @@ -72,7 +72,7 @@ class ArtificialDataset: return TrainDatasets( metadata=self.metadata, train=ListDataset(self.train, self.freq), - test=ListDataset(self.test, self.freq, is_train=False), + test=ListDataset(self.test, self.freq, shuffle=False), ) @@ -776,7 +776,6 @@ def constant_dataset() -> Tuple[DatasetInfo, Dataset, Dataset]: for i in range(10) ], freq=metadata.freq, - is_train=False ) info = DatasetInfo( diff --git a/pts/dataset/file_dataset.py b/pts/dataset/file_dataset.py index fac1d0d..cdef44c 100644 --- a/pts/dataset/file_dataset.py +++ b/pts/dataset/file_dataset.py @@ -37,14 +37,14 @@ class JsonLinesFile: JSON Lines file. """ - def __init__(self, path: Path, is_train: bool = True) -> None: + def __init__(self, path: Path, shuffle: bool = True) -> None: self.path = path - self.is_train = is_train + self.shuffle = shuffle def __iter__(self): with open(self.path) as jsonl_file: lines = jsonl_file.read().splitlines() - if self.is_train: + if self.shuffle: random.shuffle(lines) for line_number, raw in enumerate(lines, start=1): @@ -81,9 +81,9 @@ class FileDataset(Dataset): """ def __init__( - self, path: Path, freq: str, one_dim_target: bool = True, is_train: bool = True + self, path: Path, freq: str, one_dim_target: bool = True, shuffle: bool = False ) -> None: - self.is_train = is_train + self.shuffle = shuffle self.path = path self.process = ProcessDataEntry(freq, one_dim_target=one_dim_target) if not self.files(): @@ -91,7 +91,7 @@ class FileDataset(Dataset): def __iter__(self) -> Iterator[DataEntry]: for path in self.files(): - for line in JsonLinesFile(path, self.is_train): + for line in JsonLinesFile(path, self.shuffle): data = self.process(line.content) data["source"] = SourceContext( source=line.span.path, row=line.span.line @@ -110,4 +110,8 @@ class FileDataset(Dataset): List[Path] List of the paths of all files composing the dataset. """ - return glob.glob(str(self.path)) + files = glob.glob(str(self.path)) + if self.shuffle: + random.shuffle(files) + return files + diff --git a/pts/dataset/list_dataset.py b/pts/dataset/list_dataset.py index e352e94..a2d2170 100644 --- a/pts/dataset/list_dataset.py +++ b/pts/dataset/list_dataset.py @@ -11,14 +11,13 @@ class ListDataset(Dataset): data_iter: Iterable[DataEntry], freq: str, one_dim_target: bool = True, - is_train: bool = True, + shuffle: bool = False, ) -> None: process = ProcessDataEntry(freq, one_dim_target) self.list_data = [process(data) for data in data_iter] - if is_train: + if shuffle: random.shuffle(self.list_data) - def __iter__(self): source_name = "list_data" for row_number, data in enumerate(self.list_data, start=1): diff --git a/pts/dataset/multivariate_grouper.py b/pts/dataset/multivariate_grouper.py index e42ca18..d6aaaa2 100644 --- a/pts/dataset/multivariate_grouper.py +++ b/pts/dataset/multivariate_grouper.py @@ -122,9 +122,7 @@ class MultivariateGrouper: grouped_data[FieldName.START] = self.first_timestamp grouped_data[FieldName.FEAT_STATIC_CAT] = [0] - return ListDataset( - [grouped_data], freq=self.frequency, one_dim_target=False, is_train=False - ) + return ListDataset([grouped_data], freq=self.frequency, one_dim_target=False) def _prepare_test_data(self, dataset: Dataset) -> ListDataset: logging.info("group test time-series to datasets") diff --git a/pts/dataset/repository/_gp_copula_2019.py b/pts/dataset/repository/_gp_copula_2019.py index 416175c..2debcbf 100644 --- a/pts/dataset/repository/_gp_copula_2019.py +++ b/pts/dataset/repository/_gp_copula_2019.py @@ -137,9 +137,7 @@ def save_metadata(dataset_path: Path, ds_info: GPCopulaDataset): def save_dataset(dataset_path: Path, ds_info: GPCopulaDataset): - dataset = list( - FileDataset(dataset_path / "*.json", freq=ds_info.freq, is_train=False) - ) + dataset = list(FileDataset(dataset_path / "*.json", freq=ds_info.freq)) shutil.rmtree(dataset_path) train_file = dataset_path / "data.json" save_to_file( diff --git a/pts/dataset/utils.py b/pts/dataset/utils.py index a56cc37..e571cde 100644 --- a/pts/dataset/utils.py +++ b/pts/dataset/utils.py @@ -42,7 +42,7 @@ def to_pandas(instance: dict, freq: str = None) -> pd.Series: return pd.Series(target, index=index) -def load_datasets(metadata, train, test) -> TrainDatasets: +def load_datasets(metadata, train, test, shuffle: bool = False) -> TrainDatasets: """ Loads a dataset given metadata, train and test path. Parameters @@ -59,8 +59,8 @@ def load_datasets(metadata, train, test) -> TrainDatasets: An object collecting metadata, training data, test data. """ meta = MetaData.parse_file(metadata) - train_ds = FileDataset(train, meta.freq) - test_ds = FileDataset(test, meta.freq, is_train=False) if test else None + train_ds = FileDataset(train, meta.freq, shuffle=shuffle) + test_ds = FileDataset(test, meta.freq) if test else None return TrainDatasets(metadata=meta, train=train_ds, test=test_ds) diff --git a/test/dataset/test_multivariate_grouper.py b/test/dataset/test_multivariate_grouper.py index 12e4afd..730ea72 100644 --- a/test/dataset/test_multivariate_grouper.py +++ b/test/dataset/test_multivariate_grouper.py @@ -67,10 +67,8 @@ TRAIN_FILL_RULE = [np.mean, np.mean, np.mean, np.mean, lambda x: 0.0] def test_multivariate_grouper_train( univariate_ts, multivariate_ts, train_fill_rule ) -> None: - univariate_ds = ListDataset(univariate_ts, freq="1D", is_train=False) - multivariate_ds = ListDataset( - multivariate_ts, freq="1D", one_dim_target=False, is_train=False - ) + univariate_ds = ListDataset(univariate_ts, freq="1D") + multivariate_ds = ListDataset(multivariate_ts, freq="1D", one_dim_target=False) grouper = MultivariateGrouper(train_fill_rule=train_fill_rule) assert ( @@ -117,10 +115,9 @@ MAX_TARGET_DIM = [2, 1] def test_multivariate_grouper_test( univariate_ts, multivariate_ts, test_fill_rule, max_target_dim ) -> None: - univariate_ds = ListDataset(univariate_ts, freq="1D", is_train=False) + univariate_ds = ListDataset(univariate_ts, freq="1D") multivariate_ds = ListDataset( - multivariate_ts, freq="1D", one_dim_target=False, is_train=False - ) + multivariate_ts, freq="1D", one_dim_target=False) grouper = MultivariateGrouper( test_fill_rule=test_fill_rule, num_test_dates=2, max_target_dim=max_target_dim,