2019-11-05 13:43:21 +00:00
|
|
|
import logging
|
2019-12-04 11:48:53 +00:00
|
|
|
import os
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
import torch
|
|
|
|
|
2019-12-04 11:48:53 +00:00
|
|
|
import tests.utils as tutils
|
2019-10-23 10:10:13 +00:00
|
|
|
from pytorch_lightning import Trainer
|
|
|
|
from pytorch_lightning.callbacks import ModelCheckpoint
|
|
|
|
from pytorch_lightning.testing import LightningTestModel
|
|
|
|
|
|
|
|
|
2019-12-03 13:01:04 +00:00
|
|
|
def test_running_test_pretrained_model_ddp(tmpdir):
|
2019-12-04 11:48:53 +00:00
|
|
|
"""Verify `test()` on pretrained model."""
|
2019-11-28 17:06:05 +00:00
|
|
|
if not tutils.can_run_gpu_test():
|
2019-10-23 10:10:13 +00:00
|
|
|
return
|
|
|
|
|
2019-11-28 17:06:05 +00:00
|
|
|
tutils.reset_seed()
|
|
|
|
tutils.set_random_master_port()
|
2019-10-23 10:10:13 +00:00
|
|
|
|
2019-11-28 17:06:05 +00:00
|
|
|
hparams = tutils.get_hparams()
|
2019-10-23 10:10:13 +00:00
|
|
|
model = LightningTestModel(hparams)
|
|
|
|
|
|
|
|
# exp file to get meta
|
2019-12-04 11:48:53 +00:00
|
|
|
logger = tutils.get_test_tube_logger(tmpdir, False)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
# exp file to get weights
|
2019-11-28 17:06:05 +00:00
|
|
|
checkpoint = tutils.init_checkpoint_callback(logger)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
trainer_options = dict(
|
|
|
|
show_progress_bar=False,
|
2019-12-07 13:50:21 +00:00
|
|
|
max_epochs=1,
|
2019-10-23 10:10:13 +00:00
|
|
|
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
|
2019-11-05 13:43:21 +00:00
|
|
|
logging.info(os.listdir(exp.get_data_path(exp.name, exp.version)))
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
# correct result and ok accuracy
|
|
|
|
assert result == 1, 'training failed to complete'
|
2019-11-28 17:06:05 +00:00
|
|
|
pretrained_model = tutils.load_model(logger.experiment,
|
|
|
|
trainer.checkpoint_callback.filepath,
|
|
|
|
module_class=LightningTestModel)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
# run test set
|
|
|
|
new_trainer = Trainer(**trainer_options)
|
|
|
|
new_trainer.test(pretrained_model)
|
|
|
|
|
|
|
|
for dataloader in model.test_dataloader():
|
2019-11-28 17:06:05 +00:00
|
|
|
tutils.run_prediction(dataloader, pretrained_model)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
|
2019-12-03 13:01:04 +00:00
|
|
|
def test_running_test_pretrained_model(tmpdir):
|
2019-11-28 17:06:05 +00:00
|
|
|
tutils.reset_seed()
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
"""Verify test() on pretrained model"""
|
2019-11-28 17:06:05 +00:00
|
|
|
hparams = tutils.get_hparams()
|
2019-10-23 10:10:13 +00:00
|
|
|
model = LightningTestModel(hparams)
|
|
|
|
|
|
|
|
# logger file to get meta
|
2019-12-04 11:48:53 +00:00
|
|
|
logger = tutils.get_test_tube_logger(tmpdir, False)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
# logger file to get weights
|
2019-11-28 17:06:05 +00:00
|
|
|
checkpoint = tutils.init_checkpoint_callback(logger)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
trainer_options = dict(
|
|
|
|
show_progress_bar=False,
|
2019-12-07 13:50:21 +00:00
|
|
|
max_epochs=4,
|
2019-10-23 10:10:13 +00:00
|
|
|
train_percent_check=0.4,
|
|
|
|
val_percent_check=0.2,
|
|
|
|
checkpoint_callback=checkpoint,
|
|
|
|
logger=logger
|
|
|
|
)
|
|
|
|
|
|
|
|
# fit model
|
|
|
|
trainer = Trainer(**trainer_options)
|
|
|
|
result = trainer.fit(model)
|
|
|
|
|
|
|
|
# correct result and ok accuracy
|
|
|
|
assert result == 1, 'training failed to complete'
|
2019-11-28 17:06:05 +00:00
|
|
|
pretrained_model = tutils.load_model(
|
2019-10-23 10:10:13 +00:00
|
|
|
logger.experiment, trainer.checkpoint_callback.filepath, module_class=LightningTestModel
|
|
|
|
)
|
|
|
|
|
|
|
|
new_trainer = Trainer(**trainer_options)
|
|
|
|
new_trainer.test(pretrained_model)
|
|
|
|
|
|
|
|
# test we have good test accuracy
|
2019-11-28 17:06:05 +00:00
|
|
|
tutils.assert_ok_test_acc(new_trainer)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
|
2019-12-03 13:01:04 +00:00
|
|
|
def test_load_model_from_checkpoint(tmpdir):
|
2019-11-28 17:06:05 +00:00
|
|
|
tutils.reset_seed()
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
"""Verify test() on pretrained model"""
|
2019-11-28 17:06:05 +00:00
|
|
|
hparams = tutils.get_hparams()
|
2019-10-23 10:10:13 +00:00
|
|
|
model = LightningTestModel(hparams)
|
|
|
|
|
|
|
|
trainer_options = dict(
|
|
|
|
show_progress_bar=False,
|
2020-01-14 03:25:27 +00:00
|
|
|
max_epochs=2,
|
2019-10-23 10:10:13 +00:00
|
|
|
train_percent_check=0.4,
|
|
|
|
val_percent_check=0.2,
|
2020-01-14 03:20:01 +00:00
|
|
|
checkpoint_callback=ModelCheckpoint(tmpdir, save_top_k=-1),
|
2019-10-23 10:10:13 +00:00
|
|
|
logger=False,
|
2019-12-04 11:48:53 +00:00
|
|
|
default_save_path=tmpdir,
|
2019-10-23 10:10:13 +00:00
|
|
|
)
|
|
|
|
|
|
|
|
# fit model
|
|
|
|
trainer = Trainer(**trainer_options)
|
|
|
|
result = trainer.fit(model)
|
|
|
|
|
|
|
|
# correct result and ok accuracy
|
|
|
|
assert result == 1, 'training failed to complete'
|
2020-01-14 03:25:27 +00:00
|
|
|
|
|
|
|
# load last checkpoint
|
|
|
|
last_checkpoint = os.path.join(trainer.checkpoint_callback.filepath, "_ckpt_epoch_1.ckpt")
|
|
|
|
if not os.path.isfile(last_checkpoint):
|
|
|
|
last_checkpoint = os.path.join(trainer.checkpoint_callback.filepath, "_ckpt_epoch_0.ckpt")
|
|
|
|
pretrained_model = LightningTestModel.load_from_checkpoint(last_checkpoint)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
# test that hparams loaded correctly
|
|
|
|
for k, v in vars(hparams).items():
|
|
|
|
assert getattr(pretrained_model.hparams, k) == v
|
|
|
|
|
|
|
|
new_trainer = Trainer(**trainer_options)
|
|
|
|
new_trainer.test(pretrained_model)
|
|
|
|
|
|
|
|
# test we have good test accuracy
|
2019-11-28 17:06:05 +00:00
|
|
|
tutils.assert_ok_test_acc(new_trainer)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
|
2019-12-03 13:01:04 +00:00
|
|
|
def test_running_test_pretrained_model_dp(tmpdir):
|
2019-11-28 17:06:05 +00:00
|
|
|
tutils.reset_seed()
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
"""Verify test() on pretrained model"""
|
2019-11-28 17:06:05 +00:00
|
|
|
if not tutils.can_run_gpu_test():
|
2019-10-23 10:10:13 +00:00
|
|
|
return
|
|
|
|
|
2019-11-28 17:06:05 +00:00
|
|
|
hparams = tutils.get_hparams()
|
2019-10-23 10:10:13 +00:00
|
|
|
model = LightningTestModel(hparams)
|
|
|
|
|
|
|
|
# logger file to get meta
|
2019-12-04 11:48:53 +00:00
|
|
|
logger = tutils.get_test_tube_logger(tmpdir, False)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
# logger file to get weights
|
2019-11-28 17:06:05 +00:00
|
|
|
checkpoint = tutils.init_checkpoint_callback(logger)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
trainer_options = dict(
|
|
|
|
show_progress_bar=True,
|
2020-01-14 03:09:47 +00:00
|
|
|
max_epochs=4,
|
2019-10-23 10:10:13 +00:00
|
|
|
train_percent_check=0.4,
|
|
|
|
val_percent_check=0.2,
|
|
|
|
checkpoint_callback=checkpoint,
|
|
|
|
logger=logger,
|
|
|
|
gpus=[0, 1],
|
|
|
|
distributed_backend='dp'
|
|
|
|
)
|
|
|
|
|
|
|
|
# fit model
|
|
|
|
trainer = Trainer(**trainer_options)
|
|
|
|
result = trainer.fit(model)
|
|
|
|
|
|
|
|
# correct result and ok accuracy
|
|
|
|
assert result == 1, 'training failed to complete'
|
2019-11-28 17:06:05 +00:00
|
|
|
pretrained_model = tutils.load_model(logger.experiment,
|
|
|
|
trainer.checkpoint_callback.filepath,
|
|
|
|
module_class=LightningTestModel)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
new_trainer = Trainer(**trainer_options)
|
|
|
|
new_trainer.test(pretrained_model)
|
|
|
|
|
|
|
|
# test we have good test accuracy
|
2019-11-28 17:06:05 +00:00
|
|
|
tutils.assert_ok_test_acc(new_trainer)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
|
2019-12-03 13:01:04 +00:00
|
|
|
def test_dp_resume(tmpdir):
|
2019-12-04 11:48:53 +00:00
|
|
|
"""Make sure DP continues training correctly."""
|
2019-11-28 17:06:05 +00:00
|
|
|
if not tutils.can_run_gpu_test():
|
2019-10-23 10:10:13 +00:00
|
|
|
return
|
|
|
|
|
2019-11-28 17:06:05 +00:00
|
|
|
tutils.reset_seed()
|
2019-10-23 10:10:13 +00:00
|
|
|
|
2019-11-28 17:06:05 +00:00
|
|
|
hparams = tutils.get_hparams()
|
2019-10-23 10:10:13 +00:00
|
|
|
model = LightningTestModel(hparams)
|
|
|
|
|
|
|
|
trainer_options = dict(
|
|
|
|
show_progress_bar=True,
|
2019-12-07 13:50:21 +00:00
|
|
|
max_epochs=2,
|
2019-10-23 10:10:13 +00:00
|
|
|
gpus=2,
|
|
|
|
distributed_backend='dp',
|
|
|
|
)
|
|
|
|
|
|
|
|
# get logger
|
2019-12-04 11:48:53 +00:00
|
|
|
logger = tutils.get_test_tube_logger(tmpdir, debug=False)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
# exp file to get weights
|
|
|
|
# logger file to get weights
|
2019-11-28 17:06:05 +00:00
|
|
|
checkpoint = tutils.init_checkpoint_callback(logger)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
# add these to the trainer options
|
|
|
|
trainer_options['logger'] = logger
|
|
|
|
trainer_options['checkpoint_callback'] = checkpoint
|
|
|
|
|
|
|
|
# fit model
|
|
|
|
trainer = Trainer(**trainer_options)
|
|
|
|
trainer.is_slurm_managing_tasks = True
|
|
|
|
result = trainer.fit(model)
|
|
|
|
|
|
|
|
# track epoch before saving
|
|
|
|
real_global_epoch = trainer.current_epoch
|
|
|
|
|
|
|
|
# correct result and ok accuracy
|
|
|
|
assert result == 1, 'amp + dp model failed to complete'
|
|
|
|
|
|
|
|
# ---------------------------
|
|
|
|
# HPC LOAD/SAVE
|
|
|
|
# ---------------------------
|
|
|
|
# save
|
2019-12-04 11:48:53 +00:00
|
|
|
trainer.hpc_save(tmpdir, logger)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
# init new trainer
|
2019-12-04 11:48:53 +00:00
|
|
|
new_logger = tutils.get_test_tube_logger(tmpdir, version=logger.version)
|
2019-10-23 10:10:13 +00:00
|
|
|
trainer_options['logger'] = new_logger
|
2019-12-04 11:48:53 +00:00
|
|
|
trainer_options['checkpoint_callback'] = ModelCheckpoint(tmpdir)
|
2019-10-23 10:10:13 +00:00
|
|
|
trainer_options['train_percent_check'] = 0.2
|
|
|
|
trainer_options['val_percent_check'] = 0.2
|
2019-12-07 13:50:21 +00:00
|
|
|
trainer_options['max_epochs'] = 1
|
2019-10-23 10:10:13 +00:00
|
|
|
new_trainer = Trainer(**trainer_options)
|
|
|
|
|
|
|
|
# set the epoch start hook so we can predict before the model does the full training
|
|
|
|
def assert_good_acc():
|
|
|
|
assert new_trainer.current_epoch == real_global_epoch and new_trainer.current_epoch > 0
|
|
|
|
|
|
|
|
# if model and state loaded correctly, predictions will be good even though we
|
|
|
|
# haven't trained with the new loaded model
|
|
|
|
dp_model = new_trainer.model
|
|
|
|
dp_model.eval()
|
|
|
|
|
|
|
|
dataloader = trainer.get_train_dataloader()
|
2019-11-28 17:06:05 +00:00
|
|
|
tutils.run_prediction(dataloader, dp_model, dp=True)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
# new model
|
|
|
|
model = LightningTestModel(hparams)
|
2019-12-07 13:52:06 +00:00
|
|
|
model.on_train_start = assert_good_acc
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
# fit new model which should load hpc weights
|
|
|
|
new_trainer.fit(model)
|
|
|
|
|
|
|
|
# test freeze on gpu
|
|
|
|
model.freeze()
|
|
|
|
model.unfreeze()
|
|
|
|
|
|
|
|
|
2019-12-03 13:01:04 +00:00
|
|
|
def test_cpu_restore_training(tmpdir):
|
2019-12-04 11:48:53 +00:00
|
|
|
"""Verify continue training session on CPU."""
|
2019-11-28 17:06:05 +00:00
|
|
|
tutils.reset_seed()
|
2019-10-23 10:10:13 +00:00
|
|
|
|
2019-11-28 17:06:05 +00:00
|
|
|
hparams = tutils.get_hparams()
|
2019-10-23 10:10:13 +00:00
|
|
|
model = LightningTestModel(hparams)
|
|
|
|
|
|
|
|
# logger file to get meta
|
|
|
|
test_logger_version = 10
|
2019-12-04 11:48:53 +00:00
|
|
|
logger = tutils.get_test_tube_logger(tmpdir, False, version=test_logger_version)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
trainer_options = dict(
|
2020-01-14 03:09:47 +00:00
|
|
|
max_epochs=8,
|
2019-10-23 10:10:13 +00:00
|
|
|
val_check_interval=0.50,
|
|
|
|
val_percent_check=0.2,
|
|
|
|
train_percent_check=0.2,
|
|
|
|
logger=logger,
|
2020-01-05 19:37:09 +00:00
|
|
|
checkpoint_callback=ModelCheckpoint(tmpdir, save_top_k=-1)
|
2019-10-23 10:10:13 +00:00
|
|
|
)
|
|
|
|
|
|
|
|
# fit model
|
|
|
|
trainer = Trainer(**trainer_options)
|
|
|
|
result = trainer.fit(model)
|
|
|
|
real_global_epoch = trainer.current_epoch
|
|
|
|
|
|
|
|
# traning complete
|
|
|
|
assert result == 1, 'amp + ddp model failed to complete'
|
|
|
|
|
|
|
|
# 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
|
2019-12-04 11:48:53 +00:00
|
|
|
new_logger = tutils.get_test_tube_logger(tmpdir, False, version=test_logger_version)
|
2019-10-23 10:10:13 +00:00
|
|
|
trainer_options = dict(
|
2020-01-14 03:09:47 +00:00
|
|
|
max_epochs=2,
|
2019-10-23 10:10:13 +00:00
|
|
|
val_check_interval=0.50,
|
|
|
|
val_percent_check=0.2,
|
|
|
|
train_percent_check=0.2,
|
|
|
|
logger=new_logger,
|
2019-12-04 11:48:53 +00:00
|
|
|
checkpoint_callback=ModelCheckpoint(tmpdir),
|
2019-10-23 10:10:13 +00:00
|
|
|
)
|
|
|
|
trainer = Trainer(**trainer_options)
|
|
|
|
model = LightningTestModel(hparams)
|
|
|
|
|
|
|
|
# set the epoch start hook so we can predict before the model does the full training
|
|
|
|
def assert_good_acc():
|
2019-10-23 15:41:00 +00:00
|
|
|
assert trainer.current_epoch == real_global_epoch
|
2019-10-24 00:18:26 +00:00
|
|
|
assert trainer.current_epoch >= 0
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
# if model and state loaded correctly, predictions will be good even though we
|
|
|
|
# haven't trained with the new loaded model
|
|
|
|
trainer.model.eval()
|
|
|
|
for dataloader in trainer.get_val_dataloaders():
|
2019-11-28 17:06:05 +00:00
|
|
|
tutils.run_prediction(dataloader, trainer.model)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
2019-12-07 13:52:06 +00:00
|
|
|
model.on_train_start = assert_good_acc
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
# by calling fit again, we trigger training, loading weights from the cluster
|
|
|
|
# and our hook to predict using current model before any more weight updates
|
|
|
|
trainer.fit(model)
|
|
|
|
|
|
|
|
|
2019-12-03 13:01:04 +00:00
|
|
|
def test_model_saving_loading(tmpdir):
|
2019-12-04 11:48:53 +00:00
|
|
|
"""Tests use case where trainer saves the model, and user loads it from tags independently."""
|
2019-11-28 17:06:05 +00:00
|
|
|
tutils.reset_seed()
|
2019-10-23 10:10:13 +00:00
|
|
|
|
2019-11-28 17:06:05 +00:00
|
|
|
hparams = tutils.get_hparams()
|
2019-10-23 10:10:13 +00:00
|
|
|
model = LightningTestModel(hparams)
|
|
|
|
|
|
|
|
# logger file to get meta
|
2019-12-04 11:48:53 +00:00
|
|
|
logger = tutils.get_test_tube_logger(tmpdir, False)
|
2019-10-23 10:10:13 +00:00
|
|
|
|
|
|
|
trainer_options = dict(
|
2019-12-07 13:50:21 +00:00
|
|
|
max_epochs=1,
|
2019-10-23 10:10:13 +00:00
|
|
|
logger=logger,
|
2019-12-04 11:48:53 +00:00
|
|
|
checkpoint_callback=ModelCheckpoint(tmpdir)
|
2019-10-23 10:10:13 +00:00
|
|
|
)
|
|
|
|
|
|
|
|
# fit model
|
|
|
|
trainer = Trainer(**trainer_options)
|
|
|
|
result = trainer.fit(model)
|
|
|
|
|
|
|
|
# traning complete
|
|
|
|
assert result == 1, 'amp + ddp model failed to complete'
|
|
|
|
|
|
|
|
# make a prediction
|
|
|
|
for dataloader in model.test_dataloader():
|
|
|
|
for batch in dataloader:
|
|
|
|
break
|
|
|
|
|
|
|
|
x, y = batch
|
|
|
|
x = x.view(x.size(0), -1)
|
|
|
|
|
|
|
|
# generate preds before saving model
|
|
|
|
model.eval()
|
|
|
|
pred_before_saving = model(x)
|
|
|
|
|
|
|
|
# save model
|
2019-12-04 11:48:53 +00:00
|
|
|
new_weights_path = os.path.join(tmpdir, 'save_test.ckpt')
|
2019-10-23 10:10:13 +00:00
|
|
|
trainer.save_checkpoint(new_weights_path)
|
|
|
|
|
|
|
|
# load new model
|
|
|
|
tags_path = logger.experiment.get_data_path(logger.experiment.name, logger.experiment.version)
|
|
|
|
tags_path = os.path.join(tags_path, 'meta_tags.csv')
|
|
|
|
model_2 = LightningTestModel.load_from_metrics(weights_path=new_weights_path,
|
|
|
|
tags_csv=tags_path)
|
|
|
|
model_2.eval()
|
|
|
|
|
|
|
|
# make prediction
|
|
|
|
# assert that both predictions are the same
|
|
|
|
new_pred = model_2(x)
|
|
|
|
assert torch.all(torch.eq(pred_before_saving, new_pred)).item() == 1
|
|
|
|
|
2019-12-04 11:48:53 +00:00
|
|
|
# if __name__ == '__main__':
|
|
|
|
# pytest.main([__file__])
|