mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
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:
committed by
William Falcon
parent
32e74b8f36
commit
36f0b5bbd0
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user