From 5c0118fe9d049f29c2259a428975c25dfc2bd83c Mon Sep 17 00:00:00 2001 From: William Falcon Date: Mon, 27 Apr 2020 07:41:30 -0400 Subject: [PATCH 1/6] ddp pickle --- tests/callbacks/test_callbacks.py | 9 +++++++++ tests/loggers/test_all.py | 5 +++++ tests/trainer/test_trainer.py | 8 ++++++++ 3 files changed, 22 insertions(+) diff --git a/tests/callbacks/test_callbacks.py b/tests/callbacks/test_callbacks.py index 4731d435..c6c36ca5 100644 --- a/tests/callbacks/test_callbacks.py +++ b/tests/callbacks/test_callbacks.py @@ -240,6 +240,15 @@ def test_early_stopping_no_val_step(tmpdir): assert trainer.current_epoch < trainer.max_epochs +def test_pickling(tmpdir): + import pickle + early_stopping = EarlyStopping() + ckpt = ModelCheckpoint(tmpdir) + + pickle.dumps(ckpt) + pickle.dumps(early_stopping) + + def test_model_checkpoint_with_non_string_input(tmpdir): """ Test that None in checkpoint callback is valid and that chkp_path is set correctly """ diff --git a/tests/loggers/test_all.py b/tests/loggers/test_all.py index d9bb804b..85751204 100644 --- a/tests/loggers/test_all.py +++ b/tests/loggers/test_all.py @@ -77,6 +77,8 @@ def test_loggers_fit_test(tmpdir, monkeypatch, logger_class): # WandbLogger, # TODO: add this one ]) def test_loggers_pickle(tmpdir, monkeypatch, logger_class): + import pickle + """Verify that pickling trainer with logger works.""" tutils.reset_seed() @@ -88,6 +90,9 @@ def test_loggers_pickle(tmpdir, monkeypatch, logger_class): logger_args = _get_logger_args(logger_class, tmpdir) logger = logger_class(**logger_args) + # test pickling loggers + pickle.dumps(logger) + trainer = Trainer( max_epochs=1, logger=logger diff --git a/tests/trainer/test_trainer.py b/tests/trainer/test_trainer.py index 6876a693..9d8217f9 100644 --- a/tests/trainer/test_trainer.py +++ b/tests/trainer/test_trainer.py @@ -24,6 +24,14 @@ from tests.base import ( LightTestDataloader, LightValidationMixin, ) +from tests.base import TestModelBase + + +def test_model_pickle(tmpdir): + import pickle + + model = TestModelBase() + pickle.dumps(model) def test_hparams_save_load(tmpdir): From a3cebb44698cdcbb2c5cb1eef0839d65dc81a3e4 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Mon, 27 Apr 2020 07:44:45 -0400 Subject: [PATCH 2/6] ddp pickle --- tests/trainer/test_trainer_cli.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/tests/trainer/test_trainer_cli.py b/tests/trainer/test_trainer_cli.py index 93cbb8e2..b2d1da95 100644 --- a/tests/trainer/test_trainer_cli.py +++ b/tests/trainer/test_trainer_cli.py @@ -1,6 +1,7 @@ import inspect from argparse import ArgumentParser, Namespace from unittest import mock +import pickle import pytest @@ -42,14 +43,14 @@ def test_add_argparse_args_redefined(cli_args): args = parser.parse_args(cli_args) + # make sure we can pickle args + pickle.dumps(args) + # Check few deprecated args are not in namespace: for depr_name in ('gradient_clip', 'nb_gpu_nodes', 'max_nb_epochs'): assert depr_name not in args trainer = Trainer.from_argparse_args(args=args) - - # make sure trainer can be pickled - import pickle pickle.dumps(trainer) assert isinstance(trainer, Trainer) From 710dbf5a12808d639fd5a2a01d9e821f763fd8a0 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Mon, 27 Apr 2020 07:57:09 -0400 Subject: [PATCH 3/6] ddp pickle --- tests/trainer/test_trainer.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/trainer/test_trainer.py b/tests/trainer/test_trainer.py index 9d8217f9..116cafcc 100644 --- a/tests/trainer/test_trainer.py +++ b/tests/trainer/test_trainer.py @@ -30,7 +30,7 @@ from tests.base import TestModelBase def test_model_pickle(tmpdir): import pickle - model = TestModelBase() + model = TestModelBase(tutils.get_default_hparams()) pickle.dumps(model) From a24c88ab08392d3c7cb4863df8f6dfcea7958d07 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Mon, 27 Apr 2020 08:19:19 -0400 Subject: [PATCH 4/6] ddp pickle --- pytorch_lightning/callbacks/early_stopping.py | 13 +++++++++++-- pytorch_lightning/trainer/distrib_data_parallel.py | 1 + 2 files changed, 12 insertions(+), 2 deletions(-) diff --git a/pytorch_lightning/callbacks/early_stopping.py b/pytorch_lightning/callbacks/early_stopping.py index 4c30b677..d383a2fb 100644 --- a/pytorch_lightning/callbacks/early_stopping.py +++ b/pytorch_lightning/callbacks/early_stopping.py @@ -57,6 +57,7 @@ class EarlyStopping(Callback): self.min_delta = min_delta self.wait = 0 self.stopped_epoch = 0 + self.mode = mode mode_dict = { 'min': torch.lt, @@ -67,9 +68,8 @@ class EarlyStopping(Callback): if mode not in mode_dict: if self.verbose > 0: log.info(f'EarlyStopping mode {mode} is unknown, fallback to auto mode.') - mode = 'auto' + self.mode = 'auto' - self.monitor_op = mode_dict[mode] self.min_delta *= 1 if self.monitor_op == torch.gt else -1 def _validate_condition_metric(self, logs): @@ -94,6 +94,15 @@ class EarlyStopping(Callback): return True + @property + def monitor_op(self): + mode_dict = { + 'min': torch.lt, + 'max': torch.gt, + 'auto': torch.gt if 'acc' in self.monitor else torch.lt + } + return mode_dict[self.mode] + def on_train_start(self, trainer, pl_module): # Allow instances to be re-used self.wait = 0 diff --git a/pytorch_lightning/trainer/distrib_data_parallel.py b/pytorch_lightning/trainer/distrib_data_parallel.py index 659aa7a0..f26901c0 100644 --- a/pytorch_lightning/trainer/distrib_data_parallel.py +++ b/pytorch_lightning/trainer/distrib_data_parallel.py @@ -378,6 +378,7 @@ class TrainerDDPMixin(ABC): :param model: :return: """ + import pdb; pdb.set_trace() if self.proc_rank == 0: path = os.path.join(self.default_root_dir, '__temp_weight_ddp_end.ckpt') self.save_checkpoint(path) From 63addd091cfbcc76456c150431fe0ea168bf024c Mon Sep 17 00:00:00 2001 From: William Falcon Date: Mon, 27 Apr 2020 08:28:39 -0400 Subject: [PATCH 5/6] ddp fix --- pytorch_lightning/trainer/distrib_data_parallel.py | 1 - 1 file changed, 1 deletion(-) diff --git a/pytorch_lightning/trainer/distrib_data_parallel.py b/pytorch_lightning/trainer/distrib_data_parallel.py index f26901c0..659aa7a0 100644 --- a/pytorch_lightning/trainer/distrib_data_parallel.py +++ b/pytorch_lightning/trainer/distrib_data_parallel.py @@ -378,7 +378,6 @@ class TrainerDDPMixin(ABC): :param model: :return: """ - import pdb; pdb.set_trace() if self.proc_rank == 0: path = os.path.join(self.default_root_dir, '__temp_weight_ddp_end.ckpt') self.save_checkpoint(path) From e7ea564df2bf11e8efd7a7853bed067781e7d487 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Mon, 27 Apr 2020 08:47:19 -0400 Subject: [PATCH 6/6] pep8 --- tests/loggers/test_all.py | 2 -- tests/trainer/test_trainer.py | 1 - 2 files changed, 3 deletions(-) diff --git a/tests/loggers/test_all.py b/tests/loggers/test_all.py index 85751204..383ca263 100644 --- a/tests/loggers/test_all.py +++ b/tests/loggers/test_all.py @@ -77,8 +77,6 @@ def test_loggers_fit_test(tmpdir, monkeypatch, logger_class): # WandbLogger, # TODO: add this one ]) def test_loggers_pickle(tmpdir, monkeypatch, logger_class): - import pickle - """Verify that pickling trainer with logger works.""" tutils.reset_seed() diff --git a/tests/trainer/test_trainer.py b/tests/trainer/test_trainer.py index 116cafcc..18cc2586 100644 --- a/tests/trainer/test_trainer.py +++ b/tests/trainer/test_trainer.py @@ -24,7 +24,6 @@ from tests.base import ( LightTestDataloader, LightValidationMixin, ) -from tests.base import TestModelBase def test_model_pickle(tmpdir):