move TransformedIterableDataset into its own file

This commit is contained in:
Dr. Kashif Rasul
2020-01-16 10:01:32 +01:00
parent 3ec57b5218
commit 1583ce1436
2 changed files with 46 additions and 30 deletions
-30
View File
@@ -95,36 +95,6 @@ class DataLoader(Iterable[DataEntry]):
self.dtype = dtype
class TransformedIterableDataset(torch.utils.data.IterableDataset):
def __init__(self, dataset, is_train, transform):
self.dataset = dataset
self.transform = transform
self.is_train = is_train
self._cur_iter = None
def _iterate_forever(self, collection: Iterable[DataEntry]) -> Iterator[DataEntry]:
# iterate forever over the collection, the collection must be non empty
while True:
try:
first = next(iter(collection))
except StopIteration:
raise Exception("empty dataset")
else:
for x in itertools.chain([first], collection):
yield x
def __iter__(self):
if self._cur_iter is None:
self._cur_iter = self.transform(
self._iterate_forever(self.dataset), is_train=self.is_train
)
assert self._cur_iter is not None
while True:
data_entry = next(self._cur_iter)
yield {k:(v.astype(np.float32) if v.dtype.kind == "f" else v) for k,v in data_entry.items() if isinstance(v, np.ndarray)==True}
class TrainDataLoader(DataLoader):
"""
An Iterable type for iterating and transforming a dataset, in batches of a
@@ -0,0 +1,46 @@
import itertools
from typing import Dict, Iterable, Iterator, Optional
import numpy as np
import torch
from pts.transform import Transformation
from .common import DataEntry, Dataset
class TransformedIterableDataset(torch.utils.data.IterableDataset):
def __init__(
self, dataset: Dataset, is_train: bool, transform: Transformation
) -> None:
self.dataset = dataset
self.transform = transform
self.is_train = is_train
self._cur_iter: Optional[Iterator] = None
def _iterate_forever(self, collection: Iterable[DataEntry]) -> Iterator[DataEntry]:
# iterate forever over the collection, the collection must be non empty
while True:
try:
first = next(iter(collection))
except StopIteration:
raise Exception("empty dataset")
else:
for x in itertools.chain([first], collection):
yield x
def __iter__(self) -> Dict[str, np.ndarray]:
if self._cur_iter is None:
self._cur_iter = self.transform(
self._iterate_forever(self.dataset), is_train=self.is_train
)
assert self._cur_iter is not None
while True:
data_entry = next(self._cur_iter)
yield {
k: (v.astype(np.float32) if v.dtype.kind == "f" else v)
for k, v in data_entry.items()
if isinstance(v, np.ndarray) == True
}
def __len__(self) -> int:
return len(self.dataset)