mirror of
https://github.com/wassname/pytorch-ts.git
synced 2026-08-12 12:20:53 +08:00
use class
This commit is contained in:
@@ -5,8 +5,7 @@ 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
|
||||
|
||||
from gluonts.itertools import Cyclic, PseudoShuffled
|
||||
|
||||
class TransformedIterableDataset(IterableDataset):
|
||||
def __init__(
|
||||
@@ -20,16 +19,16 @@ class TransformedIterableDataset(IterableDataset):
|
||||
self.shuffle_buffer_length = shuffle_buffer_length
|
||||
|
||||
self.transformed_dataset = TransformedDataset(
|
||||
cyclic(dataset),
|
||||
Cyclic(dataset),
|
||||
transform,
|
||||
is_train=is_train,
|
||||
)
|
||||
|
||||
def __iter__(self):
|
||||
if self.shuffle_buffer_length is None:
|
||||
return iter(self.transformed_dataset)
|
||||
return self.transformed_dataset
|
||||
else:
|
||||
return pseudo_shuffled(
|
||||
iter(self.transformed_dataset),
|
||||
return PseudoShuffled(
|
||||
self.transformed_dataset,
|
||||
shuffle_buffer_length=self.shuffle_buffer_length,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user