mirror of
https://github.com/wassname/pytorch-ts.git
synced 2026-07-25 13:30:12 +08:00
added back TransformedIterableDataset
This commit is contained in:
+69
-9
File diff suppressed because one or more lines are too long
@@ -0,0 +1 @@
|
||||
from .loader import TransformedIterableDataset
|
||||
@@ -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
@@ -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
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user