From 01b8991c5a014fb1a8970bc6902a0e054f52e5df Mon Sep 17 00:00:00 2001 From: So Uchida Date: Thu, 19 Mar 2020 22:15:47 +0900 Subject: [PATCH] Support hierarchical dict (#1152) * Add support for hierarchical dict * Support nested Namespace * Add docstring * Migrate hparam flattening to each logger * Modify URLs in CHANGELOG * typo * Simplify the conditional branch about Namespace Co-Authored-By: Jirka Borovec * Update CHANGELOG.md Co-Authored-By: Jirka Borovec * added examples section to docstring * renamed _dict -> input_dict Co-authored-by: Jirka Borovec --- CHANGELOG.md | 1 + pytorch_lightning/loggers/base.py | 33 ++++++++++++++++++++++++ pytorch_lightning/loggers/comet.py | 1 + pytorch_lightning/loggers/mlflow.py | 1 + pytorch_lightning/loggers/neptune.py | 1 + pytorch_lightning/loggers/tensorboard.py | 1 + pytorch_lightning/loggers/test_tube.py | 1 + pytorch_lightning/loggers/trains.py | 8 +++--- tests/loggers/test_tensorboard.py | 3 ++- 9 files changed, 45 insertions(+), 5 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 176b1826..8489152b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,6 +8,7 @@ The format is based on [Keep a Changelog](http://keepachangelog.com/en/1.0.0/). ### Added +- Added support for hierarchical `dict` ([#1152](https://github.com/PyTorchLightning/pytorch-lightning/pull/1152)) - Added `TrainsLogger` class ([#1122](https://github.com/PyTorchLightning/pytorch-lightning/pull/1122)) - Added type hints to `pytorch_lightning.core` ([#946](https://github.com/PyTorchLightning/pytorch-lightning/pull/946)) - Added support for IterableDataset in validation and testing ([#1104](https://github.com/PyTorchLightning/pytorch-lightning/pull/1104)) diff --git a/pytorch_lightning/loggers/base.py b/pytorch_lightning/loggers/base.py index 1d0c41cf..38800e49 100644 --- a/pytorch_lightning/loggers/base.py +++ b/pytorch_lightning/loggers/base.py @@ -53,6 +53,39 @@ class LightningLoggerBase(ABC): return params + @staticmethod + def _flatten_dict(params: Dict[str, Any], delimiter: str = '/') -> Dict[str, Any]: + """Flatten hierarchical dict e.g. {'a': {'b': 'c'}} -> {'a/b': 'c'}. + + Args: + params: Dictionary contains hparams + delimiter: Delimiter to express the hierarchy. Defaults to '/'. + + Returns: + Flatten dict. + + Examples: + >>> LightningLoggerBase._flatten_dict({'a': {'b': 'c'}}) + {'a/b': 'c'} + >>> LightningLoggerBase._flatten_dict({'a': {'b': 123}}) + {'a/b': 123} + """ + + def _dict_generator(input_dict, prefixes=None): + prefixes = prefixes[:] if prefixes else [] + if isinstance(input_dict, dict): + for key, value in input_dict.items(): + if isinstance(value, (dict, Namespace)): + value = vars(value) if isinstance(value, Namespace) else value + for d in _dict_generator(value, prefixes + [key]): + yield d + else: + yield prefixes + [key, value if value is not None else str(None)] + else: + yield prefixes + [input_dict if input_dict is None else str(input_dict)] + + return {delimiter.join(keys): val for *keys, val in _dict_generator(params)} + @staticmethod def _sanitize_params(params: Dict[str, Any]) -> Dict[str, Any]: """Returns params with non-primitvies converted to strings for logging diff --git a/pytorch_lightning/loggers/comet.py b/pytorch_lightning/loggers/comet.py index e67b1b97..0109f9db 100644 --- a/pytorch_lightning/loggers/comet.py +++ b/pytorch_lightning/loggers/comet.py @@ -163,6 +163,7 @@ class CometLogger(LightningLoggerBase): @rank_zero_only def log_hyperparams(self, params: Union[Dict[str, Any], Namespace]) -> None: params = self._convert_params(params) + params = self._flatten_dict(params) self.experiment.log_parameters(params) @rank_zero_only diff --git a/pytorch_lightning/loggers/mlflow.py b/pytorch_lightning/loggers/mlflow.py index 36f50b41..6006cab4 100644 --- a/pytorch_lightning/loggers/mlflow.py +++ b/pytorch_lightning/loggers/mlflow.py @@ -89,6 +89,7 @@ class MLFlowLogger(LightningLoggerBase): @rank_zero_only def log_hyperparams(self, params: Union[Dict[str, Any], Namespace]) -> None: params = self._convert_params(params) + params = self._flatten_dict(params) for k, v in params.items(): self.experiment.log_param(self.run_id, k, v) diff --git a/pytorch_lightning/loggers/neptune.py b/pytorch_lightning/loggers/neptune.py index 12062026..2282ff93 100644 --- a/pytorch_lightning/loggers/neptune.py +++ b/pytorch_lightning/loggers/neptune.py @@ -222,6 +222,7 @@ class NeptuneLogger(LightningLoggerBase): @rank_zero_only def log_hyperparams(self, params: Union[Dict[str, Any], Namespace]) -> None: params = self._convert_params(params) + params = self._flatten_dict(params) for key, val in params.items(): self.experiment.set_property(f'param__{key}', val) diff --git a/pytorch_lightning/loggers/tensorboard.py b/pytorch_lightning/loggers/tensorboard.py index b0598a9e..07aebe19 100644 --- a/pytorch_lightning/loggers/tensorboard.py +++ b/pytorch_lightning/loggers/tensorboard.py @@ -99,6 +99,7 @@ class TensorBoardLogger(LightningLoggerBase): @rank_zero_only def log_hyperparams(self, params: Union[Dict[str, Any], Namespace]) -> None: params = self._convert_params(params) + params = self._flatten_dict(params) sanitized_params = self._sanitize_params(params) if parse_version(torch.__version__) < parse_version("1.3.0"): diff --git a/pytorch_lightning/loggers/test_tube.py b/pytorch_lightning/loggers/test_tube.py index 4fc421a3..f6ec3315 100644 --- a/pytorch_lightning/loggers/test_tube.py +++ b/pytorch_lightning/loggers/test_tube.py @@ -96,6 +96,7 @@ class TestTubeLogger(LightningLoggerBase): # TODO: HACK figure out where this is being set to true self.experiment.debug = self.debug params = self._convert_params(params) + params = self._flatten_dict(params) self.experiment.argparse(Namespace(**params)) @rank_zero_only diff --git a/pytorch_lightning/loggers/trains.py b/pytorch_lightning/loggers/trains.py index 2f49928d..80f5f022 100644 --- a/pytorch_lightning/loggers/trains.py +++ b/pytorch_lightning/loggers/trains.py @@ -130,10 +130,10 @@ class TrainsLogger(LightningLoggerBase): return None if not params: return - if isinstance(params, dict): - self._trains.connect(params) - else: - self._trains.connect(vars(params)) + + params = self._convert_params(params) + params = self._flatten_dict(params) + self._trains.connect(params) @rank_zero_only def log_metrics(self, metrics: Dict[str, float], step: Optional[int] = None) -> None: diff --git a/tests/loggers/test_tensorboard.py b/tests/loggers/test_tensorboard.py index ba7fd5d4..220cdeb5 100644 --- a/tests/loggers/test_tensorboard.py +++ b/tests/loggers/test_tensorboard.py @@ -108,8 +108,9 @@ def test_tensorboard_log_hyperparams(tmpdir): "int": 1, "string": "abc", "bool": True, + "dict": {'a': {'b': 'c'}}, "list": [1, 2, 3], - "namespace": Namespace(foo=3), + "namespace": Namespace(foo=Namespace(bar='buzz')), "layer": torch.nn.BatchNorm1d } logger.log_hyperparams(hparams)