mirror of
https://github.com/wassname/pytorch-ts.git
synced 2026-07-19 11:27:25 +08:00
set shuffling of time series in file and list dataset to false by default
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user