Use getter instead of python property for the dataloaders (#275)

* Use getter instead of python property for the dataloaders

* Fix lint

* Update trainer.py
This commit is contained in:
Hendrik Schröter
2019-10-04 15:35:02 -04:00
committed by William Falcon
parent 32e74b8f36
commit 36f0b5bbd0
3 changed files with 57 additions and 53 deletions
+8 -3
View File
@@ -10,13 +10,18 @@ def data_loader(fn):
attr_name = '_lazy_' + fn.__name__
@property
def _data_loader(self):
def _get_data_loader(self):
try:
value = getattr(self, attr_name)
except AttributeError:
try:
value = fn(self) # Lazy evaluation, done only once.
if (
value is not None and
not isinstance(value, list) and
fn.__name__ in['test_dataloader', 'val_dataloader']
):
value = [value]
except AttributeError as e:
# Guard against AttributeError suppression. (Issue #142)
traceback.print_exc()
@@ -25,4 +30,4 @@ def data_loader(fn):
setattr(self, attr_name, value) # Memoize evaluation.
return value
return _data_loader
return _get_data_loader
+30 -37
View File
@@ -24,6 +24,7 @@ from pytorch_lightning.utilities.debugging import MisconfigurationException
import pdb
from pytorch_lightning.trainer import ignored_warnings
try:
from apex import amp
APEX_AVAILABLE = True
@@ -141,9 +142,9 @@ class Trainer(TrainerIO):
self.nb_val_batches = 0
self.nb_training_batches = 0
self.nb_test_batches = 0
self.train_dataloader = None
self.test_dataloader = None
self.val_dataloader = None
self.get_train_dataloader = None
self.get_test_dataloaders = None
self.get_val_dataloaders = None
# training state
self.model = None
@@ -450,19 +451,21 @@ class Trainer(TrainerIO):
def __layout_bookeeping(self):
# determine number of training batches
self.nb_training_batches = len(self.train_dataloader)
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.val_dataloader is not None:
self.nb_val_batches = sum(len(dataloader) for dataloader in self.val_dataloader)
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.test_dataloader is not None:
self.nb_test_batches = sum(len(dataloader) for dataloader in self.test_dataloader)
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)
@@ -481,10 +484,10 @@ class Trainer(TrainerIO):
# make dataloader_idx arg in validation_step optional
args = [batch, batch_idx]
if test and len(self.test_dataloader) > 1:
if test and len(self.get_test_dataloaders()) > 1:
args.append(dataloader_idx)
elif not test and len(self.val_dataloader) > 1:
elif not test and len(self.get_val_dataloaders()) > 1:
args.append(dataloader_idx)
# handle DP, DDP forward
@@ -530,9 +533,9 @@ class Trainer(TrainerIO):
outputs = []
# run training
for dataloader_idx, dl in enumerate(dataloaders):
for dataloader_idx, dataloader in enumerate(dataloaders):
dl_outputs = []
for batch_idx, batch in enumerate(dl):
for batch_idx, batch in enumerate(dataloader):
if batch is None: # pragma: no cover
continue
@@ -582,21 +585,11 @@ class Trainer(TrainerIO):
:param model:
:return:
"""
self.get_train_dataloader = model.train_dataloader
self.get_test_dataloaders = model.test_dataloader
self.get_val_dataloaders = model.val_dataloader
self.train_dataloader = model.train_dataloader
self.test_dataloader = model.test_dataloader
self.val_dataloader = model.val_dataloader
# handle returning an actual dataloader instead of a list of loaders
have_test_loaders = self.test_dataloader is not None
if have_test_loaders and not isinstance(self.test_dataloader, list):
self.test_dataloader = [self.test_dataloader]
have_val_loaders = self.val_dataloader is not None
if have_val_loaders and not isinstance(self.val_dataloader, list):
self.val_dataloader = [self.val_dataloader]
if self.use_ddp and not isinstance(self.train_dataloader.sampler, DistributedSampler):
if self.use_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
@@ -615,8 +608,8 @@ class Trainer(TrainerIO):
"""
warnings.warn(msg)
if self.use_ddp and self.val_dataloader is not None:
for dataloader in self.val_dataloader:
if self.use_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.
@@ -638,8 +631,8 @@ class Trainer(TrainerIO):
warnings.warn(msg)
break
if self.use_ddp and self.test_dataloader is not None:
for dataloader in self.test_dataloader:
if self.use_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.
@@ -954,12 +947,12 @@ class Trainer(TrainerIO):
# run tiny validation (if validation defined)
# to make sure program won't crash during val
ref_model.on_sanity_check_start()
if self.val_dataloader is not None and self.nb_sanity_val_steps > 0:
if self.get_val_dataloaders() is not None and self.nb_sanity_val_steps > 0:
# reset progress_bar limit for sanity check
if self.show_progress_bar:
self.progress_bar.reset(self.nb_sanity_val_steps)
self.evaluate(model, self.val_dataloader, self.nb_sanity_val_steps, self.testing)
self.evaluate(model, self.get_val_dataloaders(), self.nb_sanity_val_steps, self.testing)
# ---------------------------
# CORE TRAINING LOOP
@@ -970,8 +963,8 @@ class Trainer(TrainerIO):
# run all epochs
for epoch_nb in range(self.current_epoch, self.max_nb_epochs):
# set seed for distributed sampler (enables shuffling for each epoch)
if self.use_ddp and hasattr(self.train_dataloader.sampler, 'set_epoch'):
self.train_dataloader.sampler.set_epoch(epoch_nb)
if self.use_ddp and hasattr(self.get_train_dataloader().sampler, 'set_epoch'):
self.get_train_dataloader().sampler.set_epoch(epoch_nb)
# get model
model = self.__get_model()
@@ -1016,7 +1009,7 @@ class Trainer(TrainerIO):
model.on_epoch_start()
# run epoch
for batch_nb, batch in enumerate(self.train_dataloader):
for batch_nb, batch in enumerate(self.get_train_dataloader()):
self.batch_nb = batch_nb
self.global_step += 1
@@ -1325,12 +1318,12 @@ class Trainer(TrainerIO):
model.on_pre_performance_check()
# select dataloaders
dataloaders = self.val_dataloader
dataloaders = self.get_val_dataloaders()
max_batches = self.nb_val_batches
# calculate max batches to use
if test:
dataloaders = self.test_dataloader
dataloaders = self.get_test_dataloaders()
max_batches = self.nb_test_batches
# cap max batches to 1 when using fast_dev_run
+19 -13
View File
@@ -128,7 +128,8 @@ def test_dp_resume():
dp_model = new_trainer.model
dp_model.eval()
_ = [run_prediction(dataloader, dp_model, dp=True) for dataloader in trainer.val_dataloader]
for dataloader in trainer.get_train_dataloader():
run_prediction(dataloader, dp_model, dp=True)
# new model
model = LightningTestModel(hparams)
@@ -186,7 +187,7 @@ def test_running_test_pretrained_model_ddp():
new_trainer = Trainer(**trainer_options)
new_trainer.test(pretrained_model)
run_prediction(model.test_dataloader, pretrained_model)
[run_prediction(dataloader, pretrained_model) for dataloader in model.test_dataloader()]
# test we have good test accuracy
clear_save_dir()
@@ -784,7 +785,8 @@ def test_cpu_restore_training():
# if model and state loaded correctly, predictions will be good even though we
# haven't trained with the new loaded model
trainer.model.eval()
_ = [run_prediction(dataloader, trainer.model) for dataloader in trainer.val_dataloader]
for dataloader in trainer.get_val_dataloaders():
run_prediction(dataloader, trainer.model)
model.on_sanity_check_start = assert_good_acc
@@ -852,8 +854,9 @@ def test_cpu_slurm_save_load():
# predict with trained model before saving
# make a prediction
for batch in model.test_dataloader:
break
for dataloader in model.test_dataloader():
for batch in dataloader:
break
x, y = batch
x = x.view(x.size(0), -1)
@@ -968,8 +971,9 @@ def test_model_saving_loading():
assert result == 1, 'amp + ddp model failed to complete'
# make a prediction
for batch in model.test_dataloader:
break
for dataloader in model.test_dataloader():
for batch in dataloader:
break
x, y = batch
x = x.view(x.size(0), -1)
@@ -1060,7 +1064,7 @@ def test_amp_gpu_ddp_slurm_managed():
pretrained_model = load_model(logger.experiment, save_dir, True)
# test model preds
run_prediction(model.test_dataloader, pretrained_model)
[run_prediction(dataloader, pretrained_model) for dataloader in trainer.get_test_dataloaders()]
if trainer.use_ddp:
# on hpc this would work fine... but need to hack it for the purpose of the test
@@ -1287,10 +1291,11 @@ def test_multiple_val_dataloader():
assert result == 1
# verify there are 2 val loaders
assert len(trainer.val_dataloader) == 2, 'Multiple val_dataloaders not initiated properly'
assert len(trainer.get_val_dataloaders()) == 2, \
'Multiple val_dataloaders not initiated properly'
# make sure predictions are good for each val set
[run_prediction(dataloader, trainer.model) for dataloader in trainer.val_dataloader]
[run_prediction(dataloader, trainer.model) for dataloader in trainer.get_val_dataloaders()]
def test_multiple_test_dataloader():
@@ -1318,10 +1323,11 @@ def test_multiple_test_dataloader():
result = trainer.fit(model)
# verify there are 2 val loaders
assert len(trainer.test_dataloader) == 2, 'Multiple test_dataloaders not initiated properly'
assert len(trainer.get_test_dataloaders()) == 2, \
'Multiple test_dataloaders not initiated properly'
# make sure predictions are good for each test set
[run_prediction(dataloader, trainer.model) for dataloader in trainer.test_dataloader]
[run_prediction(dataloader, trainer.model) for dataloader in trainer.get_test_dataloaders()]
# run the test method
trainer.test()
@@ -1356,7 +1362,7 @@ def run_gpu_model_test(trainer_options, model, hparams, on_gpu=True):
pretrained_model = load_model(logger.experiment, save_dir, on_gpu)
# test new model accuracy
run_prediction(model.test_dataloader, pretrained_model)
[run_prediction(dataloader, pretrained_model) for dataloader in model.test_dataloader()]
if trainer.use_ddp:
# on hpc this would work fine... but need to hack it for the purpose of the test