unify tests (#1940)

* unify tests

* Apply suggestions from code review

Co-authored-by: Adrian Wälchli <aedu.waelchli@gmail.com>

Co-authored-by: Adrian Wälchli <aedu.waelchli@gmail.com>
This commit is contained in:
Jirka Borovec
2020-05-27 22:45:23 -04:00
committed by GitHub
co-authored by Adrian Wälchli
parent 7c19c373ac
commit df78e84060
3 changed files with 25 additions and 46 deletions
+25
View File
@@ -96,3 +96,28 @@ def test_loggers_pickle(tmpdir, monkeypatch, logger_class):
trainer2 = pickle.loads(pkl_bytes)
trainer2.logger.log_metrics({'acc': 1.0})
@pytest.mark.parametrize("extra_params", [
pytest.param(dict(max_epochs=1, auto_scale_batch_size=True), id='Batch-size-Finder'),
pytest.param(dict(max_epochs=10, auto_lr_find=True), id='LR-Finder'),
])
def test_logger_reset_correctly(tmpdir, extra_params):
""" Test that the tuners do not alter the logger reference """
tutils.reset_seed()
model = EvalModelTemplate()
trainer = Trainer(
default_save_path=tmpdir,
**extra_params
)
logger1 = trainer.logger
trainer.fit(model)
logger2 = trainer.logger
logger3 = model.logger
assert logger1 == logger2, \
'Finder altered the logger of trainer'
assert logger2 == logger3, \
'Finder altered the logger of model'
-23
View File
@@ -198,26 +198,3 @@ def test_suggestion_with_non_finite_values(tmpdir):
assert before_lr == after_lr, \
'Learning rate was altered because of non-finite loss values'
def test_logger_reset_correctly(tmpdir):
""" Test that logger is updated correctly """
tutils.reset_seed()
hparams = EvalModelTemplate.get_default_hparams()
model = EvalModelTemplate(hparams)
trainer = Trainer(
default_save_path=tmpdir,
max_epochs=10,
auto_lr_find=True
)
logger1 = trainer.logger
trainer.fit(model)
logger2 = trainer.logger
logger3 = model.logger
assert logger1 == logger2, \
'Learning rate finder altered the logger of trainer'
assert logger2 == logger3, \
'Learning rate finder altered the logger of model'
-23
View File
@@ -128,26 +128,3 @@ def test_error_on_dataloader_passed_to_fit(tmpdir):
with pytest.raises(MisconfigurationException):
trainer.fit(model, **fit_options)
def test_logger_reset_correctly(tmpdir):
""" Test that logger is updated correctly """
tutils.reset_seed()
hparams = EvalModelTemplate.get_default_hparams()
model = EvalModelTemplate(hparams)
trainer = Trainer(
default_save_path=tmpdir,
max_epochs=1,
auto_scale_batch_size=True
)
logger1 = trainer.logger
trainer.fit(model)
logger2 = trainer.logger
logger3 = model.logger
assert logger1 == logger2, \
'Batch size finder altered the logger of trainer'
assert logger2 == logger3, \
'Batch size finder altered the logger of model'