mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-11 12:31:23 +08:00
Tests: refactor cleanup (#1744)
* wip * cleaning * optim imports * - * default hparams * fix restore * fix imports
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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):
|
||||
|
||||
|
||||
Reference in New Issue
Block a user