mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
prune tests (#564)
* format docstring in tests * prune unused vars * optimize imports * drop duplicated var
This commit is contained in:
committed by
William Falcon
parent
62f6f92fdf
commit
63717e8fda
@@ -1,17 +1,16 @@
|
||||
import os
|
||||
import logging
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
import tests.utils as tutils
|
||||
from pytorch_lightning import Trainer
|
||||
from pytorch_lightning.callbacks import ModelCheckpoint
|
||||
from pytorch_lightning.testing import LightningTestModel
|
||||
import tests.utils as tutils
|
||||
|
||||
|
||||
def test_running_test_pretrained_model_ddp(tmpdir):
|
||||
"""Verify test() on pretrained model"""
|
||||
"""Verify `test()` on pretrained model."""
|
||||
if not tutils.can_run_gpu_test():
|
||||
return
|
||||
|
||||
@@ -21,10 +20,8 @@ def test_running_test_pretrained_model_ddp(tmpdir):
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
save_dir = tmpdir
|
||||
|
||||
# exp file to get meta
|
||||
logger = tutils.get_test_tube_logger(save_dir, False)
|
||||
logger = tutils.get_test_tube_logger(tmpdir, False)
|
||||
|
||||
# exp file to get weights
|
||||
checkpoint = tutils.init_checkpoint_callback(logger)
|
||||
@@ -68,10 +65,8 @@ def test_running_test_pretrained_model(tmpdir):
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
save_dir = tmpdir
|
||||
|
||||
# logger file to get meta
|
||||
logger = tutils.get_test_tube_logger(save_dir, False)
|
||||
logger = tutils.get_test_tube_logger(tmpdir, False)
|
||||
|
||||
# logger file to get weights
|
||||
checkpoint = tutils.init_checkpoint_callback(logger)
|
||||
@@ -109,8 +104,6 @@ def test_load_model_from_checkpoint(tmpdir):
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
save_dir = tmpdir
|
||||
|
||||
trainer_options = dict(
|
||||
show_progress_bar=False,
|
||||
max_nb_epochs=1,
|
||||
@@ -118,7 +111,7 @@ def test_load_model_from_checkpoint(tmpdir):
|
||||
val_percent_check=0.2,
|
||||
checkpoint_callback=True,
|
||||
logger=False,
|
||||
default_save_path=save_dir
|
||||
default_save_path=tmpdir,
|
||||
)
|
||||
|
||||
# fit model
|
||||
@@ -152,10 +145,8 @@ def test_running_test_pretrained_model_dp(tmpdir):
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
save_dir = tmpdir
|
||||
|
||||
# logger file to get meta
|
||||
logger = tutils.get_test_tube_logger(save_dir, False)
|
||||
logger = tutils.get_test_tube_logger(tmpdir, False)
|
||||
|
||||
# logger file to get weights
|
||||
checkpoint = tutils.init_checkpoint_callback(logger)
|
||||
@@ -189,10 +180,7 @@ def test_running_test_pretrained_model_dp(tmpdir):
|
||||
|
||||
|
||||
def test_dp_resume(tmpdir):
|
||||
"""
|
||||
Make sure DP continues training correctly
|
||||
:return:
|
||||
"""
|
||||
"""Make sure DP continues training correctly."""
|
||||
if not tutils.can_run_gpu_test():
|
||||
return
|
||||
|
||||
@@ -208,10 +196,8 @@ def test_dp_resume(tmpdir):
|
||||
distributed_backend='dp',
|
||||
)
|
||||
|
||||
save_dir = tmpdir
|
||||
|
||||
# get logger
|
||||
logger = tutils.get_test_tube_logger(save_dir, debug=False)
|
||||
logger = tutils.get_test_tube_logger(tmpdir, debug=False)
|
||||
|
||||
# exp file to get weights
|
||||
# logger file to get weights
|
||||
@@ -236,12 +222,12 @@ def test_dp_resume(tmpdir):
|
||||
# HPC LOAD/SAVE
|
||||
# ---------------------------
|
||||
# save
|
||||
trainer.hpc_save(save_dir, logger)
|
||||
trainer.hpc_save(tmpdir, logger)
|
||||
|
||||
# init new trainer
|
||||
new_logger = tutils.get_test_tube_logger(save_dir, version=logger.version)
|
||||
new_logger = tutils.get_test_tube_logger(tmpdir, version=logger.version)
|
||||
trainer_options['logger'] = new_logger
|
||||
trainer_options['checkpoint_callback'] = ModelCheckpoint(save_dir)
|
||||
trainer_options['checkpoint_callback'] = ModelCheckpoint(tmpdir)
|
||||
trainer_options['train_percent_check'] = 0.2
|
||||
trainer_options['val_percent_check'] = 0.2
|
||||
trainer_options['max_nb_epochs'] = 1
|
||||
@@ -272,20 +258,15 @@ def test_dp_resume(tmpdir):
|
||||
|
||||
|
||||
def test_cpu_restore_training(tmpdir):
|
||||
"""
|
||||
Verify continue training session on CPU
|
||||
:return:
|
||||
"""
|
||||
"""Verify continue training session on CPU."""
|
||||
tutils.reset_seed()
|
||||
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
save_dir = tmpdir
|
||||
|
||||
# logger file to get meta
|
||||
test_logger_version = 10
|
||||
logger = tutils.get_test_tube_logger(save_dir, False, version=test_logger_version)
|
||||
logger = tutils.get_test_tube_logger(tmpdir, False, version=test_logger_version)
|
||||
|
||||
trainer_options = dict(
|
||||
max_nb_epochs=2,
|
||||
@@ -293,7 +274,7 @@ def test_cpu_restore_training(tmpdir):
|
||||
val_percent_check=0.2,
|
||||
train_percent_check=0.2,
|
||||
logger=logger,
|
||||
checkpoint_callback=ModelCheckpoint(save_dir)
|
||||
checkpoint_callback=ModelCheckpoint(tmpdir)
|
||||
)
|
||||
|
||||
# fit model
|
||||
@@ -307,14 +288,14 @@ def test_cpu_restore_training(tmpdir):
|
||||
# wipe-out trainer and model
|
||||
# retrain with not much data... this simulates picking training back up after slurm
|
||||
# we want to see if the weights come back correctly
|
||||
new_logger = tutils.get_test_tube_logger(save_dir, False, version=test_logger_version)
|
||||
new_logger = tutils.get_test_tube_logger(tmpdir, False, version=test_logger_version)
|
||||
trainer_options = dict(
|
||||
max_nb_epochs=2,
|
||||
val_check_interval=0.50,
|
||||
val_percent_check=0.2,
|
||||
train_percent_check=0.2,
|
||||
logger=new_logger,
|
||||
checkpoint_callback=ModelCheckpoint(save_dir),
|
||||
checkpoint_callback=ModelCheckpoint(tmpdir),
|
||||
)
|
||||
trainer = Trainer(**trainer_options)
|
||||
model = LightningTestModel(hparams)
|
||||
@@ -338,24 +319,19 @@ def test_cpu_restore_training(tmpdir):
|
||||
|
||||
|
||||
def test_model_saving_loading(tmpdir):
|
||||
"""
|
||||
Tests use case where trainer saves the model, and user loads it from tags independently
|
||||
:return:
|
||||
"""
|
||||
"""Tests use case where trainer saves the model, and user loads it from tags independently."""
|
||||
tutils.reset_seed()
|
||||
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
save_dir = tmpdir
|
||||
|
||||
# logger file to get meta
|
||||
logger = tutils.get_test_tube_logger(save_dir, False)
|
||||
logger = tutils.get_test_tube_logger(tmpdir, False)
|
||||
|
||||
trainer_options = dict(
|
||||
max_nb_epochs=1,
|
||||
logger=logger,
|
||||
checkpoint_callback=ModelCheckpoint(save_dir)
|
||||
checkpoint_callback=ModelCheckpoint(tmpdir)
|
||||
)
|
||||
|
||||
# fit model
|
||||
@@ -378,7 +354,7 @@ def test_model_saving_loading(tmpdir):
|
||||
pred_before_saving = model(x)
|
||||
|
||||
# save model
|
||||
new_weights_path = os.path.join(save_dir, 'save_test.ckpt')
|
||||
new_weights_path = os.path.join(tmpdir, 'save_test.ckpt')
|
||||
trainer.save_checkpoint(new_weights_path)
|
||||
|
||||
# load new model
|
||||
@@ -394,5 +370,5 @@ def test_model_saving_loading(tmpdir):
|
||||
assert torch.all(torch.eq(pred_before_saving, new_pred)).item() == 1
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
pytest.main([__file__])
|
||||
# if __name__ == '__main__':
|
||||
# pytest.main([__file__])
|
||||
|
||||
Reference in New Issue
Block a user