From e309b55b38b22a60cd4f5380635a34e307459777 Mon Sep 17 00:00:00 2001 From: Justus Schock <12886177+justusschock@users.noreply.github.com> Date: Mon, 27 Apr 2020 09:44:26 +0200 Subject: [PATCH 1/6] Update tensorboard.py --- pytorch_lightning/loggers/tensorboard.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/pytorch_lightning/loggers/tensorboard.py b/pytorch_lightning/loggers/tensorboard.py index 613262dd..36169e0a 100644 --- a/pytorch_lightning/loggers/tensorboard.py +++ b/pytorch_lightning/loggers/tensorboard.py @@ -101,7 +101,8 @@ class TensorBoardLogger(LightningLoggerBase): return self._experiment @rank_zero_only - def log_hyperparams(self, params: Union[Dict[str, Any], Namespace]) -> None: + def log_hyperparams(self, params: Union[Dict[str, Any], Namespace], + metrics: Optional[Dict[str, Any]] = None) -> None: params = self._convert_params(params) params = self._flatten_dict(params) sanitized_params = self._sanitize_params(params) @@ -114,7 +115,9 @@ class TensorBoardLogger(LightningLoggerBase): ) else: from torch.utils.tensorboard.summary import hparams - exp, ssi, sei = hparams(sanitized_params, {}) + if metrics is None: + metrics = {} + exp, ssi, sei = hparams(sanitized_params, metrics) writer = self.experiment._get_file_writer() writer.add_summary(exp) writer.add_summary(ssi) From 0ae7f479d302aca59994d2b6bc67e23ec8f4afbb Mon Sep 17 00:00:00 2001 From: Justus Schock <12886177+justusschock@users.noreply.github.com> Date: Mon, 27 Apr 2020 09:47:00 +0200 Subject: [PATCH 2/6] Update CHANGELOG.md --- CHANGELOG.md | 2 ++ 1 file changed, 2 insertions(+) diff --git a/CHANGELOG.md b/CHANGELOG.md index 5b8e2f29..212c7833 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,8 @@ The format is based on [Keep a Changelog](http://keepachangelog.com/en/1.0.0/). ### Added ### Changed + +- Allow logging of metrics togther with hparams ([#1630](https://github.com/PyTorchLightning/pytorch-lightning/pull/1630)) ### Deprecated From cdbf2f4a37b21a572205bb2c0bcab8ad41117863 Mon Sep 17 00:00:00 2001 From: Justus Schock <12886177+justusschock@users.noreply.github.com> Date: Mon, 27 Apr 2020 09:47:35 +0200 Subject: [PATCH 3/6] Update tensorboard.py --- pytorch_lightning/loggers/tensorboard.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pytorch_lightning/loggers/tensorboard.py b/pytorch_lightning/loggers/tensorboard.py index 36169e0a..fc33c9e9 100644 --- a/pytorch_lightning/loggers/tensorboard.py +++ b/pytorch_lightning/loggers/tensorboard.py @@ -101,7 +101,7 @@ class TensorBoardLogger(LightningLoggerBase): return self._experiment @rank_zero_only - def log_hyperparams(self, params: Union[Dict[str, Any], Namespace], + def log_hyperparams(self, params: Union[Dict[str, Any], Namespace], metrics: Optional[Dict[str, Any]] = None) -> None: params = self._convert_params(params) params = self._flatten_dict(params) From 312e394654f7fd4b259e84bec6e5475e4d837f80 Mon Sep 17 00:00:00 2001 From: Justus Schock <12886177+justusschock@users.noreply.github.com> Date: Mon, 27 Apr 2020 09:50:01 +0200 Subject: [PATCH 4/6] Update test_tensorboard.py --- tests/loggers/test_tensorboard.py | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) diff --git a/tests/loggers/test_tensorboard.py b/tests/loggers/test_tensorboard.py index 937a233c..4e5974b2 100644 --- a/tests/loggers/test_tensorboard.py +++ b/tests/loggers/test_tensorboard.py @@ -77,3 +77,19 @@ def test_tensorboard_log_hyperparams(tmpdir): "layer": torch.nn.BatchNorm1d } logger.log_hyperparams(hparams) + +def test_tensorboard_log_hparams_and_metrics + logger = TensorBoardLogger(tmpdir) + hparams = { + "float": 0.3, + "int": 1, + "string": "abc", + "bool": True, + "dict": {'a': {'b': 'c'}}, + "list": [1, 2, 3], + "namespace": Namespace(foo=Namespace(bar='buzz')), + "layer": torch.nn.BatchNorm1d + } + metrics = {'abc': torch.tensor([0.54])} + logger.log_hyperparams(hparams, metrics) + From 335819e1b7cd3e206a2cf46f4621e1bd771b4c2b Mon Sep 17 00:00:00 2001 From: Justus Schock <12886177+justusschock@users.noreply.github.com> Date: Mon, 27 Apr 2020 09:52:31 +0200 Subject: [PATCH 5/6] Update test_tensorboard.py --- tests/loggers/test_tensorboard.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/loggers/test_tensorboard.py b/tests/loggers/test_tensorboard.py index 4e5974b2..27ee234b 100644 --- a/tests/loggers/test_tensorboard.py +++ b/tests/loggers/test_tensorboard.py @@ -77,7 +77,8 @@ def test_tensorboard_log_hyperparams(tmpdir): "layer": torch.nn.BatchNorm1d } logger.log_hyperparams(hparams) - + + def test_tensorboard_log_hparams_and_metrics logger = TensorBoardLogger(tmpdir) hparams = { From ccd49cfbc5dd4f5290c3e04c3036a8d5ce47af7f Mon Sep 17 00:00:00 2001 From: Justus Schock Date: Mon, 27 Apr 2020 09:53:59 +0200 Subject: [PATCH 6/6] tests pep8 --- tests/loggers/test_tensorboard.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/tests/loggers/test_tensorboard.py b/tests/loggers/test_tensorboard.py index 27ee234b..a17cedc4 100644 --- a/tests/loggers/test_tensorboard.py +++ b/tests/loggers/test_tensorboard.py @@ -79,7 +79,7 @@ def test_tensorboard_log_hyperparams(tmpdir): logger.log_hyperparams(hparams) -def test_tensorboard_log_hparams_and_metrics +def test_tensorboard_log_hparams_and_metrics(tmpdir): logger = TensorBoardLogger(tmpdir) hparams = { "float": 0.3, @@ -93,4 +93,3 @@ def test_tensorboard_log_hparams_and_metrics } metrics = {'abc': torch.tensor([0.54])} logger.log_hyperparams(hparams, metrics) -