mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
Fix logger, tensorboard (#610)
* fix logger tests * fix missing flush * fix tensorboard * fix namespace * fix flush * fix add_hparams
This commit is contained in:
committed by
William Falcon
parent
4c7cfd3f12
commit
5d00e62047
@@ -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
|
||||
|
||||
@@ -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 (<v1.2) which does not have implemented flush
|
||||
self.experiment._get_file_writer().flush()
|
||||
|
||||
@rank_zero_only
|
||||
def finalize(self, status):
|
||||
|
||||
@@ -21,7 +21,7 @@ class TrainerLoggingMixin(ABC):
|
||||
self.use_ddp2 = None
|
||||
self.num_gpus = None
|
||||
|
||||
def log_metrics(self, metrics, grad_norm_dic):
|
||||
def log_metrics(self, metrics, grad_norm_dic, step=None):
|
||||
"""Logs the metric dict passed in.
|
||||
|
||||
:param metrics:
|
||||
@@ -41,9 +41,10 @@ class TrainerLoggingMixin(ABC):
|
||||
# turn all tensors to scalars
|
||||
scalar_metrics = self.metrics_to_scalars(metrics)
|
||||
|
||||
step = step if step is not None else self.global_step
|
||||
# log actual metrics
|
||||
if self.proc_rank == 0 and self.logger is not None:
|
||||
self.logger.log_metrics(scalar_metrics, step=self.global_step)
|
||||
self.logger.log_metrics(scalar_metrics, step=step)
|
||||
self.logger.save()
|
||||
|
||||
def add_tqdm_metrics(self, metrics):
|
||||
|
||||
+13
-15
@@ -9,7 +9,7 @@ from pytorch_lightning import Trainer
|
||||
from pytorch_lightning.logging import (
|
||||
LightningLoggerBase,
|
||||
rank_zero_only,
|
||||
TensorboardLogger,
|
||||
TensorBoardLogger,
|
||||
)
|
||||
from pytorch_lightning.testing import LightningTestModel
|
||||
|
||||
@@ -169,8 +169,8 @@ def test_comet_pickle(tmpdir, monkeypatch):
|
||||
except ModuleNotFoundError:
|
||||
return
|
||||
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
# hparams = tutils.get_hparams()
|
||||
# model = LightningTestModel(hparams)
|
||||
|
||||
comet_dir = os.path.join(tmpdir, "cometruns")
|
||||
|
||||
@@ -199,9 +199,9 @@ def test_tensorboard_logger(tmpdir):
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
logger = TensorboardLogger(save_dir=tmpdir, name="tensorboard_logger_test")
|
||||
logger = TensorBoardLogger(save_dir=tmpdir, name="tensorboard_logger_test")
|
||||
|
||||
trainer_options = dict(max_num_epochs=1, train_percent_check=0.01, logger=logger)
|
||||
trainer_options = dict(max_epochs=1, train_percent_check=0.01, logger=logger)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
result = trainer.fit(model)
|
||||
@@ -213,14 +213,12 @@ def test_tensorboard_logger(tmpdir):
|
||||
def test_tensorboard_pickle(tmpdir):
|
||||
"""Verify that pickling trainer with Tensorboard logger works."""
|
||||
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
# hparams = tutils.get_hparams()
|
||||
# model = LightningTestModel(hparams)
|
||||
|
||||
comet_dir = os.path.join(tmpdir, "cometruns")
|
||||
logger = TensorBoardLogger(save_dir=tmpdir, name="tensorboard_pickle_test")
|
||||
|
||||
logger = TensorboardLogger(save_dir=tmpdir, name="tensorboard_pickle_test")
|
||||
|
||||
trainer_options = dict(max_num_epochs=1, logger=logger)
|
||||
trainer_options = dict(max_epochs=1, logger=logger)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
pkl_bytes = pickle.dumps(trainer)
|
||||
@@ -235,7 +233,7 @@ def test_tensorboard_automatic_versioning(tmpdir):
|
||||
root_dir.mkdir("0")
|
||||
root_dir.mkdir("1")
|
||||
|
||||
logger = TensorboardLogger(save_dir=tmpdir, name="tb_versioning")
|
||||
logger = TensorBoardLogger(save_dir=tmpdir, name="tb_versioning")
|
||||
|
||||
assert logger.version == 2
|
||||
|
||||
@@ -248,14 +246,14 @@ def test_tensorboard_manual_versioning(tmpdir):
|
||||
root_dir.mkdir("1")
|
||||
root_dir.mkdir("2")
|
||||
|
||||
logger = TensorboardLogger(save_dir=tmpdir, name="tb_versioning", version=1)
|
||||
logger = TensorBoardLogger(save_dir=tmpdir, name="tb_versioning", version=1)
|
||||
|
||||
assert logger.version == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("step_idx", [10, None])
|
||||
def test_tensorboard_log_metrics(tmpdir, step_idx):
|
||||
logger = TensorboardLogger(tmpdir)
|
||||
logger = TensorBoardLogger(tmpdir)
|
||||
metrics = {
|
||||
"float": 0.3,
|
||||
"int": 1,
|
||||
@@ -266,7 +264,7 @@ def test_tensorboard_log_metrics(tmpdir, step_idx):
|
||||
|
||||
|
||||
def test_tensorboard_log_hyperparams(tmpdir):
|
||||
logger = TensorboardLogger(tmpdir)
|
||||
logger = TensorBoardLogger(tmpdir)
|
||||
hparams = {
|
||||
"float": 0.3,
|
||||
"int": 1,
|
||||
|
||||
Reference in New Issue
Block a user