fix for pyTorch 1.2 (#549)

* min pytorch 1.2

* fix IterableDataset

* upgrade torchvision

* fix msg
This commit is contained in:
Jirka Borovec
2019-11-26 10:58:50 -05:00
committed by William Falcon
parent 55f3ffd7c7
commit f2191b0cdf
2 changed files with 16 additions and 5 deletions
@@ -1,7 +1,17 @@
import warnings
import torch.distributed as dist
from torch.utils.data import IterableDataset
try:
# loading for pyTorch 1.3
from torch.utils.data import IterableDataset
except ImportError:
# loading for pyTorch 1.1
import torch
warnings.warn('Your version of pyTorch %s does not support `IterableDataset`,'
' please upgrade to 1.2+' % torch.__version__, ImportWarning)
EXIST_ITER_DATASET = False
else:
EXIST_ITER_DATASET = True
from torch.utils.data.distributed import DistributedSampler
from pytorch_lightning.utilities.debugging import MisconfigurationException
@@ -24,7 +34,7 @@ class TrainerDataLoadingMixin(object):
self.get_train_dataloader = model.train_dataloader
# determine number of training batches
if isinstance(self.get_train_dataloader().dataset, IterableDataset):
if EXIST_ITER_DATASET and 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 +177,8 @@ class TrainerDataLoadingMixin(object):
self.get_val_dataloaders()
# support IterableDataset for train data
self.is_iterable_train_dataloader = isinstance(self.get_train_dataloader().dataset, IterableDataset)
self.is_iterable_train_dataloader = (
EXIST_ITER_DATASET and 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,
+2 -2
View File
@@ -1,8 +1,8 @@
scikit-learn>=0.20.2
tqdm>=4.35.0
numpy>=1.16.4
torch>=1.1
torchvision>=0.3.0
torch>=1.2
torchvision>=0.4.0
pandas>=0.24 # lower version do not support py3.7
test-tube>=0.6.9
# future>=0.17.1 # required for buildins in setup.py