From df78e84060cd140895698a1fe186b24df1d428cd Mon Sep 17 00:00:00 2001 From: Jirka Borovec Date: Thu, 28 May 2020 04:45:23 +0200 Subject: [PATCH] unify tests (#1940) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * unify tests * Apply suggestions from code review Co-authored-by: Adrian Wälchli Co-authored-by: Adrian Wälchli --- tests/loggers/test_all.py | 25 +++++++++++++++++++++++++ tests/trainer/test_lr_finder.py | 23 ----------------------- tests/trainer/test_trainer_tricks.py | 23 ----------------------- 3 files changed, 25 insertions(+), 46 deletions(-) diff --git a/tests/loggers/test_all.py b/tests/loggers/test_all.py index 84d5c793..c001d4ac 100644 --- a/tests/loggers/test_all.py +++ b/tests/loggers/test_all.py @@ -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' diff --git a/tests/trainer/test_lr_finder.py b/tests/trainer/test_lr_finder.py index 9450e980..4134b587 100755 --- a/tests/trainer/test_lr_finder.py +++ b/tests/trainer/test_lr_finder.py @@ -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' diff --git a/tests/trainer/test_trainer_tricks.py b/tests/trainer/test_trainer_tricks.py index 81eb7e13..a66e8bbd 100755 --- a/tests/trainer/test_trainer_tricks.py +++ b/tests/trainer/test_trainer_tricks.py @@ -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'