mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-11 12:31:23 +08:00
fixing bug in testing for IterableDataset (#547)
This commit is contained in:
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,
|
||||
|
||||
Reference in New Issue
Block a user