mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-11 12:31:23 +08:00
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:
co-authored by
William Falcon
parent
b9b5a93f0f
commit
ffd6e693de
@@ -2,7 +2,7 @@
|
||||
|
||||
import torch
|
||||
|
||||
from .base import LightningTestModelBase
|
||||
from .base import LightningTestModelBase, LightningTestModelBaseWithoutDataloader
|
||||
from .mixins import (
|
||||
LightningValidationStepMixin,
|
||||
LightningValidationMixin,
|
||||
|
||||
+14
-6
@@ -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
|
||||
|
||||
@@ -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__])
|
||||
|
||||
Reference in New Issue
Block a user