Fixing tests (#936)

* abs import

* rename test model

* update trainer

* revert test_step check

* move tags

* fix test_step

* clean tests

* fix template

* update dataset path

* fix parent order
This commit is contained in:
Jirka Borovec
2020-02-25 13:06:24 -05:00
committed by GitHub
parent 20d15c8023
commit 5dd2afeab1
15 changed files with 264 additions and 209 deletions
+21 -17
View File
@@ -2,27 +2,31 @@
import torch
from .base import LightningTestModelBase, LightningTestModelBaseWithoutDataloader
from .base import TestModelBase
from .mixins import (
LightningValidationStepMixin,
LightningValidationMixin,
LightningValidationStepMultipleDataloadersMixin,
LightningValidationMultipleDataloadersMixin,
LightningTestStepMixin,
LightningTestMixin,
LightningTestStepMultipleDataloadersMixin,
LightningTestMultipleDataloadersMixin,
LightningTestFitSingleTestDataloadersMixin,
LightningTestFitMultipleTestDataloadersMixin,
LightningValStepFitSingleDataloaderMixin,
LightningValStepFitMultipleDataloadersMixin
LightEmptyTestStep,
LightValidationStepMixin,
LightValidationMixin,
LightValidationStepMultipleDataloadersMixin,
LightValidationMultipleDataloadersMixin,
LightTestStepMixin,
LightTestMixin,
LightTestStepMultipleDataloadersMixin,
LightTestMultipleDataloadersMixin,
LightTestFitSingleTestDataloadersMixin,
LightTestFitMultipleTestDataloadersMixin,
LightValStepFitSingleDataloaderMixin,
LightValStepFitMultipleDataloadersMixin,
LightTrainDataloader,
LightTestDataloader,
)
class LightningTestModel(LightningValidationMixin, LightningTestMixin, LightningTestModelBase):
"""
Most common test case. Validation and test dataloaders.
"""
class LightningTestModel(LightTrainDataloader,
LightValidationMixin,
LightTestMixin,
TestModelBase):
"""Most common test case. Validation and test dataloaders."""
def on_training_metrics(self, logs):
logs['some_tensor_to_test'] = torch.rand(1)
+5 -22
View File
@@ -24,7 +24,7 @@ class TestingMNIST(MNIST):
def __init__(self, root, train=True, transform=None, target_transform=None,
download=False, num_samples=8000):
super(TestingMNIST, self).__init__(
super().__init__(
root,
train=train,
transform=transform,
@@ -48,7 +48,7 @@ class TestModelBase(LightningModule):
:param hparams:
"""
# init superclass
super(TestModelBase, self).__init__()
super().__init__()
self.hparams = hparams
self.batch_size = hparams.batch_size
@@ -87,7 +87,6 @@ class TestModelBase(LightningModule):
:param x:
:return:
"""
x = self.c_d1(x)
x = torch.tanh(x)
x = self.c_d1_bn(x)
@@ -153,10 +152,8 @@ class TestModelBase(LightningModule):
def prepare_data(self):
transform = transforms.Compose([transforms.ToTensor(),
transforms.Normalize((0.5,), (1.0,))])
dataset = TestingMNIST(root=self.hparams.data_root, train=True,
transform=transform, download=True, num_samples=2000)
dataset = TestingMNIST(root=self.hparams.data_root, train=False,
transform=transform, download=True, num_samples=2000)
_ = TestingMNIST(root=self.hparams.data_root, train=True,
transform=transform, download=True, num_samples=2000)
def _dataloader(self, train):
# init data generators
@@ -194,31 +191,17 @@ class TestModelBase(LightningModule):
parser.add_argument('--out_features', default=10, type=int)
# use 500 for CPU, 50000 for GPU to see speed difference
parser.add_argument('--hidden_dim', default=50000, type=int)
# data
parser.add_argument('--data_root', default=os.path.join(root_dir, 'mnist'), type=str)
# training params (opt)
parser.opt_list('--learning_rate', default=0.001 * 8, type=float,
options=[0.0001, 0.0005, 0.001, 0.005],
tunable=False)
parser.opt_list('--optimizer_name', default='adam', type=str,
options=['adam'], tunable=False)
# if using 2 nodes with 4 gpus each the batch size here
# (256) will be 256 / (2*8) = 16 per gpu
parser.opt_list('--batch_size', default=256 * 8, type=int,
options=[32, 64, 128, 256], tunable=False,
help='batch size will be divided over all gpus being used across all nodes')
help='batch size will be divided over all GPUs being used across all nodes')
return parser
class LightningTestModelBase(TestModelBase):
""" with pre-defined train dataloader """
def train_dataloader(self):
return self._dataloader(train=True)
class LightningTestModelBaseWithoutDataloader(TestModelBase):
""" without pre-defined train dataloader """
pass
+61 -24
View File
@@ -5,7 +5,7 @@ import torch
from pytorch_lightning.core.decorators import data_loader
class LightningValidationStepMixin:
class LightValidationStepMixin:
"""
Add val_dataloader and validation_step methods for the case
when val_dataloader returns a single dataloader
@@ -14,7 +14,7 @@ class LightningValidationStepMixin:
def val_dataloader(self):
return self._dataloader(train=False)
def validation_step(self, batch, batch_idx):
def validation_step(self, batch, batch_idx, *args, **kwargs):
"""
Lightning calls this inside the validation loop
:param batch:
@@ -58,7 +58,7 @@ class LightningValidationStepMixin:
return output
class LightningValidationMixin(LightningValidationStepMixin):
class LightValidationMixin(LightValidationStepMixin):
"""
Add val_dataloader, validation_step, and validation_end methods for the case
when val_dataloader returns a single dataloader
@@ -76,7 +76,7 @@ class LightningValidationMixin(LightningValidationStepMixin):
val_loss_mean = 0
val_acc_mean = 0
for output in outputs:
val_loss = output['val_loss']
val_loss = _get_output_metric(output, 'val_loss')
# reduce manually when using dp
if self.trainer.use_dp or self.trainer.use_ddp2:
@@ -84,7 +84,7 @@ class LightningValidationMixin(LightningValidationStepMixin):
val_loss_mean += val_loss
# reduce manually when using dp
val_acc = output['val_acc']
val_acc = _get_output_metric(output, 'val_acc')
if self.trainer.use_dp or self.trainer.use_ddp2:
val_acc = torch.mean(val_acc)
@@ -98,7 +98,7 @@ class LightningValidationMixin(LightningValidationStepMixin):
return results
class LightningValidationStepMultipleDataloadersMixin:
class LightValidationStepMultipleDataloadersMixin:
"""
Add val_dataloader and validation_step methods for the case
when val_dataloader returns multiple dataloaders
@@ -107,7 +107,7 @@ class LightningValidationStepMultipleDataloadersMixin:
def val_dataloader(self):
return [self._dataloader(train=False), self._dataloader(train=False)]
def validation_step(self, batch, batch_idx, dataloader_idx):
def validation_step(self, batch, batch_idx, dataloader_idx, **kwargs):
"""
Lightning calls this inside the validation loop
:param batch:
@@ -157,7 +157,7 @@ class LightningValidationStepMultipleDataloadersMixin:
return output
class LightningValidationMultipleDataloadersMixin(LightningValidationStepMultipleDataloadersMixin):
class LightValidationMultipleDataloadersMixin(LightValidationStepMultipleDataloadersMixin):
"""
Add val_dataloader, validation_step, and validation_end methods for the case
when val_dataloader returns multiple dataloaders
@@ -200,12 +200,31 @@ class LightningValidationMultipleDataloadersMixin(LightningValidationStepMultipl
return result
class LightningTestStepMixin:
class LightTrainDataloader:
"""Simple train dataloader."""
def train_dataloader(self):
return self._dataloader(train=True)
class LightTestDataloader:
"""Simple test dataloader."""
def test_dataloader(self):
return self._dataloader(train=False)
def test_step(self, batch, batch_idx):
class LightEmptyTestStep:
"""Empty test step."""
def test_step(self, *args, **kwargs):
return dict()
class LightTestStepMixin(LightTestDataloader):
"""Test step mixin."""
def test_step(self, batch, batch_idx, *args, **kwargs):
"""
Lightning calls this inside the validation loop
:param batch:
@@ -249,7 +268,9 @@ class LightningTestStepMixin:
return output
class LightningTestMixin(LightningTestStepMixin):
class LightTestMixin(LightTestStepMixin):
"""Ritch test mixin."""
def test_end(self, outputs):
"""
Called at the end of validation to aggregate outputs
@@ -262,7 +283,7 @@ class LightningTestMixin(LightningTestStepMixin):
test_loss_mean = 0
test_acc_mean = 0
for output in outputs:
test_loss = output['test_loss']
test_loss = _get_output_metric(output, 'test_loss')
# reduce manually when using dp
if self.trainer.use_dp:
@@ -270,7 +291,7 @@ class LightningTestMixin(LightningTestStepMixin):
test_loss_mean += test_loss
# reduce manually when using dp
test_acc = output['test_acc']
test_acc = _get_output_metric(output, 'test_acc')
if self.trainer.use_dp:
test_acc = torch.mean(test_acc)
@@ -284,12 +305,13 @@ class LightningTestMixin(LightningTestStepMixin):
return result
class LightningTestStepMultipleDataloadersMixin:
class LightTestStepMultipleDataloadersMixin:
"""Test step multiple dataloaders mixin."""
def test_dataloader(self):
return [self._dataloader(train=False), self._dataloader(train=False)]
def test_step(self, batch, batch_idx, dataloader_idx):
def test_step(self, batch, batch_idx, dataloader_idx, **kwargs):
"""
Lightning calls this inside the validation loop
:param batch:
@@ -339,8 +361,10 @@ class LightningTestStepMultipleDataloadersMixin:
return output
class LightningTestFitSingleTestDataloadersMixin:
def test_step(self, batch, batch_idx):
class LightTestFitSingleTestDataloadersMixin:
"""Test fit single test dataloaders mixin."""
def test_step(self, batch, batch_idx, *args, **kwargs):
"""
Lightning calls this inside the validation loop
:param batch:
@@ -384,8 +408,10 @@ class LightningTestFitSingleTestDataloadersMixin:
return output
class LightningTestFitMultipleTestDataloadersMixin:
def test_step(self, batch, batch_idx, dataloader_idx):
class LightTestFitMultipleTestDataloadersMixin:
"""Test fit multiple test dataloaders mixin."""
def test_step(self, batch, batch_idx, dataloader_idx, **kwargs):
"""
Lightning calls this inside the validation loop
:param batch:
@@ -435,8 +461,9 @@ class LightningTestFitMultipleTestDataloadersMixin:
return output
class LightningValStepFitSingleDataloaderMixin:
def validation_step(self, batch, batch_idx):
class LightValStepFitSingleDataloaderMixin:
def validation_step(self, batch, batch_idx, *args, **kwargs):
"""
Lightning calls this inside the validation loop
:param batch:
@@ -480,8 +507,9 @@ class LightningValStepFitSingleDataloaderMixin:
return output
class LightningValStepFitMultipleDataloadersMixin:
def validation_step(self, batch, batch_idx, dataloader_idx):
class LightValStepFitMultipleDataloadersMixin:
def validation_step(self, batch, batch_idx, dataloader_idx, **kwargs):
"""
Lightning calls this inside the validation loop
:param batch:
@@ -531,7 +559,8 @@ class LightningValStepFitMultipleDataloadersMixin:
return output
class LightningTestMultipleDataloadersMixin(LightningTestStepMultipleDataloadersMixin):
class LightTestMultipleDataloadersMixin(LightTestStepMultipleDataloadersMixin):
def test_end(self, outputs):
"""
Called at the end of validation to aggregate outputs
@@ -567,3 +596,11 @@ class LightningTestMultipleDataloadersMixin(LightningTestStepMultipleDataloaders
tqdm_dict = {'test_loss': test_loss_mean.item(), 'test_acc': test_acc_mean.item()}
result = {'progress_bar': tqdm_dict}
return result
def _get_output_metric(output, name):
if isinstance(output, dict):
val = output[name]
else: # if it is 2level deep -> per dataloader and per batch
val = sum(out[name] for out in output) / len(output)
return val
+2 -2
View File
@@ -90,7 +90,7 @@ def run_model_test(trainer_options, model, on_gpu=True):
def get_hparams(continue_training=False, hpc_exp_number=0):
root_dir = os.path.dirname(os.path.realpath(__file__))
tests_dir = os.path.dirname(os.path.dirname(os.path.realpath(__file__)))
args = {
'drop_prob': 0.2,
@@ -98,7 +98,7 @@ def get_hparams(continue_training=False, hpc_exp_number=0):
'in_features': 28 * 28,
'learning_rate': 0.001 * 8,
'optimizer_name': 'adam',
'data_root': os.path.join(root_dir, 'mnist'),
'data_root': os.path.join(tests_dir, 'datasets'),
'out_features': 10,
'hidden_dim': 1000,
}
+5 -4
View File
@@ -8,9 +8,10 @@ from pytorch_lightning.callbacks import (
EarlyStopping,
)
from tests.models import (
TestModelBase,
LightTrainDataloader,
LightningTestModel,
LightningTestModelBase,
LightningTestMixin,
LightTestMixin,
)
@@ -121,7 +122,7 @@ def test_running_test_without_val(tmpdir):
"""Verify `test()` works on a model with no `val_loader`."""
tutils.reset_seed()
class CurrentTestModel(LightningTestMixin, LightningTestModelBase):
class CurrentTestModel(LightTrainDataloader, LightTestMixin, TestModelBase):
pass
hparams = tutils.get_hparams()
@@ -281,7 +282,7 @@ def test_tbptt_cpu_model(tmpdir):
def __len__(self):
return 1
class BpttTestModel(LightningTestModelBase):
class BpttTestModel(LightTrainDataloader, TestModelBase):
def __init__(self, hparams):
super().__init__(hparams)
self.test_hidden = None
+68 -50
View File
@@ -11,19 +11,22 @@ from pytorch_lightning.callbacks import (
ModelCheckpoint,
)
from tests.models import (
TestModelBase,
LightningTestModel,
LightningTestModelBase,
LightningTestModelBaseWithoutDataloader,
LightningValidationStepMixin,
LightningValidationMultipleDataloadersMixin,
LightningTestMultipleDataloadersMixin,
LightningTestFitSingleTestDataloadersMixin,
LightningTestFitMultipleTestDataloadersMixin,
LightningValStepFitMultipleDataloadersMixin,
LightningValStepFitSingleDataloaderMixin
LightEmptyTestStep,
LightValidationStepMixin,
LightValidationMultipleDataloadersMixin,
LightTestMultipleDataloadersMixin,
LightTestFitSingleTestDataloadersMixin,
LightTestFitMultipleTestDataloadersMixin,
LightValStepFitMultipleDataloadersMixin,
LightValStepFitSingleDataloaderMixin,
LightTrainDataloader,
LightTestDataloader,
)
from pytorch_lightning.core.lightning import load_hparams_from_tags_csv
from pytorch_lightning.trainer.logging import TrainerLoggingMixin
from pytorch_lightning.utilities.debugging import MisconfigurationException
def test_no_val_module(tmpdir):
@@ -32,7 +35,7 @@ def test_no_val_module(tmpdir):
hparams = tutils.get_hparams()
class CurrentTestModel(LightningTestModelBase):
class CurrentTestModel(LightTrainDataloader, TestModelBase):
pass
model = CurrentTestModel(hparams)
@@ -69,7 +72,7 @@ def test_no_val_end_module(tmpdir):
"""Tests use case where trainer saves the model, and user loads it from tags independently."""
tutils.reset_seed()
class CurrentTestModel(LightningValidationStepMixin, LightningTestModelBase):
class CurrentTestModel(LightTrainDataloader, LightValidationStepMixin, TestModelBase):
pass
hparams = tutils.get_hparams()
@@ -385,8 +388,9 @@ def test_multiple_val_dataloader(tmpdir):
tutils.reset_seed()
class CurrentTestModel(
LightningValidationMultipleDataloadersMixin,
LightningTestModelBase
LightTrainDataloader,
LightValidationMultipleDataloadersMixin,
TestModelBase,
):
pass
@@ -490,8 +494,10 @@ def test_multiple_test_dataloader(tmpdir):
tutils.reset_seed()
class CurrentTestModel(
LightningTestMultipleDataloadersMixin,
LightningTestModelBase
LightTrainDataloader,
LightTestMultipleDataloadersMixin,
LightEmptyTestStep,
TestModelBase,
):
pass
@@ -508,8 +514,7 @@ def test_multiple_test_dataloader(tmpdir):
# fit model
trainer = Trainer(**trainer_options)
result = trainer.fit(model)
trainer.fit(model)
trainer.test()
# verify there are 2 val loaders
@@ -528,9 +533,7 @@ def test_train_dataloaders_passed_to_fit(tmpdir):
""" Verify that train dataloader can be passed to fit """
tutils.reset_seed()
class CurrentTestModel(
LightningTestModelBaseWithoutDataloader,
):
class CurrentTestModel(LightTrainDataloader, TestModelBase):
pass
hparams = tutils.get_hparams()
@@ -555,8 +558,9 @@ def test_train_val_dataloaders_passed_to_fit(tmpdir):
tutils.reset_seed()
class CurrentTestModel(
LightningValStepFitSingleDataloaderMixin,
LightningTestModelBaseWithoutDataloader,
LightTrainDataloader,
LightValStepFitSingleDataloaderMixin,
TestModelBase,
):
pass
@@ -586,9 +590,11 @@ def test_all_dataloaders_passed_to_fit(tmpdir):
tutils.reset_seed()
class CurrentTestModel(
LightningValStepFitSingleDataloaderMixin,
LightningTestFitSingleTestDataloadersMixin,
LightningTestModelBaseWithoutDataloader,
LightTrainDataloader,
LightValStepFitSingleDataloaderMixin,
LightTestFitSingleTestDataloadersMixin,
LightEmptyTestStep,
TestModelBase,
):
pass
@@ -624,9 +630,9 @@ def test_multiple_dataloaders_passed_to_fit(tmpdir):
tutils.reset_seed()
class CurrentTestModel(
LightningValStepFitMultipleDataloadersMixin,
LightningTestFitMultipleTestDataloadersMixin,
LightningTestModelBaseWithoutDataloader,
LightningTestModel,
LightValStepFitMultipleDataloadersMixin,
LightTestFitMultipleTestDataloadersMixin,
):
pass
@@ -662,9 +668,10 @@ def test_mixing_of_dataloader_options(tmpdir):
tutils.reset_seed()
class CurrentTestModel(
LightningValStepFitSingleDataloaderMixin,
LightningTestFitSingleTestDataloadersMixin,
LightningTestModelBase,
LightTrainDataloader,
LightValStepFitSingleDataloaderMixin,
LightTestFitSingleTestDataloadersMixin,
TestModelBase,
):
pass
@@ -688,7 +695,7 @@ def test_mixing_of_dataloader_options(tmpdir):
trainer = Trainer(**trainer_options)
fit_options = dict(val_dataloaders=model._dataloader(train=False),
test_dataloaders=model._dataloader(train=False))
results = trainer.fit(model, **fit_options)
_ = trainer.fit(model, **fit_options)
trainer.test()
assert len(trainer.val_dataloaders) == 1, \
@@ -719,6 +726,7 @@ def test_trainer_max_steps_and_epochs(tmpdir):
# define less train steps than epochs
trainer_options.update(dict(
default_save_path=tmpdir,
max_epochs=5,
max_steps=num_train_samples + 10
))
@@ -732,8 +740,10 @@ def test_trainer_max_steps_and_epochs(tmpdir):
assert trainer.global_step == trainer.max_steps, "Model did not stop at max_steps"
# define less train epochs than steps
trainer_options['max_epochs'] = 2
trainer_options['max_steps'] = trainer_options['max_epochs'] * 2 * num_train_samples
trainer_options.update(dict(
max_epochs=2,
max_steps=trainer_options['max_epochs'] * 2 * num_train_samples
))
# fit model
trainer = Trainer(**trainer_options)
@@ -741,8 +751,8 @@ def test_trainer_max_steps_and_epochs(tmpdir):
assert result == 1, "Training did not complete"
# check training stopped at max_epochs
assert trainer.global_step == num_train_samples * trainer.max_nb_epochs \
and trainer.current_epoch == trainer.max_nb_epochs - 1, "Model did not stop at max_epochs"
assert trainer.global_step == num_train_samples * trainer.max_epochs \
and trainer.current_epoch == trainer.max_epochs - 1, "Model did not stop at max_epochs"
def test_trainer_min_steps_and_epochs(tmpdir):
@@ -750,12 +760,13 @@ def test_trainer_min_steps_and_epochs(tmpdir):
model, trainer_options, num_train_samples = _init_steps_model()
# define callback for stopping the model and default epochs
trainer_options.update({
'early_stop_callback': EarlyStopping(monitor='val_loss', min_delta=1.0),
'val_check_interval': 20,
'min_epochs': 1,
'max_epochs': 10
})
trainer_options.update(dict(
default_save_path=tmpdir,
early_stop_callback=EarlyStopping(monitor='val_loss', min_delta=1.0),
val_check_interval=20,
min_epochs=1,
max_epochs=10
))
# define less min steps than 1 epoch
trainer_options['min_steps'] = math.floor(num_train_samples / 2)
@@ -784,22 +795,29 @@ def test_trainer_min_steps_and_epochs(tmpdir):
def test_testpass_overrides(tmpdir):
hparams = tutils.get_hparams()
from pytorch_lightning.utilities.debugging import MisconfigurationException
class TestModelNoEnd(LightningTestModelBase):
def test_step(self, *args, **kwargs):
class LocalModel(LightTrainDataloader, TestModelBase):
pass
class LocalModelNoEnd(LightTrainDataloader, LightTestDataloader, LightEmptyTestStep, TestModelBase):
pass
class LocalModelNoStep(LightTrainDataloader, TestModelBase):
def test_end(self, outputs):
return {}
def test_dataloader(self):
return self.train_dataloader()
# Misconfig when neither test_step or test_end is implemented
with pytest.raises(MisconfigurationException):
model = LightningTestModelBase(hparams)
model = LocalModel(hparams)
Trainer().test(model)
# Misconfig when neither test_step or test_end is implemented
with pytest.raises(MisconfigurationException):
model = LocalModelNoStep(hparams)
Trainer().test(model)
# No exceptions when one or both of test_step or test_end are implemented
model = TestModelNoEnd(hparams)
model = LocalModelNoEnd(hparams)
Trainer().test(model)
model = LightningTestModel(hparams)