diff --git a/pytorch_lightning/trainer/data_loading_mixin.py b/pytorch_lightning/trainer/data_loading_mixin.py index 066c21dc..7755bc7d 100644 --- a/pytorch_lightning/trainer/data_loading_mixin.py +++ b/pytorch_lightning/trainer/data_loading_mixin.py @@ -15,8 +15,13 @@ except ImportError: class TrainerDataLoadingMixin(object): - - def layout_bookeeping(self): + def init_train_dataloader(self, model): + """ + Dataloaders are provided by the model + :param model: + :return: + """ + self.get_train_dataloader = model.train_dataloader # determine number of training batches if isinstance(self.get_train_dataloader(), IterableDataset): @@ -25,21 +30,6 @@ class TrainerDataLoadingMixin(object): self.nb_training_batches = len(self.get_train_dataloader()) self.nb_training_batches = int(self.nb_training_batches * self.train_percent_check) - # determine number of validation batches - # val datasets could be none, 1 or 2+ - if self.get_val_dataloaders() is not None: - self.nb_val_batches = sum(len(dataloader) for dataloader in self.get_val_dataloaders()) - self.nb_val_batches = int(self.nb_val_batches * self.val_percent_check) - self.nb_val_batches = max(1, self.nb_val_batches) - - # determine number of test batches - if self.get_test_dataloaders() is not None: - self.nb_test_batches = sum( - len(dataloader) for dataloader in self.get_test_dataloaders() - ) - self.nb_test_batches = int(self.nb_test_batches * self.test_percent_check) - self.nb_test_batches = max(1, self.nb_test_batches) - # determine when to check validation # if int passed in, val checks that often # otherwise, it checks in [0, 1.0] % range of a training epoch @@ -49,86 +39,123 @@ class TrainerDataLoadingMixin(object): self.val_check_batch = int(self.nb_training_batches * self.val_check_interval) self.val_check_batch = max(1, self.val_check_batch) + on_ddp = self.use_ddp or self.use_ddp2 + if on_ddp and not isinstance(self.get_train_dataloader().sampler, DistributedSampler): + msg = """ + You're using multiple gpus and multiple nodes without using a DistributedSampler + to assign a subset of your data to each process. To silence this warning, pass a + DistributedSampler to your DataLoader. + + ie: this: + dataset = myDataset() + dataloader = Dataloader(dataset) + + becomes: + dataset = myDataset() + dist_sampler = torch.utils.data.distributed.DistributedSampler(dataset) + dataloader = Dataloader(dataset, sampler=dist_sampler) + + If you want each process to load the full dataset, ignore this warning. + """ + if msg not in self.shown_warnings and self.proc_rank == 0: + self.shown_warnings.add(msg) + warnings.warn(msg) + + def init_val_dataloader(self, model): + """ + Dataloaders are provided by the model + :param model: + :return: + """ + self.get_val_dataloaders = model.val_dataloader + + # determine number of validation batches + # val datasets could be none, 1 or 2+ + if self.get_val_dataloaders() is not None: + self.nb_val_batches = sum(len(dataloader) for dataloader in self.get_val_dataloaders()) + self.nb_val_batches = int(self.nb_val_batches * self.val_percent_check) + self.nb_val_batches = max(1, self.nb_val_batches) + + on_ddp = self.use_ddp or self.use_ddp2 + if on_ddp and self.get_val_dataloaders() is not None: + for dataloader in self.get_val_dataloaders(): + if not isinstance(dataloader.sampler, DistributedSampler): + msg = """ + Your val_dataloader(s) don't use DistributedSampler. + + You're using multiple gpus and multiple nodes without using a + DistributedSampler to assign a subset of your data to each process. + To silence this warning, pass a DistributedSampler to your DataLoader. + + ie: this: + dataset = myDataset() + dataloader = Dataloader(dataset) + + becomes: + dataset = myDataset() + dist_sampler = torch.utils.data.distributed.DistributedSampler(dataset) + dataloader = Dataloader(dataset, sampler=dist_sampler) + + If you want each process to load the full dataset, ignore this warning. + """ + if msg not in self.shown_warnings and self.proc_rank == 0: + self.shown_warnings.add(msg) + warnings.warn(msg) + break + + def init_test_dataloader(self, model): + """ + Dataloaders are provided by the model + :param model: + :return: + """ + + self.get_test_dataloaders = model.test_dataloader + + # determine number of test batches + if self.get_test_dataloaders() is not None: + len_sum = sum(len(dataloader) for dataloader in self.get_test_dataloaders()) + self.nb_test_batches = len_sum + self.nb_test_batches = int(self.nb_test_batches * self.test_percent_check) + self.nb_test_batches = max(1, self.nb_test_batches) + + on_ddp = self.use_ddp or self.use_ddp2 + if on_ddp and self.get_test_dataloaders() is not None: + for dataloader in self.get_test_dataloaders(): + if not isinstance(dataloader.sampler, DistributedSampler): + msg = """ + Your test_dataloader(s) don't use DistributedSampler. + + You're using multiple gpus and multiple nodes without using a + DistributedSampler to assign a subset of your data to each process. + To silence this warning, pass a DistributedSampler to your DataLoader. + + ie: this: + dataset = myDataset() + dataloader = Dataloader(dataset) + + becomes: + dataset = myDataset() + dist_sampler = torch.utils.data.distributed.DistributedSampler(dataset) + dataloader = Dataloader(dataset, sampler=dist_sampler) + + If you want each process to load the full dataset, ignore this warning. + """ + if msg not in self.shown_warnings and self.proc_rank == 0: + self.shown_warnings.add(msg) + warnings.warn(msg) + break + def get_dataloaders(self, model): """ Dataloaders are provided by the model :param model: :return: """ - self.get_train_dataloader = model.train_dataloader - self.get_test_dataloaders = model.test_dataloader - self.get_val_dataloaders = model.val_dataloader - # call warnings from proc zero only which triggers dataloaders - # if those have to download data it will only happen on proc 0 - if self.proc_rank == 0: - on_ddp = self.use_ddp or self.use_ddp2 - if on_ddp and not isinstance(self.get_train_dataloader().sampler, DistributedSampler): - msg = """ - You're using multiple gpus and multiple nodes without using a DistributedSampler - to assign a subset of your data to each process. To silence this warning, pass a - DistributedSampler to your DataLoader. - - ie: this: - dataset = myDataset() - dataloader = Dataloader(dataset) - - becomes: - dataset = myDataset() - dist_sampler = torch.utils.data.distributed.DistributedSampler(dataset) - dataloader = Dataloader(dataset, sampler=dist_sampler) - - If you want each process to load the full dataset, ignore this warning. - """ - warnings.warn(msg) - - if on_ddp and self.get_val_dataloaders() is not None: - for dataloader in self.get_val_dataloaders(): - if not isinstance(dataloader.sampler, DistributedSampler): - msg = """ - Your val_dataloader(s) don't use DistributedSampler. - - You're using multiple gpus and multiple nodes without using a - DistributedSampler to assign a subset of your data to each process. - To silence this warning, pass a DistributedSampler to your DataLoader. - - ie: this: - dataset = myDataset() - dataloader = Dataloader(dataset) - - becomes: - dataset = myDataset() - dist_sampler = torch.utils.data.distributed.DistributedSampler(dataset) - dataloader = Dataloader(dataset, sampler=dist_sampler) - - If you want each process to load the full dataset, ignore this warning. - """ - warnings.warn(msg) - break - - if on_ddp and self.get_test_dataloaders() is not None: - for dataloader in self.get_test_dataloaders(): - if not isinstance(dataloader.sampler, DistributedSampler): - msg = """ - Your test_dataloader(s) don't use DistributedSampler. - - You're using multiple gpus and multiple nodes without using a - DistributedSampler to assign a subset of your data to each process. - To silence this warning, pass a DistributedSampler to your DataLoader. - - ie: this: - dataset = myDataset() - dataloader = Dataloader(dataset) - - becomes: - dataset = myDataset() - dist_sampler = torch.utils.data.distributed.DistributedSampler(dataset) - dataloader = Dataloader(dataset, sampler=dist_sampler) - - If you want each process to load the full dataset, ignore this warning. - """ - warnings.warn(msg) - break + self.init_train_dataloader(model) + self.init_test_dataloader(model) + self.init_val_dataloader(model) if self.use_ddp or self.use_ddp2: # wait for all processes to catch up @@ -147,7 +174,7 @@ class TrainerDataLoadingMixin(object): Trainer(val_check_interval) must be an int. An int k specifies checking validation every k training batches ''' - raise MisconfigurationException('when using ') + raise MisconfigurationException(m) def determine_data_use_amount(self, train_percent_check, val_percent_check, test_percent_check, overfit_pct): diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index 19f31034..cbe9af37 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -135,6 +135,7 @@ class Trainer(TrainerIOMixin, self.min_nb_epochs = min_nb_epochs self.nb_sanity_val_steps = nb_sanity_val_steps self.print_nan_grads = print_nan_grads + self.shown_warnings = set() self.fast_dev_run = fast_dev_run if self.fast_dev_run: @@ -410,9 +411,6 @@ class Trainer(TrainerIOMixin, # transfer data loaders from model self.get_dataloaders(ref_model) - # init training constants - self.layout_bookeeping() - # print model summary if self.proc_rank == 0 and self.weights_summary is not None: if self.weights_summary in ['full', 'top']: