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
@@ -52,54 +52,3 @@ class CoolModel(pl.LightningModule):
|
||||
@pl.data_loader
|
||||
def test_dataloader(self):
|
||||
return DataLoader(MNIST('path/to/save', train=False), batch_size=32)
|
||||
|
||||
#
|
||||
# def main():
|
||||
# reset_seed()
|
||||
# set_random_master_port()
|
||||
#
|
||||
# hparams = get_hparams()
|
||||
# model = LightningTestModel(hparams)
|
||||
#
|
||||
# save_dir = init_save_dir()
|
||||
#
|
||||
# # exp file to get meta
|
||||
# logger = get_test_tube_logger(False)
|
||||
#
|
||||
# print(logger.debug)
|
||||
#
|
||||
# # exp file to get weights
|
||||
# checkpoint = init_checkpoint_callback(logger)
|
||||
#
|
||||
# trainer_options = dict(
|
||||
# show_progress_bar=False,
|
||||
# max_nb_epochs=1,
|
||||
# train_percent_check=0.4,
|
||||
# val_percent_check=0.2,
|
||||
# checkpoint_callback=checkpoint,
|
||||
# logger=logger,
|
||||
# gpus=[0, 1],
|
||||
# distributed_backend='ddp'
|
||||
# )
|
||||
#
|
||||
# # fit model
|
||||
# trainer = Trainer(**trainer_options)
|
||||
# result = trainer.fit(model)
|
||||
#
|
||||
# exp = logger.experiment
|
||||
# print(os.listdir(exp.get_data_path(exp.name, exp.version)))
|
||||
#
|
||||
# # correct result and ok accuracy
|
||||
# assert result == 1, 'training failed to complete'
|
||||
# pretrained_model = load_model(logger.experiment, save_dir,
|
||||
# module_class=LightningTestModel)
|
||||
#
|
||||
# # run test set
|
||||
# new_trainer = Trainer(**trainer_options)
|
||||
# new_trainer.test(pretrained_model)
|
||||
#
|
||||
# # test we have good test accuracy
|
||||
# clear_save_dir()
|
||||
#
|
||||
# if __name__ == '__main__':
|
||||
# main()
|
||||
|
||||
+14
-36
@@ -1,22 +1,17 @@
|
||||
import os
|
||||
import warnings
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
import tests.utils as tutils
|
||||
from pytorch_lightning import Trainer
|
||||
from pytorch_lightning.testing import (
|
||||
LightningTestModel,
|
||||
)
|
||||
from pytorch_lightning.utilities.debugging import MisconfigurationException
|
||||
import tests.utils as tutils
|
||||
|
||||
|
||||
def test_amp_single_gpu(tmpdir):
|
||||
"""
|
||||
Make sure DDP + AMP work
|
||||
:return:
|
||||
"""
|
||||
"""Make sure DDP + AMP work."""
|
||||
tutils.reset_seed()
|
||||
|
||||
if not tutils.can_run_gpu_test():
|
||||
@@ -34,14 +29,11 @@ def test_amp_single_gpu(tmpdir):
|
||||
use_amp=True
|
||||
)
|
||||
|
||||
tutils.run_model_test(trainer_options, model, hparams)
|
||||
tutils.run_model_test(trainer_options, model)
|
||||
|
||||
|
||||
def test_no_amp_single_gpu(tmpdir):
|
||||
"""
|
||||
Make sure DDP + AMP work
|
||||
:return:
|
||||
"""
|
||||
"""Make sure DDP + AMP work."""
|
||||
tutils.reset_seed()
|
||||
|
||||
if not tutils.can_run_gpu_test():
|
||||
@@ -60,14 +52,11 @@ def test_no_amp_single_gpu(tmpdir):
|
||||
)
|
||||
|
||||
with pytest.raises((MisconfigurationException, ModuleNotFoundError)):
|
||||
tutils.run_model_test(trainer_options, model, hparams)
|
||||
tutils.run_model_test(trainer_options, model)
|
||||
|
||||
|
||||
def test_amp_gpu_ddp(tmpdir):
|
||||
"""
|
||||
Make sure DDP + AMP work
|
||||
:return:
|
||||
"""
|
||||
"""Make sure DDP + AMP work."""
|
||||
if not tutils.can_run_gpu_test():
|
||||
return
|
||||
|
||||
@@ -86,14 +75,11 @@ def test_amp_gpu_ddp(tmpdir):
|
||||
use_amp=True
|
||||
)
|
||||
|
||||
tutils.run_model_test(trainer_options, model, hparams)
|
||||
tutils.run_model_test(trainer_options, model)
|
||||
|
||||
|
||||
def test_amp_gpu_ddp_slurm_managed(tmpdir):
|
||||
"""
|
||||
Make sure DDP + AMP work
|
||||
:return:
|
||||
"""
|
||||
"""Make sure DDP + AMP work."""
|
||||
if not tutils.can_run_gpu_test():
|
||||
return
|
||||
|
||||
@@ -114,10 +100,8 @@ def test_amp_gpu_ddp_slurm_managed(tmpdir):
|
||||
use_amp=True
|
||||
)
|
||||
|
||||
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)
|
||||
@@ -153,8 +137,8 @@ def test_amp_gpu_ddp_slurm_managed(tmpdir):
|
||||
trainer.optimizers, trainer.lr_schedulers = pretrained_model.configure_optimizers()
|
||||
|
||||
# test HPC loading / saving
|
||||
trainer.hpc_save(save_dir, logger)
|
||||
trainer.hpc_load(save_dir, on_gpu=True)
|
||||
trainer.hpc_save(tmpdir, logger)
|
||||
trainer.hpc_load(tmpdir, on_gpu=True)
|
||||
|
||||
# test freeze on gpu
|
||||
model.freeze()
|
||||
@@ -162,10 +146,7 @@ def test_amp_gpu_ddp_slurm_managed(tmpdir):
|
||||
|
||||
|
||||
def test_cpu_model_with_amp(tmpdir):
|
||||
"""
|
||||
Make sure model trains on CPU
|
||||
:return:
|
||||
"""
|
||||
"""Make sure model trains on CPU."""
|
||||
tutils.reset_seed()
|
||||
|
||||
trainer_options = dict(
|
||||
@@ -181,14 +162,11 @@ def test_cpu_model_with_amp(tmpdir):
|
||||
model, hparams = tutils.get_model()
|
||||
|
||||
with pytest.raises((MisconfigurationException, ModuleNotFoundError)):
|
||||
tutils.run_model_test(trainer_options, model, hparams, on_gpu=False)
|
||||
tutils.run_model_test(trainer_options, model, on_gpu=False)
|
||||
|
||||
|
||||
def test_amp_gpu_dp(tmpdir):
|
||||
"""
|
||||
Make sure DP + AMP work
|
||||
:return:
|
||||
"""
|
||||
"""Make sure DP + AMP work."""
|
||||
tutils.reset_seed()
|
||||
|
||||
if not tutils.can_run_gpu_test():
|
||||
|
||||
+19
-46
@@ -1,8 +1,8 @@
|
||||
import warnings
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
import tests.utils as tutils
|
||||
from pytorch_lightning import Trainer, data_loader
|
||||
from pytorch_lightning.callbacks import (
|
||||
EarlyStopping,
|
||||
@@ -12,14 +12,10 @@ from pytorch_lightning.testing import (
|
||||
LightningTestModelBase,
|
||||
LightningTestMixin,
|
||||
)
|
||||
import tests.utils as tutils
|
||||
|
||||
|
||||
def test_early_stopping_cpu_model(tmpdir):
|
||||
"""
|
||||
Test each of the trainer options
|
||||
:return:
|
||||
"""
|
||||
"""Test each of the trainer options."""
|
||||
tutils.reset_seed()
|
||||
|
||||
stopping = EarlyStopping(monitor='val_loss', min_delta=0.1)
|
||||
@@ -37,7 +33,7 @@ def test_early_stopping_cpu_model(tmpdir):
|
||||
)
|
||||
|
||||
model, hparams = tutils.get_model()
|
||||
tutils.run_model_test(trainer_options, model, hparams, on_gpu=False)
|
||||
tutils.run_model_test(trainer_options, model, on_gpu=False)
|
||||
|
||||
# test freeze on cpu
|
||||
model.freeze()
|
||||
@@ -45,10 +41,7 @@ def test_early_stopping_cpu_model(tmpdir):
|
||||
|
||||
|
||||
def test_lbfgs_cpu_model(tmpdir):
|
||||
"""
|
||||
Test each of the trainer options
|
||||
:return:
|
||||
"""
|
||||
"""Test each of the trainer options."""
|
||||
tutils.reset_seed()
|
||||
|
||||
trainer_options = dict(
|
||||
@@ -62,15 +55,11 @@ def test_lbfgs_cpu_model(tmpdir):
|
||||
)
|
||||
|
||||
model, hparams = tutils.get_model(use_test_model=True, lbfgs=True)
|
||||
tutils.run_model_test_no_loggers(trainer_options, model, hparams,
|
||||
on_gpu=False, min_acc=0.30)
|
||||
tutils.run_model_test_no_loggers(trainer_options, model, min_acc=0.30)
|
||||
|
||||
|
||||
def test_default_logger_callbacks_cpu_model(tmpdir):
|
||||
"""
|
||||
Test each of the trainer options
|
||||
:return:
|
||||
"""
|
||||
"""Test each of the trainer options."""
|
||||
tutils.reset_seed()
|
||||
|
||||
trainer_options = dict(
|
||||
@@ -85,7 +74,7 @@ def test_default_logger_callbacks_cpu_model(tmpdir):
|
||||
)
|
||||
|
||||
model, hparams = tutils.get_model()
|
||||
tutils.run_model_test_no_loggers(trainer_options, model, hparams, on_gpu=False)
|
||||
tutils.run_model_test_no_loggers(trainer_options, model)
|
||||
|
||||
# test freeze on cpu
|
||||
model.freeze()
|
||||
@@ -93,7 +82,7 @@ def test_default_logger_callbacks_cpu_model(tmpdir):
|
||||
|
||||
|
||||
def test_running_test_after_fitting(tmpdir):
|
||||
"""Verify test() on fitted model"""
|
||||
"""Verify test() on fitted model."""
|
||||
tutils.reset_seed()
|
||||
|
||||
hparams = tutils.get_hparams()
|
||||
@@ -129,10 +118,9 @@ def test_running_test_after_fitting(tmpdir):
|
||||
|
||||
|
||||
def test_running_test_without_val(tmpdir):
|
||||
"""Verify `test()` works on a model with no `val_loader`."""
|
||||
tutils.reset_seed()
|
||||
|
||||
"""Verify test() works on a model with no val_loader"""
|
||||
|
||||
class CurrentTestModel(LightningTestMixin, LightningTestModelBase):
|
||||
pass
|
||||
|
||||
@@ -212,10 +200,7 @@ def test_single_gpu_batch_parse():
|
||||
|
||||
|
||||
def test_simple_cpu(tmpdir):
|
||||
"""
|
||||
Verify continue training session on CPU
|
||||
:return:
|
||||
"""
|
||||
"""Verify continue training session on CPU."""
|
||||
tutils.reset_seed()
|
||||
|
||||
hparams = tutils.get_hparams()
|
||||
@@ -238,10 +223,7 @@ def test_simple_cpu(tmpdir):
|
||||
|
||||
|
||||
def test_cpu_model(tmpdir):
|
||||
"""
|
||||
Make sure model trains on CPU
|
||||
:return:
|
||||
"""
|
||||
"""Make sure model trains on CPU."""
|
||||
tutils.reset_seed()
|
||||
|
||||
trainer_options = dict(
|
||||
@@ -255,14 +237,11 @@ def test_cpu_model(tmpdir):
|
||||
|
||||
model, hparams = tutils.get_model()
|
||||
|
||||
tutils.run_model_test(trainer_options, model, hparams, on_gpu=False)
|
||||
tutils.run_model_test(trainer_options, model, on_gpu=False)
|
||||
|
||||
|
||||
def test_all_features_cpu_model(tmpdir):
|
||||
"""
|
||||
Test each of the trainer options
|
||||
:return:
|
||||
"""
|
||||
"""Test each of the trainer options."""
|
||||
tutils.reset_seed()
|
||||
|
||||
trainer_options = dict(
|
||||
@@ -280,14 +259,11 @@ def test_all_features_cpu_model(tmpdir):
|
||||
)
|
||||
|
||||
model, hparams = tutils.get_model()
|
||||
tutils.run_model_test(trainer_options, model, hparams, on_gpu=False)
|
||||
tutils.run_model_test(trainer_options, model, on_gpu=False)
|
||||
|
||||
|
||||
def test_tbptt_cpu_model(tmpdir):
|
||||
"""
|
||||
Test truncated back propagation through time works.
|
||||
:return:
|
||||
"""
|
||||
"""Test truncated back propagation through time works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
truncated_bptt_steps = 2
|
||||
@@ -360,10 +336,7 @@ def test_tbptt_cpu_model(tmpdir):
|
||||
|
||||
|
||||
def test_single_gpu_model(tmpdir):
|
||||
"""
|
||||
Make sure single GPU works (DP mode)
|
||||
:return:
|
||||
"""
|
||||
"""Make sure single GPU works (DP mode)."""
|
||||
tutils.reset_seed()
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
@@ -381,8 +354,8 @@ def test_single_gpu_model(tmpdir):
|
||||
gpus=1
|
||||
)
|
||||
|
||||
tutils.run_model_test(trainer_options, model, hparams)
|
||||
tutils.run_model_test(trainer_options, model)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
pytest.main([__file__])
|
||||
# if __name__ == '__main__':
|
||||
# pytest.main([__file__])
|
||||
|
||||
+21
-41
@@ -1,7 +1,9 @@
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
import tests.utils as tutils
|
||||
from pytorch_lightning import Trainer
|
||||
from pytorch_lightning.callbacks import (
|
||||
ModelCheckpoint,
|
||||
@@ -15,16 +17,12 @@ from pytorch_lightning.trainer.dp_mixin import (
|
||||
determine_root_gpu_device,
|
||||
)
|
||||
from pytorch_lightning.utilities.debugging import MisconfigurationException
|
||||
import tests.utils as tutils
|
||||
|
||||
PRETEND_N_OF_GPUS = 16
|
||||
|
||||
|
||||
def test_multi_gpu_model_ddp2(tmpdir):
|
||||
"""
|
||||
Make sure DDP2 works
|
||||
:return:
|
||||
"""
|
||||
"""Make sure DDP2 works."""
|
||||
if not tutils.can_run_gpu_test():
|
||||
return
|
||||
|
||||
@@ -43,14 +41,11 @@ def test_multi_gpu_model_ddp2(tmpdir):
|
||||
distributed_backend='ddp2'
|
||||
)
|
||||
|
||||
tutils.run_model_test(trainer_options, model, hparams)
|
||||
tutils.run_model_test(trainer_options, model)
|
||||
|
||||
|
||||
def test_multi_gpu_model_ddp(tmpdir):
|
||||
"""
|
||||
Make sure DDP works
|
||||
:return:
|
||||
"""
|
||||
"""Make sure DDP works."""
|
||||
if not tutils.can_run_gpu_test():
|
||||
return
|
||||
|
||||
@@ -68,7 +63,7 @@ def test_multi_gpu_model_ddp(tmpdir):
|
||||
distributed_backend='ddp'
|
||||
)
|
||||
|
||||
tutils.run_model_test(trainer_options, model, hparams)
|
||||
tutils.run_model_test(trainer_options, model)
|
||||
|
||||
|
||||
def test_optimizer_return_options():
|
||||
@@ -103,26 +98,20 @@ def test_optimizer_return_options():
|
||||
|
||||
|
||||
def test_cpu_slurm_save_load(tmpdir):
|
||||
"""
|
||||
Verify model save/load/checkpoint on CPU
|
||||
:return:
|
||||
"""
|
||||
"""Verify model save/load/checkpoint on CPU."""
|
||||
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)
|
||||
version = logger.version
|
||||
|
||||
trainer_options = dict(
|
||||
max_nb_epochs=1,
|
||||
logger=logger,
|
||||
checkpoint_callback=ModelCheckpoint(save_dir)
|
||||
checkpoint_callback=ModelCheckpoint(tmpdir)
|
||||
)
|
||||
|
||||
# fit model
|
||||
@@ -147,16 +136,16 @@ def test_cpu_slurm_save_load(tmpdir):
|
||||
|
||||
# test HPC saving
|
||||
# simulate snapshot on slurm
|
||||
saved_filepath = trainer.hpc_save(save_dir, logger)
|
||||
saved_filepath = trainer.hpc_save(tmpdir, logger)
|
||||
assert os.path.exists(saved_filepath)
|
||||
|
||||
# new logger file to get meta
|
||||
logger = tutils.get_test_tube_logger(save_dir, False, version=version)
|
||||
logger = tutils.get_test_tube_logger(tmpdir, False, version=version)
|
||||
|
||||
trainer_options = dict(
|
||||
max_nb_epochs=1,
|
||||
logger=logger,
|
||||
checkpoint_callback=ModelCheckpoint(save_dir),
|
||||
checkpoint_callback=ModelCheckpoint(tmpdir),
|
||||
)
|
||||
trainer = Trainer(**trainer_options)
|
||||
model = LightningTestModel(hparams)
|
||||
@@ -178,11 +167,7 @@ def test_cpu_slurm_save_load(tmpdir):
|
||||
|
||||
|
||||
def test_multi_gpu_none_backend(tmpdir):
|
||||
"""
|
||||
Make sure when using multiple GPUs the user can't use
|
||||
distributed_backend = None
|
||||
:return:
|
||||
"""
|
||||
"""Make sure when using multiple GPUs the user can't use `distributed_backend = None`."""
|
||||
tutils.reset_seed()
|
||||
|
||||
if not tutils.can_run_gpu_test():
|
||||
@@ -199,14 +184,11 @@ def test_multi_gpu_none_backend(tmpdir):
|
||||
)
|
||||
|
||||
with pytest.raises(MisconfigurationException):
|
||||
tutils.run_model_test(trainer_options, model, hparams)
|
||||
tutils.run_model_test(trainer_options, model)
|
||||
|
||||
|
||||
def test_multi_gpu_model_dp(tmpdir):
|
||||
"""
|
||||
Make sure DP works
|
||||
:return:
|
||||
"""
|
||||
"""Make sure DP works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
if not tutils.can_run_gpu_test():
|
||||
@@ -223,17 +205,14 @@ def test_multi_gpu_model_dp(tmpdir):
|
||||
gpus='-1'
|
||||
)
|
||||
|
||||
tutils.run_model_test(trainer_options, model, hparams)
|
||||
tutils.run_model_test(trainer_options, model)
|
||||
|
||||
# test memory helper functions
|
||||
memory.get_memory_profile('min_max')
|
||||
|
||||
|
||||
def test_ddp_sampler_error(tmpdir):
|
||||
"""
|
||||
Make sure DDP + AMP work
|
||||
:return:
|
||||
"""
|
||||
"""Make sure DDP + AMP work."""
|
||||
if not tutils.can_run_gpu_test():
|
||||
return
|
||||
|
||||
@@ -374,7 +353,8 @@ test_parse_gpu_ids_data = [
|
||||
pytest.param(1, [0]),
|
||||
pytest.param(-1, list(range(PRETEND_N_OF_GPUS)), id="-1 - use all gpus"),
|
||||
pytest.param('-1', list(range(PRETEND_N_OF_GPUS)), id="'-1' - use all gpus"),
|
||||
pytest.param(3, [0, 1, 2])]
|
||||
pytest.param(3, [0, 1, 2]),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.gpus_param_tests
|
||||
@@ -403,5 +383,5 @@ def test_parse_gpu_returns_None_when_no_devices_are_available(mocked_device_coun
|
||||
parse_gpu_ids(gpus)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
pytest.main([__file__])
|
||||
# if __name__ == '__main__':
|
||||
# pytest.main([__file__])
|
||||
|
||||
+10
-29
@@ -1,26 +1,19 @@
|
||||
import os
|
||||
import pickle
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from pytorch_lightning import Trainer
|
||||
from pytorch_lightning.testing import LightningTestModel
|
||||
from pytorch_lightning.logging import LightningLoggerBase, rank_zero_only
|
||||
import tests.utils as tutils
|
||||
from pytorch_lightning import Trainer
|
||||
from pytorch_lightning.logging import LightningLoggerBase, rank_zero_only
|
||||
from pytorch_lightning.testing import LightningTestModel
|
||||
|
||||
|
||||
def test_testtube_logger(tmpdir):
|
||||
"""
|
||||
verify that basic functionality of test tube logger works
|
||||
"""
|
||||
"""Verify that basic functionality of test tube logger works."""
|
||||
tutils.reset_seed()
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
save_dir = tmpdir
|
||||
|
||||
logger = tutils.get_test_tube_logger(save_dir, False)
|
||||
logger = tutils.get_test_tube_logger(tmpdir, False)
|
||||
|
||||
trainer_options = dict(
|
||||
max_nb_epochs=1,
|
||||
@@ -35,16 +28,12 @@ def test_testtube_logger(tmpdir):
|
||||
|
||||
|
||||
def test_testtube_pickle(tmpdir):
|
||||
"""
|
||||
Verify that pickling a trainer containing a test tube logger works
|
||||
"""
|
||||
"""Verify that pickling a trainer containing a test tube logger works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
save_dir = tmpdir
|
||||
|
||||
logger = tutils.get_test_tube_logger(tmpdir, False)
|
||||
logger.log_hyperparams(hparams)
|
||||
logger.save()
|
||||
@@ -62,9 +51,7 @@ def test_testtube_pickle(tmpdir):
|
||||
|
||||
|
||||
def test_mlflow_logger(tmpdir):
|
||||
"""
|
||||
verify that basic functionality of mlflow logger works
|
||||
"""
|
||||
"""Verify that basic functionality of mlflow logger works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
try:
|
||||
@@ -93,9 +80,7 @@ def test_mlflow_logger(tmpdir):
|
||||
|
||||
|
||||
def test_mlflow_pickle(tmpdir):
|
||||
"""
|
||||
verify that pickling trainer with mlflow logger works
|
||||
"""
|
||||
"""Verify that pickling trainer with mlflow logger works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
try:
|
||||
@@ -122,9 +107,7 @@ def test_mlflow_pickle(tmpdir):
|
||||
|
||||
|
||||
def test_comet_logger(tmpdir):
|
||||
"""
|
||||
verify that basic functionality of Comet.ml logger works
|
||||
"""
|
||||
"""Verify that basic functionality of Comet.ml logger works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
try:
|
||||
@@ -158,9 +141,7 @@ def test_comet_logger(tmpdir):
|
||||
|
||||
|
||||
def test_comet_pickle(tmpdir):
|
||||
"""
|
||||
verify that pickling trainer with comet logger works
|
||||
"""
|
||||
"""Verify that pickling trainer with comet logger works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
try:
|
||||
|
||||
@@ -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__])
|
||||
|
||||
+16
-38
@@ -1,7 +1,9 @@
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
import tests.utils as tutils
|
||||
from pytorch_lightning import Trainer
|
||||
from pytorch_lightning.callbacks import (
|
||||
ModelCheckpoint,
|
||||
@@ -11,19 +13,14 @@ from pytorch_lightning.testing import (
|
||||
LightningTestModelBase,
|
||||
LightningValidationStepMixin,
|
||||
LightningValidationMultipleDataloadersMixin,
|
||||
LightningTestMixin,
|
||||
LightningTestMultipleDataloadersMixin,
|
||||
)
|
||||
from pytorch_lightning.trainer import trainer_io
|
||||
from pytorch_lightning.trainer.logging_mixin import TrainerLoggingMixin
|
||||
import tests.utils as tutils
|
||||
|
||||
|
||||
def test_no_val_module(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()
|
||||
@@ -33,15 +30,13 @@ def test_no_val_module(tmpdir):
|
||||
|
||||
model = CurrentTestModel(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
|
||||
@@ -52,7 +47,7 @@ def test_no_val_module(tmpdir):
|
||||
assert result == 1, 'amp + ddp model failed to complete'
|
||||
|
||||
# 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
|
||||
@@ -64,10 +59,7 @@ def test_no_val_module(tmpdir):
|
||||
|
||||
|
||||
def test_no_val_end_module(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()
|
||||
|
||||
class CurrentTestModel(LightningValidationStepMixin, LightningTestModelBase):
|
||||
@@ -76,15 +68,13 @@ def test_no_val_end_module(tmpdir):
|
||||
hparams = tutils.get_hparams()
|
||||
model = CurrentTestModel(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
|
||||
@@ -95,7 +85,7 @@ def test_no_val_end_module(tmpdir):
|
||||
assert result == 1, 'amp + ddp model failed to complete'
|
||||
|
||||
# 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
|
||||
@@ -226,18 +216,12 @@ def test_dp_output_reduce():
|
||||
|
||||
|
||||
def test_model_checkpoint_options(tmp_path):
|
||||
"""
|
||||
Test ModelCheckpoint options
|
||||
:return:
|
||||
"""
|
||||
|
||||
# TODO split this up into multiple tests
|
||||
|
||||
"""Test ModelCheckpoint options."""
|
||||
def mock_save_function(filepath):
|
||||
open(filepath, 'a').close()
|
||||
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
_ = LightningTestModel(hparams)
|
||||
|
||||
# simulated losses
|
||||
save_dir = tmp_path / "1"
|
||||
@@ -355,10 +339,7 @@ def test_model_freeze_unfreeze():
|
||||
|
||||
|
||||
def test_multiple_val_dataloader(tmpdir):
|
||||
"""
|
||||
Verify multiple val_dataloader
|
||||
:return:
|
||||
"""
|
||||
"""Verify multiple val_dataloader."""
|
||||
tutils.reset_seed()
|
||||
|
||||
class CurrentTestModel(
|
||||
@@ -395,10 +376,7 @@ def test_multiple_val_dataloader(tmpdir):
|
||||
|
||||
|
||||
def test_multiple_test_dataloader(tmpdir):
|
||||
"""
|
||||
Verify multiple test_dataloader
|
||||
:return:
|
||||
"""
|
||||
"""Verify multiple test_dataloader."""
|
||||
tutils.reset_seed()
|
||||
|
||||
class CurrentTestModel(
|
||||
@@ -434,5 +412,5 @@ def test_multiple_test_dataloader(tmpdir):
|
||||
trainer.test()
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
pytest.main([__file__])
|
||||
# if __name__ == '__main__':
|
||||
# pytest.main([__file__])
|
||||
|
||||
+5
-4
@@ -24,7 +24,7 @@ np.random.seed(ROOT_SEED)
|
||||
RANDOM_SEEDS = list(np.random.randint(0, 10000, 1000))
|
||||
|
||||
|
||||
def run_model_test_no_loggers(trainer_options, model, hparams, on_gpu=True, min_acc=0.50):
|
||||
def run_model_test_no_loggers(trainer_options, model, min_acc=0.50):
|
||||
save_dir = trainer_options['default_save_path']
|
||||
|
||||
# fit model
|
||||
@@ -48,7 +48,7 @@ def run_model_test_no_loggers(trainer_options, model, hparams, on_gpu=True, min_
|
||||
trainer.optimizers, trainer.lr_schedulers = pretrained_model.configure_optimizers()
|
||||
|
||||
|
||||
def run_model_test(trainer_options, model, hparams, on_gpu=True):
|
||||
def run_model_test(trainer_options, model, on_gpu=True):
|
||||
save_dir = trainer_options['default_save_path']
|
||||
|
||||
# logger file to get meta
|
||||
@@ -95,7 +95,8 @@ def get_hparams(continue_training=False, hpc_exp_number=0):
|
||||
'optimizer_name': 'adam',
|
||||
'data_root': os.path.join(root_dir, 'mnist'),
|
||||
'out_features': 10,
|
||||
'hidden_dim': 1000}
|
||||
'hidden_dim': 1000,
|
||||
}
|
||||
|
||||
if continue_training:
|
||||
args['test_tube_do_checkpoint_load'] = True
|
||||
@@ -122,7 +123,7 @@ def get_model(use_test_model=False, lbfgs=False):
|
||||
|
||||
def get_test_tube_logger(save_dir, debug=True, version=None):
|
||||
# set up logger object without actually saving logs
|
||||
logger = TestTubeLogger(save_dir, name='lightning_logs', debug=False, version=version)
|
||||
logger = TestTubeLogger(save_dir, name='lightning_logs', debug=debug, version=version)
|
||||
return logger
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user