fixed dataset stuff + docs (#1599)

* Fixed dataset docs and disabled auto-sampler for iterable dataset
This commit is contained in:
William Falcon
2020-04-24 16:51:26 -04:00
committed by GitHub
parent d56b3e5e69
commit d0faf97893
2 changed files with 59 additions and 1 deletions
+14 -1
View File
@@ -10,6 +10,12 @@ from pytorch_lightning.core import LightningModule
from pytorch_lightning.utilities import rank_zero_warn
from pytorch_lightning.utilities.exceptions import MisconfigurationException
try:
from torch.utils.data import IterableDataset
ITERABLE_DATASET_EXISTS = True
except ImportError:
ITERABLE_DATASET_EXISTS = False
try:
from apex import amp
except ImportError:
@@ -95,7 +101,14 @@ class TrainerDataLoadingMixin(ABC):
def auto_add_sampler(self, dataloader: DataLoader, train: bool) -> DataLoader:
# don't do anything if it's not a dataloader
if not isinstance(dataloader, DataLoader):
# don't manipulate iterable datasets
is_dataloader = isinstance(dataloader, DataLoader)
is_iterable_ds = False
if ITERABLE_DATASET_EXISTS and hasattr(dataloader, 'dataset'):
is_iterable_ds = isinstance(dataloader.dataset, IterableDataset)
if not is_dataloader or is_iterable_ds:
return dataloader
need_dist_sampler = (self.use_ddp or self.use_ddp2 or self.use_horovod or self.use_tpu)
if self.replace_sampler_ddp and need_dist_sampler: