Fixing tests (#936)

* abs import

* rename test model

* update trainer

* revert test_step check

* move tags

* fix test_step

* clean tests

* fix template

* update dataset path

* fix parent order
This commit is contained in:
Jirka Borovec
2020-02-25 13:06:24 -05:00
committed by GitHub
parent 20d15c8023
commit 5dd2afeab1
15 changed files with 264 additions and 209 deletions
+10 -7
View File
@@ -177,7 +177,7 @@ class TrainerDataLoadingMixin(ABC):
self.is_iterable_train_dataloader = (
EXIST_ITER_DATASET and isinstance(self.train_dataloader.dataset, IterableDataset)
)
if self.is_iterable_train_dataloader and not isinstance(self.val_check_interval, int):
if self.is_iterable_dataloader(self.train_dataloader) and not isinstance(self.val_check_interval, int):
m = '''
When using an iterableDataset for `train_dataloader`,
`Trainer(val_check_interval)` must be an int.
@@ -185,6 +185,11 @@ class TrainerDataLoadingMixin(ABC):
'''
raise MisconfigurationException(m)
def is_iterable_dataloader(self, dataloader):
return (
EXIST_ITER_DATASET and isinstance(dataloader.dataset, IterableDataset)
)
def reset_val_dataloader(self, model):
"""
Dataloaders are provided by the model
@@ -200,9 +205,8 @@ class TrainerDataLoadingMixin(ABC):
self.num_val_batches = 0
# add samplers
for i, dataloader in enumerate(self.val_dataloaders):
dl = self.auto_add_sampler(dataloader, train=False)
self.val_dataloaders[i] = dl
self.val_dataloaders = [self.auto_add_sampler(dl, train=False)
for dl in self.val_dataloaders if dl]
# determine number of validation batches
# val datasets could be none, 1 or 2+
@@ -227,9 +231,8 @@ class TrainerDataLoadingMixin(ABC):
self.num_test_batches = 0
# add samplers
for i, dataloader in enumerate(self.test_dataloaders):
dl = self.auto_add_sampler(dataloader, train=False)
self.test_dataloaders[i] = dl
self.test_dataloaders = [self.auto_add_sampler(dl, train=False)
for dl in self.test_dataloaders if dl]
# determine number of test batches
if self.test_dataloaders is not None: