mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
fixed ckpt tests (#352)
* fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests * fixed ckpt tests
This commit is contained in:
@@ -26,6 +26,7 @@ class TestTubeLogger(LightningLoggerBase):
|
||||
def experiment(self):
|
||||
if self._experiment is not None:
|
||||
return self._experiment
|
||||
|
||||
self._experiment = Experiment(
|
||||
save_dir=self.save_dir,
|
||||
name=self.name,
|
||||
@@ -39,39 +40,45 @@ class TestTubeLogger(LightningLoggerBase):
|
||||
|
||||
@rank_zero_only
|
||||
def log_hyperparams(self, params):
|
||||
# TODO: HACK figure out where this is being set to true
|
||||
self.experiment.debug = self.debug
|
||||
self.experiment.argparse(params)
|
||||
|
||||
@rank_zero_only
|
||||
def log_metrics(self, metrics, step_num=None):
|
||||
# TODO: HACK figure out where this is being set to true
|
||||
self.experiment.debug = self.debug
|
||||
self.experiment.log(metrics, global_step=step_num)
|
||||
|
||||
@rank_zero_only
|
||||
def save(self):
|
||||
# TODO: HACK figure out where this is being set to true
|
||||
self.experiment.debug = self.debug
|
||||
self.experiment.save()
|
||||
|
||||
@rank_zero_only
|
||||
def finalize(self, status):
|
||||
# TODO: HACK figure out where this is being set to true
|
||||
self.experiment.debug = self.debug
|
||||
self.save()
|
||||
self.close()
|
||||
|
||||
@rank_zero_only
|
||||
def close(self):
|
||||
# TODO: HACK figure out where this is being set to true
|
||||
self.experiment.debug = self.debug
|
||||
exp = self.experiment
|
||||
exp.close()
|
||||
|
||||
@property
|
||||
def rank(self):
|
||||
if self._experiment is None:
|
||||
return self._rank
|
||||
else:
|
||||
return self.experiment.rank
|
||||
return self._rank
|
||||
|
||||
@rank.setter
|
||||
def rank(self, value):
|
||||
if self._experiment is None:
|
||||
self._rank = value
|
||||
else:
|
||||
return self.experiment.rank
|
||||
self._rank = value
|
||||
if self._experiment is not None:
|
||||
self.experiment.rank = value
|
||||
|
||||
@property
|
||||
def version(self):
|
||||
|
||||
@@ -983,6 +983,19 @@ class Trainer(TrainerIO):
|
||||
ref_model.use_amp = self.use_amp
|
||||
ref_model.testing = self.testing
|
||||
|
||||
# link up experiment object
|
||||
if self.logger is not None:
|
||||
ref_model.logger = self.logger
|
||||
|
||||
# save exp to get started
|
||||
if hasattr(ref_model, "hparams"):
|
||||
self.logger.log_hyperparams(ref_model.hparams)
|
||||
|
||||
self.logger.save()
|
||||
|
||||
if self.use_ddp or self.use_ddp2:
|
||||
dist.barrier()
|
||||
|
||||
# set up checkpoint callback
|
||||
self.__configure_checkpoint_callback()
|
||||
|
||||
@@ -1003,15 +1016,6 @@ class Trainer(TrainerIO):
|
||||
m = "weights_summary can be None, 'full' or 'top'"
|
||||
raise MisconfigurationException(m)
|
||||
|
||||
# link up experiment object
|
||||
if self.logger is not None:
|
||||
ref_model.logger = self.logger
|
||||
|
||||
# save exp to get started
|
||||
if hasattr(ref_model, "hparams"):
|
||||
self.logger.log_hyperparams(ref_model.hparams)
|
||||
self.logger.save()
|
||||
|
||||
# track model now.
|
||||
# if cluster resets state, the model will update with the saved weights
|
||||
self.model = model
|
||||
|
||||
+20
-20
@@ -15,7 +15,8 @@ from torchvision.datasets import MNIST
|
||||
import numpy as np
|
||||
import pdb
|
||||
# from test_models import assert_ok_test_acc, load_model, \
|
||||
# clear_save_dir, get_test_tube_logger, get_hparams, init_save_dir
|
||||
# clear_save_dir, get_test_tube_logger, get_hparams, init_save_dir, \
|
||||
# init_checkpoint_callback, reset_seed, set_random_master_port
|
||||
|
||||
|
||||
class CoolModel(pl.LightningModule):
|
||||
@@ -59,55 +60,54 @@ class CoolModel(pl.LightningModule):
|
||||
@pl.data_loader
|
||||
def test_dataloader(self):
|
||||
return DataLoader(MNIST('path/to/save', train=False), batch_size=32)
|
||||
#
|
||||
|
||||
#
|
||||
# def main():
|
||||
# """
|
||||
# Make sure DDP + AMP continue training correctly
|
||||
# :return:
|
||||
# """
|
||||
# """
|
||||
# Make sure DDP2 works
|
||||
# :return:
|
||||
# """
|
||||
# reset_seed()
|
||||
# set_random_master_port()
|
||||
#
|
||||
# hparams = get_hparams()
|
||||
# model = LightningTestModel(hparams)
|
||||
#
|
||||
# save_dir = init_save_dir()
|
||||
#
|
||||
# # logger file to get meta
|
||||
# # exp file to get meta
|
||||
# logger = get_test_tube_logger(False)
|
||||
# logger.log_hyperparams(hparams)
|
||||
# logger.save()
|
||||
#
|
||||
# # logger file to get weights
|
||||
# checkpoint = ModelCheckpoint(save_dir)
|
||||
# print(logger.debug)
|
||||
#
|
||||
# # exp file to get weights
|
||||
# checkpoint = init_checkpoint_callback(logger)
|
||||
#
|
||||
# trainer_options = dict(
|
||||
# show_progress_bar=True,
|
||||
# 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='dp'
|
||||
# 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)
|
||||
# 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
|
||||
# assert_ok_test_acc(new_trainer)
|
||||
# clear_save_dir()
|
||||
|
||||
#
|
||||
# if __name__ == '__main__':
|
||||
# main()
|
||||
|
||||
+75
-62
@@ -32,16 +32,68 @@ from pytorch_lightning.logging import TestTubeLogger
|
||||
from examples import LightningTemplateModel
|
||||
|
||||
# generate a list of random seeds for each test
|
||||
RANDOM_PORTS = list(np.random.randint(12000, 19000, 1000))
|
||||
ROOT_SEED = 1234
|
||||
torch.manual_seed(ROOT_SEED)
|
||||
np.random.seed(ROOT_SEED)
|
||||
RANDOM_SEEDS = list(np.random.randint(0, 10000, 1000))
|
||||
RANDOM_PORTS = list(np.random.randint(12000, 19000, 1000))
|
||||
|
||||
|
||||
# ------------------------------------------------------------------------
|
||||
# TESTS
|
||||
# ------------------------------------------------------------------------
|
||||
def test_running_test_pretrained_model_ddp():
|
||||
"""Verify test() on pretrained model"""
|
||||
if not can_run_gpu_test():
|
||||
return
|
||||
|
||||
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)
|
||||
|
||||
# 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)
|
||||
|
||||
[run_prediction(dataloader, pretrained_model) for dataloader in model.test_dataloader()]
|
||||
|
||||
# test we have good test accuracy
|
||||
clear_save_dir()
|
||||
|
||||
|
||||
def test_default_logger_callbacks_cpu_model():
|
||||
"""
|
||||
Test each of the trainer options
|
||||
@@ -77,11 +129,11 @@ def test_lbfgs_cpu_model():
|
||||
trainer_options = dict(
|
||||
max_nb_epochs=1,
|
||||
gradient_clip_val=1.0,
|
||||
overfit_pct=0.20,
|
||||
overfit_pct=0.30,
|
||||
print_nan_grads=True,
|
||||
show_progress_bar=False,
|
||||
weights_summary='top',
|
||||
train_percent_check=0.2,
|
||||
train_percent_check=0.3,
|
||||
val_percent_check=0.2
|
||||
)
|
||||
|
||||
@@ -144,7 +196,8 @@ def test_dp_resume():
|
||||
logger = get_test_tube_logger(debug=False)
|
||||
|
||||
# exp file to get weights
|
||||
checkpoint = ModelCheckpoint(save_dir)
|
||||
# logger file to get weights
|
||||
checkpoint = init_checkpoint_callback(logger)
|
||||
|
||||
# add these to the trainer options
|
||||
trainer_options['logger'] = logger
|
||||
@@ -202,55 +255,6 @@ def test_dp_resume():
|
||||
clear_save_dir()
|
||||
|
||||
|
||||
def test_running_test_pretrained_model_ddp():
|
||||
"""Verify test() on pretrained model"""
|
||||
if not can_run_gpu_test():
|
||||
return
|
||||
|
||||
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)
|
||||
|
||||
# exp file to get weights
|
||||
checkpoint = ModelCheckpoint(save_dir)
|
||||
|
||||
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)
|
||||
|
||||
# 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)
|
||||
|
||||
[run_prediction(dataloader, pretrained_model) for dataloader in model.test_dataloader()]
|
||||
|
||||
# test we have good test accuracy
|
||||
clear_save_dir()
|
||||
|
||||
|
||||
def test_running_test_after_fitting():
|
||||
"""Verify test() on fitted model"""
|
||||
reset_seed()
|
||||
@@ -264,7 +268,7 @@ def test_running_test_after_fitting():
|
||||
logger = get_test_tube_logger(False)
|
||||
|
||||
# logger file to get weights
|
||||
checkpoint = ModelCheckpoint(save_dir)
|
||||
checkpoint = init_checkpoint_callback(logger)
|
||||
|
||||
trainer_options = dict(
|
||||
show_progress_bar=False,
|
||||
@@ -305,7 +309,7 @@ def test_running_test_without_val():
|
||||
logger = get_test_tube_logger(False)
|
||||
|
||||
# logger file to get weights
|
||||
checkpoint = ModelCheckpoint(save_dir)
|
||||
checkpoint = init_checkpoint_callback(logger)
|
||||
|
||||
trainer_options = dict(
|
||||
show_progress_bar=False,
|
||||
@@ -344,7 +348,7 @@ def test_running_test_pretrained_model():
|
||||
logger = get_test_tube_logger(False)
|
||||
|
||||
# logger file to get weights
|
||||
checkpoint = ModelCheckpoint(save_dir)
|
||||
checkpoint = init_checkpoint_callback(logger)
|
||||
|
||||
trainer_options = dict(
|
||||
show_progress_bar=False,
|
||||
@@ -389,7 +393,7 @@ def test_running_test_pretrained_model_dp():
|
||||
logger = get_test_tube_logger(False)
|
||||
|
||||
# logger file to get weights
|
||||
checkpoint = ModelCheckpoint(save_dir)
|
||||
checkpoint = init_checkpoint_callback(logger)
|
||||
|
||||
trainer_options = dict(
|
||||
show_progress_bar=True,
|
||||
@@ -1115,7 +1119,7 @@ def test_amp_gpu_ddp_slurm_managed():
|
||||
logger = get_test_tube_logger(False)
|
||||
|
||||
# exp file to get weights
|
||||
checkpoint = ModelCheckpoint(save_dir)
|
||||
checkpoint = init_checkpoint_callback(logger)
|
||||
|
||||
# add these to the trainer options
|
||||
trainer_options['checkpoint_callback'] = checkpoint
|
||||
@@ -1459,7 +1463,7 @@ def run_gpu_model_test(trainer_options, model, hparams, on_gpu=True):
|
||||
logger = get_test_tube_logger(False)
|
||||
|
||||
# logger file to get weights
|
||||
checkpoint = ModelCheckpoint(save_dir)
|
||||
checkpoint = init_checkpoint_callback(logger)
|
||||
|
||||
# add these to the trainer options
|
||||
trainer_options['checkpoint_callback'] = checkpoint
|
||||
@@ -1529,7 +1533,7 @@ def get_test_tube_logger(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='test_tt_dir', debug=debug, version=version)
|
||||
logger = TestTubeLogger(save_dir, name='lightning_logs', debug=False, version=version)
|
||||
return logger
|
||||
|
||||
|
||||
@@ -1558,10 +1562,11 @@ def load_model(exp, save_dir, module_class=LightningTemplateModel):
|
||||
|
||||
# load trained model
|
||||
tags_path = exp.get_data_path(exp.name, exp.version)
|
||||
checkpoint_folder = os.path.join(tags_path, 'checkpoints')
|
||||
tags_path = os.path.join(tags_path, 'meta_tags.csv')
|
||||
|
||||
checkpoints = [x for x in os.listdir(save_dir) if '.ckpt' in x]
|
||||
weights_dir = os.path.join(save_dir, checkpoints[0])
|
||||
checkpoints = [x for x in os.listdir(checkpoint_folder) if '.ckpt' in x]
|
||||
weights_dir = os.path.join(checkpoint_folder, checkpoints[0])
|
||||
|
||||
trained_model = module_class.load_from_metrics(weights_path=weights_dir,
|
||||
tags_csv=tags_path)
|
||||
@@ -1631,5 +1636,13 @@ def set_random_master_port():
|
||||
os.environ['MASTER_PORT'] = str(port)
|
||||
|
||||
|
||||
def init_checkpoint_callback(logger):
|
||||
exp = logger.experiment
|
||||
exp_path = exp.get_data_path(exp.name, exp.version)
|
||||
ckpt_dir = os.path.join(exp_path, 'checkpoints')
|
||||
checkpoint = ModelCheckpoint(ckpt_dir)
|
||||
return checkpoint
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
pytest.main([__file__])
|
||||
|
||||
Reference in New Issue
Block a user