Files
pytorch-ts/pts/dataset/transformed_dataset.py
T
2019-07-15 14:15:12 +02:00

32 lines
917 B
Python

from typing import List, Iterator
from .common import Dataset, DataEntry
from pts.feature import Transformation, Chain
class TransformedDataset(Dataset):
"""
A dataset that corresponds to applying a list of transformations to each
element in the base_dataset.
This only supports SimpleTransformations, which do the same thing at
prediction and training time.
Parameters
----------
base_dataset
Dataset to transform
transformations
List of transformations to apply
"""
def __init__(
self, base_dataset: Dataset, transformations: List[Transformation]
) -> None:
self.base_dataset = base_dataset
self.transformations = Chain(transformations)
def __iter__(self) -> Iterator[DataEntry]:
yield from self.transformations(self.base_dataset, is_train=True)
def __len__(self):
return sum(1 for _ in self)