From 0cd5e64701148585b7957cd62d6cf764b2d0185e Mon Sep 17 00:00:00 2001 From: Jirka Borovec Date: Mon, 4 May 2020 13:13:52 +0200 Subject: [PATCH] Tests: refactor loggers (#1689) * refactor default model * drop redundant seeds * path * refactor loggers tests * imports --- tests/loggers/test_base.py | 6 +++--- tests/loggers/test_neptune.py | 4 ++-- tests/loggers/test_trains.py | 5 ++--- tests/loggers/test_wandb.py | 1 - 4 files changed, 7 insertions(+), 9 deletions(-) diff --git a/tests/loggers/test_base.py b/tests/loggers/test_base.py index 1a52dadf..595ca0ab 100644 --- a/tests/loggers/test_base.py +++ b/tests/loggers/test_base.py @@ -7,7 +7,7 @@ import tests.base.utils as tutils from pytorch_lightning import Trainer from pytorch_lightning.loggers import LightningLoggerBase, LoggerCollection from pytorch_lightning.utilities import rank_zero_only -from tests.base import LightningTestModel, EvalModelTemplate +from tests.base import EvalModelTemplate def test_logger_collection(): @@ -61,7 +61,7 @@ class CustomLogger(LightningLoggerBase): def test_custom_logger(tmpdir): hparams = tutils.get_default_hparams() - model = LightningTestModel(hparams) + model = EvalModelTemplate(tutils.get_default_hparams()) logger = CustomLogger() @@ -80,7 +80,7 @@ def test_custom_logger(tmpdir): def test_multiple_loggers(tmpdir): hparams = tutils.get_default_hparams() - model = LightningTestModel(hparams) + model = EvalModelTemplate(hparams) logger1 = CustomLogger() logger2 = CustomLogger() diff --git a/tests/loggers/test_neptune.py b/tests/loggers/test_neptune.py index 11961234..2ca3eaf5 100644 --- a/tests/loggers/test_neptune.py +++ b/tests/loggers/test_neptune.py @@ -5,7 +5,7 @@ import torch import tests.base.utils as tutils from pytorch_lightning import Trainer from pytorch_lightning.loggers import NeptuneLogger -from tests.base import LightningTestModel +from tests.base import EvalModelTemplate @patch('pytorch_lightning.loggers.neptune.neptune') @@ -61,7 +61,7 @@ def test_neptune_additional_methods(neptune): def test_neptune_leave_open_experiment_after_fit(tmpdir): """Verify that neptune experiment was closed after training""" - model = LightningTestModel(tutils.get_default_hparams()) + model = EvalModelTemplate(tutils.get_default_hparams()) def _run_training(logger): logger._experiment = MagicMock() diff --git a/tests/loggers/test_trains.py b/tests/loggers/test_trains.py index 305d0707..738a0d9b 100644 --- a/tests/loggers/test_trains.py +++ b/tests/loggers/test_trains.py @@ -3,13 +3,12 @@ import pickle import tests.base.utils as tutils from pytorch_lightning import Trainer from pytorch_lightning.loggers import TrainsLogger -from tests.base import LightningTestModel +from tests.base import EvalModelTemplate def test_trains_logger(tmpdir): """Verify that basic functionality of TRAINS logger works.""" - hparams = tutils.get_default_hparams() - model = LightningTestModel(hparams) + model = EvalModelTemplate(tutils.get_default_hparams()) TrainsLogger.set_bypass_mode(True) TrainsLogger.set_credentials(api_host='http://integration.trains.allegro.ai:8008', files_host='http://integration.trains.allegro.ai:8081', diff --git a/tests/loggers/test_wandb.py b/tests/loggers/test_wandb.py index 3a63fcb9..4cd0eff4 100644 --- a/tests/loggers/test_wandb.py +++ b/tests/loggers/test_wandb.py @@ -2,7 +2,6 @@ import os import pickle from unittest.mock import patch -import tests.base.utils as tutils from pytorch_lightning import Trainer from pytorch_lightning.loggers import WandbLogger