diff --git a/pytorch_lightning/logging/__init__.py b/pytorch_lightning/logging/__init__.py index 8a6c3d03..fca3be61 100644 --- a/pytorch_lightning/logging/__init__.py +++ b/pytorch_lightning/logging/__init__.py @@ -168,7 +168,7 @@ Every k batches, lightning will write the new logs to disk from os import environ from .base import LightningLoggerBase, rank_zero_only -from .tensorboard import TensorboardLogger +from .tensorboard import TensorBoardLogger try: from .test_tube import TestTubeLogger diff --git a/pytorch_lightning/logging/tensorboard.py b/pytorch_lightning/logging/tensorboard.py index 338b02cc..26d626d1 100644 --- a/pytorch_lightning/logging/tensorboard.py +++ b/pytorch_lightning/logging/tensorboard.py @@ -8,8 +8,8 @@ from torch.utils.tensorboard import SummaryWriter from .base import LightningLoggerBase, rank_zero_only -class TensorboardLogger(LightningLoggerBase): - r"""Log to local file system in Tensorboard format +class TensorBoardLogger(LightningLoggerBase): + r"""Log to local file system in TensorBoard format Implemented using :class:`torch.utils.tensorboard.SummaryWriter`. Logs are saved to `os.path.join(save_dir, name, version)` @@ -18,7 +18,7 @@ class TensorboardLogger(LightningLoggerBase): .. code-block:: python - logger = TensorboardLogger("tb_logs", name="my_model") + logger = TensorBoardLogger("tb_logs", name="my_model") trainer = Trainer(logger=logger) trainer.train(model) @@ -35,7 +35,7 @@ class TensorboardLogger(LightningLoggerBase): super().__init__() self.save_dir = save_dir self._name = name - self._version = version if version is not None else None + self._version = version self._experiment = None self.kwargs = kwargs @@ -59,23 +59,35 @@ class TensorboardLogger(LightningLoggerBase): def log_hyperparams(self, params): if parse_version(torch.__version__) < parse_version("1.3.0"): warn( - f"Hyperparameter logging is not available for Torch version {torch.__version__}. " - "Skipping log_hyperparams. Upgrade to Torch 1.3.0 or above to enable " - "hyperparameter logging" + f"Hyperparameter logging is not available for Torch version {torch.__version__}." + " Skipping log_hyperparams. Upgrade to Torch 1.3.0 or above to enable" + " hyperparameter logging." ) + # TODO: some alternative should be added return - self.experiment.add_hparams(hparam_dict=vars(params)) + try: + # in case converting from namespace, todo: rather test if it is namespace + params = vars(params) + except TypeError: + pass + if params is not None: + # `add_hparams` requires both - hparams and metric + self.experiment.add_hparams(hparam_dict=dict(params), metric_dict={}) @rank_zero_only - def log_metrics(self, metrics, step_idx=None): + def log_metrics(self, metrics, step=None): for k, v in metrics.items(): if isinstance(v, torch.Tensor): v = v.item() - self.experiment.add_scalar(k, v, step_idx) + self.experiment.add_scalar(k, v, step) @rank_zero_only def save(self): - self.experiment.flush() + try: + self.experiment.flush() + except AttributeError: + # you are using PT version (