mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
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:
co-authored by
Adrian Wälchli
parent
7c19c373ac
commit
df78e84060
@@ -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'
|
||||
|
||||
@@ -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'
|
||||
|
||||
@@ -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'
|
||||
|
||||
Reference in New Issue
Block a user