Tests: refactor cleanup (#1744)

* wip

* cleaning

* optim imports

* -

* default hparams

* fix restore

* fix imports
This commit is contained in:
Jirka Borovec
2020-05-10 13:15:28 -04:00
committed by GitHub
parent 4970927ec8
commit 134eb61e1a
41 changed files with 187 additions and 1150 deletions
+7 -6
View File
@@ -1,4 +1,5 @@
import pytest
import tests.base.utils as tutils
from pytorch_lightning import Callback
from pytorch_lightning import Trainer, LightningModule
@@ -11,7 +12,7 @@ from pathlib import Path
def test_trainer_callback_system(tmpdir):
"""Test the callback system."""
hparams = tutils.get_default_hparams()
hparams = EvalModelTemplate.get_default_hparams()
model = EvalModelTemplate(hparams)
def _check_args(trainer, pl_module):
@@ -209,7 +210,7 @@ def test_early_stopping_no_val_step(tmpdir):
output.update({'my_train_metric': output['loss']}) # could be anything else
return output
model = CurrentModel(tutils.get_default_hparams())
model = CurrentModel()
model.validation_step = None
model.val_dataloader = None
@@ -245,7 +246,7 @@ def test_pickling(tmpdir):
def test_model_checkpoint_with_non_string_input(tmpdir, save_top_k):
""" Test that None in checkpoint callback is valid and that chkp_path is set correctly """
tutils.reset_seed()
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
checkpoint = ModelCheckpoint(filepath=None, save_top_k=save_top_k)
@@ -267,7 +268,7 @@ def test_model_checkpoint_with_non_string_input(tmpdir, save_top_k):
def test_model_checkpoint_path(tmpdir, logger_version, expected):
"""Test that "version_" prefix is only added when logger's version is an integer"""
tutils.reset_seed()
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
logger = TensorBoardLogger(str(tmpdir), version=logger_version)
trainer = Trainer(
@@ -286,7 +287,7 @@ def test_lr_logger_single_lr(tmpdir):
""" Test that learning rates are extracted and logged for single lr scheduler"""
tutils.reset_seed()
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.configure_optimizers = model.configure_optimizers__single_scheduler
lr_logger = LearningRateLogger()
@@ -311,7 +312,7 @@ def test_lr_logger_multi_lrs(tmpdir):
""" Test that learning rates are extracted and logged for multi lr schedulers """
tutils.reset_seed()
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
model.configure_optimizers = model.configure_optimizers__multiple_schedulers
lr_logger = LearningRateLogger()
+3 -3
View File
@@ -58,7 +58,7 @@ def test_progress_bar_misconfiguration():
def test_progress_bar_totals():
"""Test that the progress finishes with the correct total steps processed."""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
trainer = Trainer(
progress_bar_refresh_rate=1,
@@ -107,7 +107,7 @@ def test_progress_bar_totals():
def test_progress_bar_fast_dev_run():
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
trainer = Trainer(
fast_dev_run=True,
@@ -140,7 +140,7 @@ def test_progress_bar_fast_dev_run():
def test_progress_bar_progress_refresh(refresh_rate):
"""Test that the three progress bars get correctly updated when using different refresh rates."""
model = EvalModelTemplate(tutils.get_default_hparams())
model = EvalModelTemplate()
class CurrentProgressBar(ProgressBar):