diff --git a/pytorch_lightning/testing/lm_test_module_base.py b/pytorch_lightning/testing/lm_test_module_base.py index d14cf534..5bd7c0bc 100644 --- a/pytorch_lightning/testing/lm_test_module_base.py +++ b/pytorch_lightning/testing/lm_test_module_base.py @@ -104,7 +104,8 @@ class LightningTestModelBase(LightningModule): if self.trainer.batch_nb % 1 == 0: output = OrderedDict({ 'loss': loss_val, - 'progress_bar': {'some_val': loss_val * loss_val} + 'progress_bar': {'some_val': loss_val * loss_val}, + 'log': {'train_some_val': loss_val * loss_val}, }) return output diff --git a/pytorch_lightning/testing/lm_test_module_mixins.py b/pytorch_lightning/testing/lm_test_module_mixins.py index feab206f..562126db 100644 --- a/pytorch_lightning/testing/lm_test_module_mixins.py +++ b/pytorch_lightning/testing/lm_test_module_mixins.py @@ -105,7 +105,7 @@ class LightningValidationMixin(LightningValidationStepMixin): val_acc_mean /= len(outputs) tqdm_dict = {'val_loss': val_loss_mean.item(), 'val_acc': val_acc_mean.item()} - results = {'progress_bar': tqdm_dict} + results = {'progress_bar': tqdm_dict, 'log': tqdm_dict} return results diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index 10cf2c57..26f109a7 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -183,6 +183,7 @@ class Trainer(TrainerIO): version=self.slurm_job_id, name='lightning_logs' ) + self.logger.rank = 0 # configure checkpoint callback self.checkpoint_callback = checkpoint_callback @@ -1159,12 +1160,14 @@ class Trainer(TrainerIO): def __metrics_to_scalars(self, metrics): new_metrics = {} for k, v in metrics.items(): - if type(v) is torch.Tensor: + if isinstance(v, torch.Tensor): v = v.item() if type(v) is dict: v = self.__metrics_to_scalars(v) + new_metrics[k] = v + return new_metrics def __log_vals_blacklist(self): @@ -1335,7 +1338,6 @@ class Trainer(TrainerIO): # track progress bar metrics self.__add_tqdm_metrics(progress_bar_metrics) - all_log_metrics.append(log_metrics) # accumulate loss @@ -1402,7 +1404,6 @@ class Trainer(TrainerIO): # collapse all metrics into one dict all_log_metrics = {k: v for d in all_log_metrics for k, v in d.items()} - return 0, grad_norm_dic, all_log_metrics def __run_evaluation(self, test=False): @@ -1443,7 +1444,6 @@ class Trainer(TrainerIO): dataloaders, max_batches, test) - _, progress_bar_metrics, log_metrics = self.__process_output(eval_results) # add metrics to prog bar diff --git a/pytorch_lightning/trainer/trainer_io.py b/pytorch_lightning/trainer/trainer_io.py index e11560ff..f6305308 100644 --- a/pytorch_lightning/trainer/trainer_io.py +++ b/pytorch_lightning/trainer/trainer_io.py @@ -5,7 +5,7 @@ import pdb from subprocess import call import torch - +import torch.distributed as dist from pytorch_lightning.pt_overrides.override_data_parallel import ( LightningDistributedDataParallel, LightningDataParallel) @@ -35,6 +35,11 @@ class TrainerIO(object): # if script called from hpc resubmit, load weights self.restore_hpc_weights_if_needed(model) + # wait for all models to restore weights + if self.use_ddp or self.use_ddp2: + # wait for all processes to catch up + dist.barrier() + def restore_state_if_checkpoint_exists(self, model): # do nothing if there's not dir or callback no_ckpt_callback = self.checkpoint_callback is None diff --git a/tests/debug.py b/tests/debug.py index bce19eda..e69afd80 100644 --- a/tests/debug.py +++ b/tests/debug.py @@ -14,7 +14,8 @@ from torch.utils.data import DataLoader from torchvision.datasets import MNIST import numpy as np import pdb -from . import test_models +# from test_models import assert_ok_test_acc, load_model, \ +# clear_save_dir, get_test_tube_logger, get_hparams, init_save_dir class CoolModel(pl.LightningModule): @@ -58,57 +59,55 @@ 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: +# """ +# hparams = get_hparams() +# model = LightningTestModel(hparams) +# +# save_dir = init_save_dir() +# +# # logger 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) +# +# trainer_options = dict( +# show_progress_bar=True, +# 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' +# ) +# +# # 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) +# +# 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() - -def main(): - """ - Make sure DDP + AMP continue training correctly - :return: - """ - """ - Make sure DDP2 works - :return: - """ - hparams = test_models.get_hparams() - model = LightningTestModel(hparams) - - save_dir = test_models.init_save_dir() - - # logger file to get meta - logger = test_models.get_test_tube_logger(False) - logger.log_hyperparams(hparams) - logger.save() - - # logger file to get weights - checkpoint = ModelCheckpoint(save_dir) - - trainer_options = dict( - show_progress_bar=True, - 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' - ) - - # fit model - trainer = Trainer(**trainer_options) - result = trainer.fit(model) - - # correct result and ok accuracy - assert result == 1, 'training failed to complete' - pretrained_model = test_models.load_model(logger.experiment, save_dir, - module_class=LightningTestModel) - - new_trainer = Trainer(**trainer_options) - new_trainer.test(pretrained_model) - - # test we have good test accuracy - test_models.assert_ok_test_acc(new_trainer) - test_models.clear_save_dir() - - -if __name__ == '__main__': - main() +# if __name__ == '__main__': +# main()