new way of passing dataloaders (#759)

* new way of passing dataloaders

* fixed docs

* fixed codestyle to follow flake8

* allow val/test be list of dataloaders and smarter checking

* added test

* fix flake error

* fix linking to new test model

* split into multiple test

* fix naming and typo

* minor documentation changes

* remove random file

* Update trainer.py

* Update trainer.py

* Update trainer.py

* Update trainer.py

* Update trainer.py

* Update trainer.py

* better error/warning message

* final adjustments

* update CHANGELOG.md

Co-authored-by: William Falcon <waf2107@columbia.edu>
This commit is contained in:
Nicki Skafte
2020-02-19 06:00:08 -05:00
committed by GitHub
co-authored by William Falcon
parent b9b5a93f0f
commit ffd6e693de
7 changed files with 267 additions and 13 deletions
+1 -1
View File
@@ -2,7 +2,7 @@
import torch
from .base import LightningTestModelBase
from .base import LightningTestModelBase, LightningTestModelBaseWithoutDataloader
from .mixins import (
LightningValidationStepMixin,
LightningValidationMixin,
+14 -6
View File
@@ -36,7 +36,7 @@ class TestingMNIST(MNIST):
self.targets = self.targets[:num_samples]
class LightningTestModelBase(LightningModule):
class TestModelBase(LightningModule):
"""
Base LightningModule for testing. Implements only the required
interface
@@ -48,7 +48,7 @@ class LightningTestModelBase(LightningModule):
:param hparams:
"""
# init superclass
super(LightningTestModelBase, self).__init__()
super(TestModelBase, self).__init__()
self.hparams = hparams
self.batch_size = hparams.batch_size
@@ -178,10 +178,6 @@ class LightningTestModelBase(LightningModule):
return loader
@data_loader
def train_dataloader(self):
return self._dataloader(train=True)
@staticmethod
def add_model_specific_args(parent_parser, root_dir): # pragma: no cover
"""
@@ -218,3 +214,15 @@ class LightningTestModelBase(LightningModule):
options=[32, 64, 128, 256], tunable=False,
help='batch size will be divided over all gpus being used across all nodes')
return parser
class LightningTestModelBase(TestModelBase):
""" with pre-defined train dataloader """
@data_loader
def train_dataloader(self):
return self._dataloader(train=True)
class LightningTestModelBaseWithoutDataloader(TestModelBase):
""" without pre-defined train dataloader """
pass
+161
View File
@@ -13,6 +13,7 @@ from pytorch_lightning.callbacks import (
from tests.models import (
LightningTestModel,
LightningTestModelBase,
LightningTestModelBaseWithoutDataloader,
LightningValidationStepMixin,
LightningValidationMultipleDataloadersMixin,
LightningTestMultipleDataloadersMixin,
@@ -449,6 +450,165 @@ def test_multiple_test_dataloader(tmpdir):
trainer.test()
def test_train_dataloaders_passed_to_fit(tmpdir):
""" Verify that train dataloader can be passed to fit """
tutils.reset_seed()
class CurrentTestModel(
LightningTestModelBaseWithoutDataloader
):
pass
hparams = tutils.get_hparams()
# logger file to get meta
trainer_options = dict(
default_save_path=tmpdir,
max_epochs=1,
val_percent_check=0.1,
train_percent_check=0.2
)
# only train passed to fit
model = CurrentTestModel(hparams)
trainer = Trainer(**trainer_options)
fit_options = dict(train_dataloader=model._dataloader(train=True))
results = trainer.fit(model, **fit_options)
def test_train_val_dataloaders_passed_to_fit(tmpdir):
""" Verify that train & val dataloader can be passed to fit """
tutils.reset_seed()
class CurrentTestModel(
LightningTestModelBaseWithoutDataloader
):
pass
hparams = tutils.get_hparams()
# logger file to get meta
trainer_options = dict(
default_save_path=tmpdir,
max_epochs=1,
val_percent_check=0.1,
train_percent_check=0.2
)
# train, val passed to fit
model = CurrentTestModel(hparams)
trainer = Trainer(**trainer_options)
fit_options = dict(train_dataloader=model._dataloader(train=True),
val_dataloader=model._dataloader(train=False))
results = trainer.fit(model, **fit_options)
assert len(trainer.get_val_dataloaders()) == 1, \
f'`val_dataloaders` not initiated properly, got {trainer.get_val_dataloaders()}'
def test_all_dataloaders_passed_to_fit(tmpdir):
""" Verify train, val & test dataloader can be passed to fit """
tutils.reset_seed()
class CurrentTestModel(
LightningTestModelBaseWithoutDataloader
):
pass
hparams = tutils.get_hparams()
# logger file to get meta
trainer_options = dict(
default_save_path=tmpdir,
max_epochs=1,
val_percent_check=0.1,
train_percent_check=0.2
)
# train, val and test passed to fit
model = CurrentTestModel(hparams)
trainer = Trainer(**trainer_options)
fit_options = dict(train_dataloader=model._dataloader(train=True),
val_dataloader=model._dataloader(train=False),
test_dataloader=model._dataloader(train=False))
results = trainer.fit(model, **fit_options)
assert len(trainer.get_val_dataloaders()) == 1, \
f'`val_dataloaders` not initiated properly, got {trainer.get_val_dataloaders()}'
assert len(trainer.get_test_dataloaders()) == 1, \
f'`test_dataloaders` not initiated properly, got {trainer.get_test_dataloaders()}'
def test_multiple_dataloaders_passed_to_fit(tmpdir):
""" Verify that multiple val & test dataloaders can be passed to fit """
tutils.reset_seed()
class CurrentTestModel(
LightningTestModelBaseWithoutDataloader
):
pass
hparams = tutils.get_hparams()
# logger file to get meta
trainer_options = dict(
default_save_path=tmpdir,
max_epochs=1,
val_percent_check=0.1,
train_percent_check=0.2
)
# train, multiple val and multiple test passed to fit
model = CurrentTestModel(hparams)
trainer = Trainer(**trainer_options)
fit_options = dict(train_dataloader=model._dataloader(train=True),
val_dataloader=[model._dataloader(train=False),
model._dataloader(train=False)],
test_dataloader=[model._dataloader(train=False),
model._dataloader(train=False)])
results = trainer.fit(model, **fit_options)
assert len(trainer.get_val_dataloaders()) == 2, \
f'Multiple `val_dataloaders` not initiated properly, got {trainer.get_val_dataloaders()}'
assert len(trainer.get_test_dataloaders()) == 2, \
f'Multiple `test_dataloaders` not initiated properly, got {trainer.get_test_dataloaders()}'
def test_mixing_of_dataloader_options(tmpdir):
"""Verify that dataloaders can be passed to fit"""
tutils.reset_seed()
class CurrentTestModel(
LightningTestModelBase
):
pass
hparams = tutils.get_hparams()
model = CurrentTestModel(hparams)
# logger file to get meta
trainer_options = dict(
default_save_path=tmpdir,
max_epochs=1,
val_percent_check=0.1,
train_percent_check=0.2
)
# fit model
trainer = Trainer(**trainer_options)
fit_options = dict(val_dataloader=model._dataloader(train=False))
results = trainer.fit(model, **fit_options)
# fit model
trainer = Trainer(**trainer_options)
fit_options = dict(val_dataloader=model._dataloader(train=False),
test_dataloader=model._dataloader(train=False))
results = trainer.fit(model, **fit_options)
assert len(trainer.get_val_dataloaders()) == 1, \
f'`val_dataloaders` not initiated properly, got {trainer.get_val_dataloaders()}'
assert len(trainer.get_test_dataloaders()) == 1, \
f'`test_dataloaders` not initiated properly, got {trainer.get_test_dataloaders()}'
def _init_steps_model():
"""private method for initializing a model with 5% train epochs"""
tutils.reset_seed()
@@ -533,5 +693,6 @@ def test_trainer_min_steps_and_epochs(tmpdir):
assert trainer.global_step >= math.floor(num_train_samples * 1.5) and \
trainer.current_epoch > 0, "Model did not train for at least min_steps"
# if __name__ == '__main__':
# pytest.main([__file__])