mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-10 12:21:57 +08:00
Tests: refactor loggers (#1689)
* refactor default model * drop redundant seeds * path * refactor loggers tests * imports
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user