fixing bug in testing for IterableDataset (#547)

This commit is contained in:
MikeScarp
2019-11-26 04:59:20 -05:00
committed by William Falcon
parent 462788738b
commit 55f3ffd7c7
@@ -24,7 +24,7 @@ class TrainerDataLoadingMixin(object):
self.get_train_dataloader = model.train_dataloader
# determine number of training batches
if isinstance(self.get_train_dataloader(), IterableDataset):
if isinstance(self.get_train_dataloader().dataset, IterableDataset):
self.nb_training_batches = float('inf')
else:
self.nb_training_batches = len(self.get_train_dataloader())
@@ -167,7 +167,7 @@ class TrainerDataLoadingMixin(object):
self.get_val_dataloaders()
# support IterableDataset for train data
self.is_iterable_train_dataloader = isinstance(self.get_train_dataloader(), IterableDataset)
self.is_iterable_train_dataloader = isinstance(self.get_train_dataloader().dataset, IterableDataset)
if self.is_iterable_train_dataloader and not isinstance(self.val_check_interval, int):
m = '''
When using an iterableDataset for train_dataloader,