mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-12 12:40:20 +08:00
Use pytest tmpdir fixture (#482)
* Use pytest tmpdir * Switch to tmpdir fixtures * Switch to tmpdir fixture * tmpdir fixture * Fix more conflicts
This commit is contained in:
committed by
William Falcon
parent
a6d64ac013
commit
62f6f92fdf
+4
-33
@@ -17,7 +17,6 @@ from pytorch_lightning.testing import (
|
||||
)
|
||||
|
||||
# generate a list of random seeds for each test
|
||||
RANDOM_FILE_PATHS = list(np.random.randint(12000, 19000, 1000))
|
||||
RANDOM_PORTS = list(np.random.randint(12000, 19000, 1000))
|
||||
ROOT_SEED = 1234
|
||||
torch.manual_seed(ROOT_SEED)
|
||||
@@ -26,8 +25,7 @@ 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):
|
||||
save_dir = init_save_dir()
|
||||
trainer_options['default_save_path'] = save_dir
|
||||
save_dir = trainer_options['default_save_path']
|
||||
|
||||
# fit model
|
||||
trainer = Trainer(**trainer_options)
|
||||
@@ -49,14 +47,12 @@ def run_model_test_no_loggers(trainer_options, model, hparams, on_gpu=True, min_
|
||||
trainer.model = pretrained_model
|
||||
trainer.optimizers, trainer.lr_schedulers = pretrained_model.configure_optimizers()
|
||||
|
||||
clear_save_dir()
|
||||
|
||||
|
||||
def run_model_test(trainer_options, model, hparams, on_gpu=True):
|
||||
save_dir = init_save_dir()
|
||||
save_dir = trainer_options['default_save_path']
|
||||
|
||||
# logger file to get meta
|
||||
logger = get_test_tube_logger(False)
|
||||
logger = get_test_tube_logger(save_dir, False)
|
||||
|
||||
# logger file to get weights
|
||||
checkpoint = init_checkpoint_callback(logger)
|
||||
@@ -87,8 +83,6 @@ def run_model_test(trainer_options, model, hparams, on_gpu=True):
|
||||
trainer.hpc_save(save_dir, logger)
|
||||
trainer.hpc_load(save_dir, on_gpu=on_gpu)
|
||||
|
||||
clear_save_dir()
|
||||
|
||||
|
||||
def get_hparams(continue_training=False, hpc_exp_number=0):
|
||||
root_dir = os.path.dirname(os.path.realpath(__file__))
|
||||
@@ -126,35 +120,12 @@ def get_model(use_test_model=False, lbfgs=False):
|
||||
return model, hparams
|
||||
|
||||
|
||||
def get_test_tube_logger(debug=True, version=None):
|
||||
def get_test_tube_logger(save_dir, debug=True, version=None):
|
||||
# set up logger object without actually saving logs
|
||||
root_dir = os.path.dirname(os.path.realpath(__file__))
|
||||
save_dir = os.path.join(root_dir, 'save_dir')
|
||||
logger = TestTubeLogger(save_dir, name='lightning_logs', debug=False, version=version)
|
||||
return logger
|
||||
|
||||
|
||||
def init_save_dir():
|
||||
root_dir = os.path.dirname(os.path.realpath(__file__))
|
||||
save_dir = os.path.join(root_dir, 'tests', 'save_dir')
|
||||
|
||||
if os.path.exists(save_dir):
|
||||
n = RANDOM_FILE_PATHS.pop()
|
||||
shutil.move(save_dir, save_dir + f'_{n}')
|
||||
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
|
||||
return save_dir
|
||||
|
||||
|
||||
def clear_save_dir():
|
||||
root_dir = os.path.dirname(os.path.realpath(__file__))
|
||||
save_dir = os.path.join(root_dir, 'tests', 'save_dir')
|
||||
if os.path.exists(save_dir):
|
||||
n = RANDOM_FILE_PATHS.pop()
|
||||
shutil.move(save_dir, save_dir + f'_{n}')
|
||||
|
||||
|
||||
def load_model(exp, root_weights_dir, module_class=LightningTemplateModel):
|
||||
# load trained model
|
||||
tags_path = exp.get_data_path(exp.name, exp.version)
|
||||
|
||||
Reference in New Issue
Block a user