From 6e86b59d21be4c895c3620a5e63701f2cefc451a Mon Sep 17 00:00:00 2001 From: William Falcon Date: Mon, 27 Apr 2020 06:51:42 -0400 Subject: [PATCH 1/7] ddp pickle --- pytorch_lightning/trainer/trainer.py | 4 ++++ pytorch_lightning/trainer/training_io.py | 29 ++---------------------- pytorch_lightning/utilities/parsing.py | 29 ++++++++++++++++++++++++ 3 files changed, 35 insertions(+), 27 deletions(-) diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index eeafb73f..d7fcd1d6 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -711,6 +711,10 @@ class Trainer( model.logger = self.logger self.copy_trainer_model_properties(model) + # clean hparams + if hasattr(model, 'hparams'): + parsing.clean_namespace(model.hparams) + # set up the passed in dataloaders (if needed) self.__attach_dataloaders(model, train_dataloader, val_dataloaders) diff --git a/pytorch_lightning/trainer/training_io.py b/pytorch_lightning/trainer/training_io.py index 4bb3c406..82bc0829 100644 --- a/pytorch_lightning/trainer/training_io.py +++ b/pytorch_lightning/trainer/training_io.py @@ -101,7 +101,7 @@ from pytorch_lightning.overrides.data_parallel import ( LightningDistributedDataParallel, LightningDataParallel, ) -from pytorch_lightning.utilities import rank_zero_warn +from pytorch_lightning.utilities import rank_zero_warn, parsing try: import torch_xla @@ -325,7 +325,7 @@ class TrainerIOMixin(ABC): checkpoint['native_amp_scaling_state'] = self.scaler.state_dict() if hasattr(model, "hparams"): - self.__clean_namespace(model.hparams) + parsing.clean_namespace(model.hparams) is_namespace = isinstance(model.hparams, Namespace) checkpoint['hparams'] = vars(model.hparams) if is_namespace else model.hparams checkpoint['hparams_type'] = 'namespace' if is_namespace else 'dict' @@ -339,31 +339,6 @@ class TrainerIOMixin(ABC): return checkpoint - def __clean_namespace(self, hparams): - """ - Removes all functions from hparams so we can pickle - :param hparams: - :return: - """ - - if isinstance(hparams, Namespace): - del_attrs = [] - for k in hparams.__dict__: - if callable(getattr(hparams, k)): - del_attrs.append(k) - - for k in del_attrs: - delattr(hparams, k) - - elif isinstance(hparams, dict): - del_attrs = [] - for k, v in hparams.items(): - if callable(v): - del_attrs.append(k) - - for k in del_attrs: - del hparams[k] - # -------------------- # HPC IO # -------------------- diff --git a/pytorch_lightning/utilities/parsing.py b/pytorch_lightning/utilities/parsing.py index 26fc410d..2549e485 100644 --- a/pytorch_lightning/utilities/parsing.py +++ b/pytorch_lightning/utilities/parsing.py @@ -1,3 +1,6 @@ +from argparse import Namespace + + def strtobool(val): """Convert a string representation of truth to true (1) or false (0). Copied from the python implementation distutils.utils.strtobool @@ -18,3 +21,29 @@ def strtobool(val): return 0 else: raise ValueError(f'invalid truth value {val}') + + +def clean_namespace(hparams): + """ + Removes all functions from hparams so we can pickle + :param hparams: + :return: + """ + + if isinstance(hparams, Namespace): + del_attrs = [] + for k in hparams.__dict__: + if callable(getattr(hparams, k)): + del_attrs.append(k) + + for k in del_attrs: + delattr(hparams, k) + + elif isinstance(hparams, dict): + del_attrs = [] + for k, v in hparams.items(): + if callable(v): + del_attrs.append(k) + + for k in del_attrs: + del hparams[k] From 4aac5568a62039f6d51088b50f7f433a68fb9ec9 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Mon, 27 Apr 2020 07:05:29 -0400 Subject: [PATCH 2/7] ddp pickle --- pytorch_lightning/trainer/trainer.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index d7fcd1d6..d3f20881 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -620,11 +620,11 @@ class Trainer( else: return int(x) - def arg_default(x): - if ',' in x: - return str(x) - else: - return int(x) + # def arg_default(x): + # if ',' in x: + # return str(x) + # else: + # return int(x) parser.add_argument( f'--{arg}', From d5dff384eb5a2b885922065ec33c39ff255cf485 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Mon, 27 Apr 2020 07:07:03 -0400 Subject: [PATCH 3/7] ddp pickle --- pytorch_lightning/trainer/trainer.py | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index d3f20881..59cdf41b 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -620,11 +620,13 @@ class Trainer( else: return int(x) - # def arg_default(x): - # if ',' in x: - # return str(x) - # else: - # return int(x) + def arg_default_fx(x): + if ',' in x: + return str(x) + else: + return int(x) + + arg_default = arg_default_fx parser.add_argument( f'--{arg}', From d2989f76e6e1474d654f830017b69662b5060aca Mon Sep 17 00:00:00 2001 From: William Falcon Date: Mon, 27 Apr 2020 07:08:34 -0400 Subject: [PATCH 4/7] ddp pickle --- pytorch_lightning/trainer/trainer.py | 27 ++++++++++++++------------- 1 file changed, 14 insertions(+), 13 deletions(-) diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index 59cdf41b..deed0098 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -614,19 +614,8 @@ class Trainer( return bool(parsing.strtobool(x)) if arg == 'gpus': - def allowed_type(x): - if ',' in x: - return str(x) - else: - return int(x) - - def arg_default_fx(x): - if ',' in x: - return str(x) - else: - return int(x) - - arg_default = arg_default_fx + allowed_type = Trainer.allowed_type + arg_default = Trainer.arg_default parser.add_argument( f'--{arg}', @@ -639,6 +628,18 @@ class Trainer( return parser + def allowed_type(x): + if ',' in x: + return str(x) + else: + return int(x) + + def arg_default(x): + if ',' in x: + return str(x) + else: + return int(x) + @classmethod def from_argparse_args(cls, args, **kwargs): From 2181ad1bc737880c3949d43735101033d126baae Mon Sep 17 00:00:00 2001 From: William Falcon Date: Mon, 27 Apr 2020 07:12:04 -0400 Subject: [PATCH 5/7] ddp pickle --- pytorch_lightning/trainer/data_loading.py | 1 + 1 file changed, 1 insertion(+) diff --git a/pytorch_lightning/trainer/data_loading.py b/pytorch_lightning/trainer/data_loading.py index 7d0b8180..92fb73e2 100644 --- a/pytorch_lightning/trainer/data_loading.py +++ b/pytorch_lightning/trainer/data_loading.py @@ -134,6 +134,7 @@ class TrainerDataLoadingMixin(ABC): 'ddp': self.num_nodes * self.num_processes, 'ddp2': self.num_nodes, } + import pdb; pdb.set_trace() sampler = DistributedSampler( dataloader.dataset, num_replicas=world_size.get(self.distributed_backend, 0), From cebc74f7bccd1f20aca969707487dca34ec5bdb5 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Mon, 27 Apr 2020 07:19:12 -0400 Subject: [PATCH 6/7] ddp pickle --- pytorch_lightning/trainer/data_loading.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pytorch_lightning/trainer/data_loading.py b/pytorch_lightning/trainer/data_loading.py index 92fb73e2..b3e15024 100644 --- a/pytorch_lightning/trainer/data_loading.py +++ b/pytorch_lightning/trainer/data_loading.py @@ -133,8 +133,8 @@ class TrainerDataLoadingMixin(ABC): world_size = { 'ddp': self.num_nodes * self.num_processes, 'ddp2': self.num_nodes, + 'ddp_cpu': self.num_processes * self.num_nodes } - import pdb; pdb.set_trace() sampler = DistributedSampler( dataloader.dataset, num_replicas=world_size.get(self.distributed_backend, 0), From b993a3ed399de2fdf2d5fd9f5646160d53fa2588 Mon Sep 17 00:00:00 2001 From: William Falcon Date: Mon, 27 Apr 2020 07:21:58 -0400 Subject: [PATCH 7/7] ddp pickle --- tests/trainer/test_trainer_cli.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/tests/trainer/test_trainer_cli.py b/tests/trainer/test_trainer_cli.py index bfc67111..93cbb8e2 100644 --- a/tests/trainer/test_trainer_cli.py +++ b/tests/trainer/test_trainer_cli.py @@ -47,6 +47,11 @@ def test_add_argparse_args_redefined(cli_args): 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)