added back TransformedIterableDataset

This commit is contained in:
Dr. Kashif Rasul
2020-12-18 13:02:26 +01:00
parent b072ab227b
commit 2726bc94ec
5 changed files with 132 additions and 33 deletions
+69 -9
View File
File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
from .loader import TransformedIterableDataset
+35
View File
@@ -0,0 +1,35 @@
from typing import Callable, Iterable, Iterator, List, Optional
from torch.utils import data
from torch.utils.data import IterableDataset
from gluonts.dataset.common import Dataset
from gluonts.transform import Transformation, TransformedDataset
from gluonts.itertools import cyclic, pseudo_shuffled
class TransformedIterableDataset(IterableDataset):
def __init__(
self,
dataset: Dataset,
transform: Transformation,
is_train: bool = True,
shuffle_buffer_length: Optional[int] = None,
):
super().__init__()
self.shuffle_buffer_length = shuffle_buffer_length
self.transformed_dataset = TransformedDataset(
cyclic(dataset),
transform,
is_train=is_train,
)
def __iter__(self):
if self.shuffle_buffer_length is None:
return iter(self.transformed_dataset)
else:
return pseudo_shuffled(
iter(self.transformed_dataset),
shuffle_buffer_length=self.shuffle_buffer_length,
)
+20 -17
View File
@@ -5,10 +5,11 @@ import numpy as np
import torch
import torch.nn as nn
from torch.utils import data
from torch.utils.data import DataLoader
from gluonts.core.component import validated
from gluonts.dataset.common import Dataset
from gluonts.dataset.loader import TrainDataLoader, ValidationDataLoader
from gluonts.model.estimator import Estimator
from gluonts.torch.model.predictor import PyTorchPredictor
from gluonts.torch.batchify import batchify
@@ -16,6 +17,7 @@ from gluonts.transform import SelectFields, Transformation
from pts import Trainer
from pts.model import get_module_forward_input_names
from pts.dataset.loader import TransformedIterableDataset
class TrainOutput(NamedTuple):
@@ -78,7 +80,7 @@ class PyTorchEstimator(Estimator):
training_data: Dataset,
validation_data: Optional[Dataset] = None,
num_workers: Optional[int] = None,
num_prefetch: Optional[int] = None,
prefetch_factor: Optional[int] = 2,
shuffle_buffer_length: Optional[int] = None,
**kwargs,
) -> TrainOutput:
@@ -88,32 +90,33 @@ class PyTorchEstimator(Estimator):
input_names = get_module_forward_input_names(trained_net)
training_data_loader = TrainDataLoader(
training_iter_dataset = TransformedIterableDataset(
dataset=training_data,
transform=transformation + SelectFields(input_names),
batch_size=self.trainer.batch_size,
stack_fn=partial(
batchify,
device=self.trainer.device,
),
num_workers=num_workers,
num_prefetch=num_prefetch,
is_train=True,
shuffle_buffer_length=shuffle_buffer_length,
)
training_data_loader = DataLoader(
training_iter_dataset,
batch_size=self.trainer.batch_size,
num_workers=num_workers,
prefetch_factor=prefetch_factor,
**kwargs,
)
validation_data_loader = None
if validation_data is not None:
validation_data_loader = ValidationDataLoader(
validation_iter_dataset = TransformedIterableDataset(
dataset=validation_data,
transform=transformation + SelectFields(input_names),
is_train=True,
)
validation_data_loader = DataLoader(
validation_iter_dataset,
batch_size=self.trainer.batch_size,
stack_fn=partial(
batchify,
device=self.trainer.device,
),
num_workers=num_workers,
num_prefetch=num_prefetch,
prefetch_factor=prefetch_factor,
**kwargs,
)
@@ -128,7 +131,7 @@ class PyTorchEstimator(Estimator):
trained_net=trained_net,
predictor=self.create_predictor(
transformation, trained_net, self.trainer.device
),
)
)
def train(
+7 -7
View File
@@ -5,9 +5,9 @@ from tqdm import tqdm
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from gluonts.core.component import validated
from gluonts.dataset.loader import TrainDataLoader, ValidationDataLoader
class Trainer:
@@ -35,8 +35,8 @@ class Trainer:
def __call__(
self,
net: nn.Module,
train_iter: TrainDataLoader,
validation_iter: Optional[ValidationDataLoader] = None,
train_iter: DataLoader,
validation_iter: Optional[DataLoader] = None,
) -> None:
optimizer = torch.optim.Adam(
net.parameters(), lr=self.learning_rate, weight_decay=self.weight_decay
@@ -50,9 +50,9 @@ class Trainer:
with tqdm(train_iter) as it:
for batch_no, data_entry in enumerate(it, start=1):
optimizer.zero_grad()
#inputs = [data_entry[k].to(self.device) for k in input_names]
inputs = [v.to(self.device) for v in data_entry.values()]
output = net(*data_entry.values())
output = net(*inputs)
if isinstance(output, (list, tuple)):
loss = output[0]
else:
@@ -67,7 +67,7 @@ class Trainer:
refresh=False,
)
n_iter = epoch_no * self.num_batches_per_epoch + batch_no
#.add_scalar("Loss/train", loss.item(), n_iter)
# .add_scalar("Loss/train", loss.item(), n_iter)
loss.backward()
optimizer.step()
@@ -82,4 +82,4 @@ class Trainer:
# mark epoch end time and log time cost of current epoch
toc = time.time()
#writer.close()
# writer.close()