From 58a467dd68b157fdba8824a437dbaf698ad88569 Mon Sep 17 00:00:00 2001 From: Jirka Borovec Date: Fri, 24 Apr 2020 23:21:00 +0200 Subject: [PATCH] model checkpint on rank_zero_only & global rank state (#1408) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * try delete in async or DDP us0-ecase * changelog * add model chekpoint rank * simple delete * flake8 * use global rank * chnagelog * fix review * fix import * proposal * proposal * proposal * improve proposal (fix problems with method call self) * cleaning Co-authored-by: Adrian Wälchli Co-authored-by: William Falcon --- CHANGELOG.md | 10 ++++++ .../callbacks/model_checkpoint.py | 9 ++++-- pytorch_lightning/loggers/__init__.py | 5 +-- pytorch_lightning/loggers/base.py | 31 +------------------ pytorch_lightning/loggers/comet.py | 3 +- pytorch_lightning/loggers/mlflow.py | 3 +- pytorch_lightning/loggers/neptune.py | 3 +- pytorch_lightning/loggers/tensorboard.py | 3 +- pytorch_lightning/loggers/test_tube.py | 15 ++------- pytorch_lightning/loggers/trains.py | 3 +- pytorch_lightning/loggers/wandb.py | 3 +- .../trainer/distrib_data_parallel.py | 11 +++---- pytorch_lightning/trainer/distrib_parts.py | 10 ++---- pytorch_lightning/trainer/logging.py | 2 -- pytorch_lightning/utilities/__init__.py | 2 +- pytorch_lightning/utilities/distributed.py | 26 ++++++++++++++++ pytorch_lightning/utilities/warnings.py | 18 ----------- tests/loggers/test_base.py | 3 +- 18 files changed, 72 insertions(+), 88 deletions(-) create mode 100644 pytorch_lightning/utilities/distributed.py delete mode 100644 pytorch_lightning/utilities/warnings.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 2d8d1049..8c61cb82 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -22,6 +22,10 @@ The format is based on [Keep a Changelog](http://keepachangelog.com/en/1.0.0/). - Added `terminate_on_nan` flag to trainer that performs a NaN check with each training iteration when set to `True` ([#1475](https://github.com/PyTorchLightning/pytorch-lightning/pull/1475)) +- Added speed parity tests (max 1 sec difference per epoch)([#1482](https://github.com/PyTorchLightning/pytorch-lightning/pull/1482)) + +- Added `terminate_on_nan` flag to trainer that performs a NaN check with each training iteration when set to `True`. ([#1475](https://github.com/PyTorchLightning/pytorch-lightning/pull/1475)) + - Added `ddp_cpu` backend for testing ddp without GPUs ([#1158](https://github.com/PyTorchLightning/pytorch-lightning/pull/1158)) - Added [Horovod](http://horovod.ai) support as a distributed backend `Trainer(distributed_backend='horovod')` ([#1529](https://github.com/PyTorchLightning/pytorch-lightning/pull/1529)) @@ -33,10 +37,13 @@ The format is based on [Keep a Changelog](http://keepachangelog.com/en/1.0.0/). ### Changed - Changed the default behaviour to no longer include a NaN check with each training iteration. ([#1475](https://github.com/PyTorchLightning/pytorch-lightning/pull/1475)) + - Decoupled the progress bar from trainer. It is a callback now and can be customized or even be replaced entirely ([#1450](https://github.com/PyTorchLightning/pytorch-lightning/pull/1450)). - Changed lr schedule step interval behavior to update every backwards pass instead of every forwards pass ([#1476](https://github.com/PyTorchLightning/pytorch-lightning/issues/1476)) +- Defines shared proc. rank, remove rank from instances (e.g. loggers) ([#1408](https://github.com/PyTorchLightning/pytorch-lightning/pull/1408)) + - Updated semantic segmentation example with custom u-net and logging ([#1371](https://github.com/PyTorchLightning/pytorch-lightning/pull/1371)) @@ -74,6 +81,8 @@ The format is based on [Keep a Changelog](http://keepachangelog.com/en/1.0.0/). - Fixed do not copy the batch when training on a single GPU ([#1576](https://github.com/PyTorchLightning/pytorch-lightning/issues/1576), [#1579](https://github.com/PyTorchLightning/pytorch-lightning/issues/1579)) +- Fixed soft checkpoint removing on DDP ([#1408](https://github.com/PyTorchLightning/pytorch-lightning/pull/1408)) + - Fixes automatic parser bug ([#1585](https://github.com/PyTorchLightning/pytorch-lightning/issues/1585)) ## [0.7.3] - 2020-04-09 @@ -90,6 +99,7 @@ The format is based on [Keep a Changelog](http://keepachangelog.com/en/1.0.0/). - Fixed gradient clipping ([#1438](https://github.com/PyTorchLightning/pytorch-lightning/pull/1438)) - Fixed pretty print ([#1441](https://github.com/PyTorchLightning/pytorch-lightning/pull/1441)) + ## [0.7.2] - 2020-04-07 ### Added diff --git a/pytorch_lightning/callbacks/model_checkpoint.py b/pytorch_lightning/callbacks/model_checkpoint.py index 298a7bb0..1d5d6d14 100644 --- a/pytorch_lightning/callbacks/model_checkpoint.py +++ b/pytorch_lightning/callbacks/model_checkpoint.py @@ -12,10 +12,10 @@ import re import numpy as np from typing import Optional +import torch from pytorch_lightning import _logger as log from pytorch_lightning.callbacks.base import Callback -from pytorch_lightning.utilities import rank_zero_warn -import torch +from pytorch_lightning.utilities import rank_zero_warn, rank_zero_only class ModelCheckpoint(Callback): @@ -91,6 +91,7 @@ class ModelCheckpoint(Callback): f"Checkpoint directory {filepath} exists and is not empty with save_top_k != 0." "All files in this directory will be deleted when a checkpoint is saved!" ) + self._rank = 0 self.monitor = monitor self.verbose = verbose @@ -129,7 +130,8 @@ class ModelCheckpoint(Callback): self.monitor_op, self.kth_value, self.mode = mode_dict[mode] def _del_model(self, filepath): - os.remove(filepath) + if os.path.isfile(filepath): + os.remove(filepath) def _save_model(self, filepath): # make paths @@ -189,6 +191,7 @@ class ModelCheckpoint(Callback): filepath = os.path.join(self.dirpath, self.prefix + filename + str_ver + '.ckpt') return filepath + @rank_zero_only def on_validation_end(self, trainer, pl_module): # only run on main process if trainer.proc_rank != 0: diff --git a/pytorch_lightning/loggers/__init__.py b/pytorch_lightning/loggers/__init__.py index edb3ffec..c753a048 100644 --- a/pytorch_lightning/loggers/__init__.py +++ b/pytorch_lightning/loggers/__init__.py @@ -30,7 +30,8 @@ You can implement your own logger by writing a class that inherits from :class:`LightningLoggerBase`. Use the :func:`~pytorch_lightning.loggers.base.rank_zero_only` decorator to make sure that only the first process in DDP training logs data. ->>> from pytorch_lightning.loggers import LightningLoggerBase, rank_zero_only +>>> from pytorch_lightning.utilities import rank_zero_only +>>> from pytorch_lightning.loggers import LightningLoggerBase >>> class MyLogger(LightningLoggerBase): ... ... @rank_zero_only @@ -80,7 +81,7 @@ Supported Loggers """ from os import environ -from pytorch_lightning.loggers.base import LightningLoggerBase, LoggerCollection, rank_zero_only +from pytorch_lightning.loggers.base import LightningLoggerBase, LoggerCollection from pytorch_lightning.loggers.tensorboard import TensorBoardLogger __all__ = [ diff --git a/pytorch_lightning/loggers/base.py b/pytorch_lightning/loggers/base.py index 27f07025..39891c44 100644 --- a/pytorch_lightning/loggers/base.py +++ b/pytorch_lightning/loggers/base.py @@ -3,26 +3,12 @@ import functools import operator from abc import ABC, abstractmethod from argparse import Namespace -from functools import wraps from typing import Union, Optional, Dict, Iterable, Any, Callable, List, Sequence, Mapping, Tuple import numpy as np import torch - -def rank_zero_only(fn: Callable): - """Decorate a logger method to run it only on the process with rank 0. - - Args: - fn: Function to decorate - """ - - @wraps(fn) - def wrapped_fn(self, *args, **kwargs): - if self.rank == 0: - fn(self, *args, **kwargs) - - return wrapped_fn +from pytorch_lightning.utilities import rank_zero_only class LightningLoggerBase(ABC): @@ -251,16 +237,6 @@ class LightningLoggerBase(ABC): """Do any cleanup that is necessary to close an experiment.""" self.save() - @property - def rank(self) -> int: - """Process rank. In general, metrics should only be logged by the process with rank 0.""" - return self._rank - - @rank.setter - def rank(self, value: int) -> None: - """Set the process rank.""" - self._rank = value - @property @abstractmethod def name(self) -> str: @@ -307,11 +283,6 @@ class LoggerCollection(LightningLoggerBase): def close(self) -> None: [logger.close() for logger in self._logger_iterable] - @LightningLoggerBase.rank.setter - def rank(self, value: int) -> None: - for logger in self._logger_iterable: - logger.rank = value - @property def name(self) -> str: return '_'.join([str(logger.name) for logger in self._logger_iterable]) diff --git a/pytorch_lightning/loggers/comet.py b/pytorch_lightning/loggers/comet.py index 8cd5dd22..1d24676d 100644 --- a/pytorch_lightning/loggers/comet.py +++ b/pytorch_lightning/loggers/comet.py @@ -24,8 +24,9 @@ import torch from torch import is_tensor from pytorch_lightning import _logger as log -from pytorch_lightning.loggers.base import LightningLoggerBase, rank_zero_only +from pytorch_lightning.loggers.base import LightningLoggerBase from pytorch_lightning.utilities.exceptions import MisconfigurationException +from pytorch_lightning.utilities import rank_zero_only class CometLogger(LightningLoggerBase): diff --git a/pytorch_lightning/loggers/mlflow.py b/pytorch_lightning/loggers/mlflow.py index 6f672c6a..74c51893 100644 --- a/pytorch_lightning/loggers/mlflow.py +++ b/pytorch_lightning/loggers/mlflow.py @@ -15,7 +15,8 @@ except ImportError: # pragma: no-cover ' install it with `pip install mlflow`.') from pytorch_lightning import _logger as log -from pytorch_lightning.loggers.base import LightningLoggerBase, rank_zero_only +from pytorch_lightning.loggers.base import LightningLoggerBase +from pytorch_lightning.utilities import rank_zero_only class MLFlowLogger(LightningLoggerBase): diff --git a/pytorch_lightning/loggers/neptune.py b/pytorch_lightning/loggers/neptune.py index e4a41757..374b513b 100644 --- a/pytorch_lightning/loggers/neptune.py +++ b/pytorch_lightning/loggers/neptune.py @@ -18,7 +18,8 @@ import torch from torch import is_tensor from pytorch_lightning import _logger as log -from pytorch_lightning.loggers.base import LightningLoggerBase, rank_zero_only +from pytorch_lightning.loggers.base import LightningLoggerBase +from pytorch_lightning.utilities import rank_zero_only class NeptuneLogger(LightningLoggerBase): diff --git a/pytorch_lightning/loggers/tensorboard.py b/pytorch_lightning/loggers/tensorboard.py index 58e26f09..613262dd 100644 --- a/pytorch_lightning/loggers/tensorboard.py +++ b/pytorch_lightning/loggers/tensorboard.py @@ -13,8 +13,9 @@ import torch from pkg_resources import parse_version from torch.utils.tensorboard import SummaryWriter -from pytorch_lightning.loggers.base import LightningLoggerBase, rank_zero_only from pytorch_lightning import _logger as log +from pytorch_lightning.loggers.base import LightningLoggerBase +from pytorch_lightning.utilities import rank_zero_only class TensorBoardLogger(LightningLoggerBase): diff --git a/pytorch_lightning/loggers/test_tube.py b/pytorch_lightning/loggers/test_tube.py index fb81dbd1..7f382c3f 100644 --- a/pytorch_lightning/loggers/test_tube.py +++ b/pytorch_lightning/loggers/test_tube.py @@ -11,7 +11,8 @@ except ImportError: # pragma: no-cover raise ImportError('You want to use `test_tube` logger which is not installed yet,' # pragma: no-cover ' install it with `pip install test-tube`.') -from pytorch_lightning.loggers.base import LightningLoggerBase, rank_zero_only +from pytorch_lightning.loggers.base import LightningLoggerBase +from pytorch_lightning.utilities.distributed import rank_zero_only class TestTubeLogger(LightningLoggerBase): @@ -92,7 +93,7 @@ class TestTubeLogger(LightningLoggerBase): version=self.version, description=self.description, create_git_tag=self.create_git_tag, - rank=self.rank, + rank=rank_zero_only.rank, ) return self._experiment @@ -134,16 +135,6 @@ class TestTubeLogger(LightningLoggerBase): exp = self.experiment exp.close() - @property - def rank(self) -> int: - return self._rank - - @rank.setter - def rank(self, value: int) -> None: - self._rank = value - if self._experiment is not None: - self.experiment.rank = value - @property def name(self) -> str: if self._experiment is None: diff --git a/pytorch_lightning/loggers/trains.py b/pytorch_lightning/loggers/trains.py index ab43bdd3..ca4f38c6 100644 --- a/pytorch_lightning/loggers/trains.py +++ b/pytorch_lightning/loggers/trains.py @@ -19,7 +19,8 @@ except ImportError: # pragma: no-cover ' install it with `pip install trains`.') from pytorch_lightning import _logger as log -from pytorch_lightning.loggers.base import LightningLoggerBase, rank_zero_only +from pytorch_lightning.loggers.base import LightningLoggerBase +from pytorch_lightning.utilities import rank_zero_only class TrainsLogger(LightningLoggerBase): diff --git a/pytorch_lightning/loggers/wandb.py b/pytorch_lightning/loggers/wandb.py index 3e844305..c3486441 100644 --- a/pytorch_lightning/loggers/wandb.py +++ b/pytorch_lightning/loggers/wandb.py @@ -15,7 +15,8 @@ except ImportError: # pragma: no-cover raise ImportError('You want to use `wandb` logger which is not installed yet,' # pragma: no-cover ' install it with `pip install wandb`.') -from pytorch_lightning.loggers.base import LightningLoggerBase, rank_zero_only +from pytorch_lightning.loggers.base import LightningLoggerBase +from pytorch_lightning.utilities import rank_zero_only class WandbLogger(LightningLoggerBase): diff --git a/pytorch_lightning/trainer/distrib_data_parallel.py b/pytorch_lightning/trainer/distrib_data_parallel.py index abb2a32f..659aa7a0 100644 --- a/pytorch_lightning/trainer/distrib_data_parallel.py +++ b/pytorch_lightning/trainer/distrib_data_parallel.py @@ -120,9 +120,10 @@ from typing import Union import torch from pytorch_lightning import _logger as log +from pytorch_lightning.callbacks import ModelCheckpoint from pytorch_lightning.loggers import LightningLoggerBase from pytorch_lightning.utilities.exceptions import MisconfigurationException -from pytorch_lightning.utilities.warnings import set_proc_rank, rank_zero_warn +from pytorch_lightning.utilities.distributed import rank_zero_only, rank_zero_warn try: from apex import amp @@ -146,6 +147,7 @@ class TrainerDDPMixin(ABC): on_gpu: bool num_gpu_nodes: int logger: Union[LightningLoggerBase, bool] + checkpoint_callback: Union[ModelCheckpoint, bool] data_parallel_device_ids: ... distributed_backend: str amp_level: str @@ -322,12 +324,9 @@ class TrainerDDPMixin(ABC): elif self.use_ddp2: self.proc_rank = self.node_rank self.world_size = self.num_nodes - # set warning rank - set_proc_rank(self.proc_rank) - # let the exp know the rank to avoid overwriting logs - if self.logger is not None: - self.logger.rank = self.proc_rank + # set warning rank + rank_zero_only.rank = self.proc_rank # set up server using proc 0's ip address # try to init for 20 times at max in case ports are taken diff --git a/pytorch_lightning/trainer/distrib_parts.py b/pytorch_lightning/trainer/distrib_parts.py index 06905be9..73efaf67 100644 --- a/pytorch_lightning/trainer/distrib_parts.py +++ b/pytorch_lightning/trainer/distrib_parts.py @@ -352,7 +352,7 @@ from pytorch_lightning.overrides.data_parallel import ( LightningDataParallel, ) from pytorch_lightning.utilities.exceptions import MisconfigurationException -from pytorch_lightning.utilities.warnings import set_proc_rank +from pytorch_lightning.utilities.distributed import rank_zero_only try: from apex import amp @@ -506,7 +506,7 @@ class TrainerDPMixin(ABC): # track current tpu self.current_tpu_idx = tpu_core_idx self.proc_rank = self.tpu_local_core_rank - set_proc_rank(self.proc_rank) + rank_zero_only.rank = self.proc_rank # CHOOSE OPTIMIZER # allow for lr schedulers as well @@ -609,11 +609,7 @@ class TrainerDPMixin(ABC): # Update logger rank info from Horovod to avoid race conditions from different ranks # creating directories / writing files in the same locations. self.proc_rank = hvd.rank() - set_proc_rank(self.proc_rank) - if self.logger: - self.logger.rank = self.proc_rank - if model.logger: - model.logger.rank = self.proc_rank + rank_zero_only.rank = self.proc_rank with ExitStack() as stack: for optimizer in self.optimizers: diff --git a/pytorch_lightning/trainer/logging.py b/pytorch_lightning/trainer/logging.py index 3ba662fb..656c2026 100644 --- a/pytorch_lightning/trainer/logging.py +++ b/pytorch_lightning/trainer/logging.py @@ -33,7 +33,6 @@ class TrainerLoggingMixin(ABC): version=self.slurm_job_id, name='lightning_logs' ) - self.logger.rank = 0 elif logger is False: self.logger = None else: @@ -41,7 +40,6 @@ class TrainerLoggingMixin(ABC): self.logger = LoggerCollection(logger) else: self.logger = logger - self.logger.rank = 0 def log_metrics(self, metrics, grad_norm_dic, step=None): """Logs the metric dict passed in. diff --git a/pytorch_lightning/utilities/__init__.py b/pytorch_lightning/utilities/__init__.py index 469ae3ca..c8bc2805 100644 --- a/pytorch_lightning/utilities/__init__.py +++ b/pytorch_lightning/utilities/__init__.py @@ -1,3 +1,3 @@ """General utilities""" -from pytorch_lightning.utilities.warnings import rank_zero_warn +from pytorch_lightning.utilities.distributed import rank_zero_only, rank_zero_warn diff --git a/pytorch_lightning/utilities/distributed.py b/pytorch_lightning/utilities/distributed.py new file mode 100644 index 00000000..f4c942b9 --- /dev/null +++ b/pytorch_lightning/utilities/distributed.py @@ -0,0 +1,26 @@ +from functools import wraps +import warnings + + +def rank_zero_only(fn): + + @wraps(fn) + def wrapped_fn(*args, **kwargs): + if rank_zero_only.rank == 0: + return fn(*args, **kwargs) + + return wrapped_fn + + +try: + # add the attribute to the function but don't overwrite in case Trainer has already set it + getattr(rank_zero_only, 'rank') +except AttributeError: + rank_zero_only.rank = 0 + + +def _warn(*args, **kwargs): + warnings.warn(*args, **kwargs) + + +rank_zero_warn = rank_zero_only(_warn) diff --git a/pytorch_lightning/utilities/warnings.py b/pytorch_lightning/utilities/warnings.py deleted file mode 100644 index c1fa6fce..00000000 --- a/pytorch_lightning/utilities/warnings.py +++ /dev/null @@ -1,18 +0,0 @@ -"""Custom Lightning warnings""" - -import warnings - -_proc_rank = 0 - - -def set_proc_rank(value: int) -> None: - """Set the (sub)process rank.""" - global _proc_rank - _proc_rank = value - - -def rank_zero_warn(*args, **kwargs) -> None: - """Warning only if (sub)process has rank 0.""" - global _proc_rank - if _proc_rank == 0: - warnings.warn(*args, **kwargs) diff --git a/tests/loggers/test_base.py b/tests/loggers/test_base.py index 87032412..9bcbc2fb 100644 --- a/tests/loggers/test_base.py +++ b/tests/loggers/test_base.py @@ -5,7 +5,8 @@ import numpy as np import tests.base.utils as tutils from pytorch_lightning import Trainer -from pytorch_lightning.loggers import LightningLoggerBase, rank_zero_only, LoggerCollection +from pytorch_lightning.loggers import LightningLoggerBase, LoggerCollection +from pytorch_lightning.utilities import rank_zero_only from tests.base import LightningTestModel