diff --git a/pytorch_lightning/trainer/data_loading.py b/pytorch_lightning/trainer/data_loading.py index 7d0b8180..b3e15024 100644 --- a/pytorch_lightning/trainer/data_loading.py +++ b/pytorch_lightning/trainer/data_loading.py @@ -133,6 +133,7 @@ class TrainerDataLoadingMixin(ABC): world_size = { 'ddp': self.num_nodes * self.num_processes, 'ddp2': self.num_nodes, + 'ddp_cpu': self.num_processes * self.num_nodes } sampler = DistributedSampler( dataloader.dataset, diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index eeafb73f..deed0098 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -614,17 +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(x): - if ',' in x: - return str(x) - else: - return int(x) + allowed_type = Trainer.allowed_type + arg_default = Trainer.arg_default parser.add_argument( f'--{arg}', @@ -637,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): @@ -711,6 +714,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] 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)