mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
Logger tests and fixes (#1009)
* Refactor logger tests * Update and add tests for wandb logger * Update and add tests for logger bases * Update and add tests for mlflow logger * Improve coverage * Updates * Update CHANGELOG * Updates * Fix style * Fix style * Updates
This commit is contained in:
@@ -47,6 +47,8 @@ The format is based on [Keep a Changelog](http://keepachangelog.com/en/1.0.0/).
|
||||
- Fixed a bug where the model checkpointer didn't write to the same directory as the logger ([#771](https://github.com/PyTorchLightning/pytorch-lightning/pull/771))
|
||||
- Fixed a bug where the `TensorBoardLogger` class would create an additional empty log file during fitting ([#777](https://github.com/PyTorchLightning/pytorch-lightning/pull/777))
|
||||
- Fixed a bug where `global_step` was advanced incorrectly when using `accumulate_grad_batches > 1` ([#832](https://github.com/PyTorchLightning/pytorch-lightning/pull/832))
|
||||
- Fixed a bug when calling `self.logger.experiment` with multiple loggers ([#1009](https://github.com/PyTorchLightning/pytorch-lightning/pull/1009))
|
||||
- Fixed a bug when calling `logger.append_tags` on a `NeptuneLogger` with a single tag ([#1009](https://github.com/PyTorchLightning/pytorch-lightning/pull/1009))
|
||||
|
||||
## [0.6.0] - 2020-01-21
|
||||
|
||||
|
||||
@@ -105,7 +105,7 @@ class LoggerCollection(LightningLoggerBase):
|
||||
|
||||
@property
|
||||
def experiment(self) -> List[Any]:
|
||||
return [logger.experiment() for logger in self._logger_iterable]
|
||||
return [logger.experiment for logger in self._logger_iterable]
|
||||
|
||||
def log_metrics(self, metrics: Dict[str, float], step: Optional[int] = None):
|
||||
[logger.log_metrics(metrics, step) for logger in self._logger_iterable]
|
||||
@@ -122,11 +122,7 @@ class LoggerCollection(LightningLoggerBase):
|
||||
def close(self):
|
||||
[logger.close() for logger in self._logger_iterable]
|
||||
|
||||
@property
|
||||
def rank(self) -> int:
|
||||
return self._rank
|
||||
|
||||
@rank.setter
|
||||
@LightningLoggerBase.rank.setter
|
||||
def rank(self, value: int):
|
||||
self._rank = value
|
||||
for logger in self._logger_iterable:
|
||||
|
||||
@@ -7,7 +7,7 @@ CometLogger
|
||||
"""
|
||||
import argparse
|
||||
from logging import getLogger
|
||||
from typing import Optional, Union, Dict
|
||||
from typing import Optional, Dict, Union
|
||||
|
||||
try:
|
||||
from comet_ml import Experiment as CometExperiment
|
||||
@@ -20,8 +20,10 @@ try:
|
||||
# For more information, see: https://www.comet.ml/docs/python-sdk/releases/#release-300
|
||||
from comet_ml.papi import API
|
||||
except ImportError:
|
||||
raise ImportError('Missing comet_ml package.')
|
||||
raise ImportError('You want to use `comet_ml` logger which is not installed yet,'
|
||||
' install it with `pip install comet-ml`.')
|
||||
|
||||
import torch
|
||||
from torch import is_tensor
|
||||
|
||||
from pytorch_lightning.utilities.debugging import MisconfigurationException
|
||||
@@ -87,11 +89,7 @@ class CometLogger(LightningLoggerBase):
|
||||
self._experiment = None
|
||||
|
||||
# Determine online or offline mode based on which arguments were passed to CometLogger
|
||||
if save_dir is not None and api_key is not None:
|
||||
# If arguments are passed for both save_dir and api_key, preference is given to online mode
|
||||
self.mode = "online"
|
||||
self.api_key = api_key
|
||||
elif api_key is not None:
|
||||
if api_key is not None:
|
||||
self.mode = "online"
|
||||
self.api_key = api_key
|
||||
elif save_dir is not None:
|
||||
@@ -168,7 +166,11 @@ class CometLogger(LightningLoggerBase):
|
||||
self.experiment.log_parameters(vars(params))
|
||||
|
||||
@rank_zero_only
|
||||
def log_metrics(self, metrics: Dict[str, float], step: Optional[int] = None):
|
||||
def log_metrics(
|
||||
self,
|
||||
metrics: Dict[str, Union[torch.Tensor, float]],
|
||||
step: Optional[int] = None
|
||||
):
|
||||
# Comet.ml expects metrics to be a dictionary of detached tensors on CPU
|
||||
for key, val in metrics.items():
|
||||
if is_tensor(val):
|
||||
|
||||
@@ -31,7 +31,8 @@ from typing import Optional, Dict, Any
|
||||
try:
|
||||
import mlflow
|
||||
except ImportError:
|
||||
raise ImportError('Missing mlflow package.')
|
||||
raise ImportError('You want to use `mlflow` logger which is not installed yet,'
|
||||
' install it with `pip install mlflow`.')
|
||||
|
||||
from .base import LightningLoggerBase, rank_zero_only
|
||||
|
||||
@@ -79,7 +80,7 @@ class MLFlowLogger(LightningLoggerBase):
|
||||
if expt:
|
||||
self._expt_id = expt.experiment_id
|
||||
else:
|
||||
logger.warning(f"Experiment with name {self.experiment_name} not found. Creating it.")
|
||||
logger.warning(f'Experiment with name {self.experiment_name} not found. Creating it.')
|
||||
self._expt_id = self._mlflow_client.create_experiment(name=self.experiment_name)
|
||||
|
||||
run = self._mlflow_client.create_run(experiment_id=self._expt_id, tags=self.tags)
|
||||
@@ -96,9 +97,7 @@ class MLFlowLogger(LightningLoggerBase):
|
||||
timestamp_ms = int(time() * 1000)
|
||||
for k, v in metrics.items():
|
||||
if isinstance(v, str):
|
||||
logger.warning(
|
||||
f"Discarding metric with string value {k}={v}"
|
||||
)
|
||||
logger.warning(f'Discarding metric with string value {k}={v}.')
|
||||
continue
|
||||
self.experiment.log_metric(self.run_id, k, v, timestamp_ms, step)
|
||||
|
||||
@@ -106,7 +105,7 @@ class MLFlowLogger(LightningLoggerBase):
|
||||
pass
|
||||
|
||||
@rank_zero_only
|
||||
def finalize(self, status: str = "FINISHED"):
|
||||
def finalize(self, status: str = 'FINISHED'):
|
||||
if status == 'success':
|
||||
status = 'FINISHED'
|
||||
self.experiment.set_terminated(self.run_id, status)
|
||||
|
||||
@@ -15,11 +15,11 @@ try:
|
||||
from neptune.experiments import Experiment
|
||||
except ImportError:
|
||||
raise ImportError('You want to use `neptune` logger which is not installed yet,'
|
||||
' please install it e.g. `pip install neptune-client`.')
|
||||
' install it with `pip install neptune-client`.')
|
||||
|
||||
import torch
|
||||
from torch import is_tensor
|
||||
|
||||
# from .base import LightningLoggerBase, rank_zero_only
|
||||
from pytorch_lightning.loggers.base import LightningLoggerBase, rank_zero_only
|
||||
|
||||
logger = getLogger(__name__)
|
||||
@@ -130,15 +130,15 @@ class NeptuneLogger(LightningLoggerBase):
|
||||
self._kwargs = kwargs
|
||||
|
||||
if offline_mode:
|
||||
self.mode = "offline"
|
||||
self.mode = 'offline'
|
||||
neptune.init(project_qualified_name='dry-run/project',
|
||||
backend=neptune.OfflineBackend())
|
||||
else:
|
||||
self.mode = "online"
|
||||
self.mode = 'online'
|
||||
neptune.init(api_token=self.api_key,
|
||||
project_qualified_name=self.project_name)
|
||||
|
||||
logger.info(f"NeptuneLogger was initialized in {self.mode} mode")
|
||||
logger.info(f'NeptuneLogger was initialized in {self.mode} mode')
|
||||
|
||||
@property
|
||||
def experiment(self) -> Experiment:
|
||||
@@ -166,25 +166,22 @@ class NeptuneLogger(LightningLoggerBase):
|
||||
@rank_zero_only
|
||||
def log_hyperparams(self, params: argparse.Namespace):
|
||||
for key, val in vars(params).items():
|
||||
self.experiment.set_property(f"param__{key}", val)
|
||||
self.experiment.set_property(f'param__{key}', val)
|
||||
|
||||
@rank_zero_only
|
||||
def log_metrics(self, metrics: Dict[str, float], step: Optional[int] = None):
|
||||
def log_metrics(
|
||||
self,
|
||||
metrics: Dict[str, Union[torch.Tensor, float]],
|
||||
step: Optional[int] = None
|
||||
):
|
||||
"""Log metrics (numeric values) in Neptune experiments
|
||||
|
||||
Args:
|
||||
metrics: Dictionary with metric names as keys and measured quantities as values
|
||||
step: Step number at which the metrics should be recorded, must be strictly increasing
|
||||
"""
|
||||
|
||||
for key, val in metrics.items():
|
||||
if is_tensor(val):
|
||||
val = val.cpu().detach()
|
||||
|
||||
if step is None:
|
||||
self.experiment.log_metric(key, val)
|
||||
else:
|
||||
self.experiment.log_metric(key, x=step, y=val)
|
||||
self.log_metric(key, val, step=step)
|
||||
|
||||
@rank_zero_only
|
||||
def finalize(self, status: str):
|
||||
@@ -192,20 +189,25 @@ class NeptuneLogger(LightningLoggerBase):
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
if self.mode == "offline":
|
||||
return "offline-name"
|
||||
if self.mode == 'offline':
|
||||
return 'offline-name'
|
||||
else:
|
||||
return self.experiment.name
|
||||
|
||||
@property
|
||||
def version(self) -> str:
|
||||
if self.mode == "offline":
|
||||
return "offline-id-1234"
|
||||
if self.mode == 'offline':
|
||||
return 'offline-id-1234'
|
||||
else:
|
||||
return self.experiment.id
|
||||
|
||||
@rank_zero_only
|
||||
def log_metric(self, metric_name: str, metric_value: float, step: Optional[int] = None):
|
||||
def log_metric(
|
||||
self,
|
||||
metric_name: str,
|
||||
metric_value: Union[torch.Tensor, float, str],
|
||||
step: Optional[int] = None
|
||||
):
|
||||
"""Log metrics (numeric values) in Neptune experiments
|
||||
|
||||
Args:
|
||||
@@ -213,6 +215,9 @@ class NeptuneLogger(LightningLoggerBase):
|
||||
metric_value: The value of the log (data-point).
|
||||
step: Step number at which the metrics should be recorded, must be strictly increasing
|
||||
"""
|
||||
if is_tensor(metric_value):
|
||||
metric_value = metric_value.cpu().detach()
|
||||
|
||||
if step is None:
|
||||
self.experiment.log_metric(metric_name, metric_value)
|
||||
else:
|
||||
@@ -227,10 +232,7 @@ class NeptuneLogger(LightningLoggerBase):
|
||||
text: The value of the log (data-point).
|
||||
step: Step number at which the metrics should be recorded, must be strictly increasing
|
||||
"""
|
||||
if step is None:
|
||||
self.experiment.log_metric(log_name, text)
|
||||
else:
|
||||
self.experiment.log_metric(log_name, x=step, y=text)
|
||||
self.log_metric(log_name, text, step=step)
|
||||
|
||||
@rank_zero_only
|
||||
def log_image(self, log_name: str, image: Union[str, Any], step: Optional[int] = None):
|
||||
@@ -277,6 +279,6 @@ class NeptuneLogger(LightningLoggerBase):
|
||||
If multiple - comma separated - str are passed, all of them are added as tags.
|
||||
If list of str is passed, all elements of the list are added as tags.
|
||||
"""
|
||||
if not isinstance(tags, Iterable):
|
||||
if str(tags) == tags:
|
||||
tags = [tags] # make it as an iterable is if it is not yet
|
||||
self.experiment.append_tags(*tags)
|
||||
|
||||
@@ -44,7 +44,10 @@ class TensorBoardLogger(LightningLoggerBase):
|
||||
"""
|
||||
NAME_CSV_TAGS = 'meta_tags.csv'
|
||||
|
||||
def __init__(self, save_dir: str, name: str = "default", version: Optional[Union[int, str]] = None, **kwargs):
|
||||
def __init__(
|
||||
self, save_dir: str, name: Optional[str] = "default",
|
||||
version: Optional[Union[int, str]] = None, **kwargs
|
||||
):
|
||||
super().__init__()
|
||||
self.save_dir = save_dir
|
||||
self._name = name
|
||||
|
||||
@@ -4,7 +4,8 @@ from typing import Optional, Dict, Any
|
||||
try:
|
||||
from test_tube import Experiment
|
||||
except ImportError:
|
||||
raise ImportError('Missing test-tube package.')
|
||||
raise ImportError('You want to use `test_tube` logger which is not installed yet,'
|
||||
' install it with `pip install test-tube`.')
|
||||
|
||||
from .base import LightningLoggerBase, rank_zero_only
|
||||
|
||||
|
||||
@@ -9,12 +9,14 @@ import argparse
|
||||
import os
|
||||
from typing import Optional, List, Dict
|
||||
|
||||
import torch.nn as nn
|
||||
|
||||
try:
|
||||
import wandb
|
||||
from wandb.wandb_run import Run
|
||||
except ImportError:
|
||||
raise ImportError('You want to use `wandb` logger which is not installed yet,'
|
||||
' please install it e.g. `pip install wandb`.')
|
||||
' install it with `pip install wandb`.')
|
||||
|
||||
from .base import LightningLoggerBase, rank_zero_only
|
||||
|
||||
@@ -50,7 +52,7 @@ class WandbLogger(LightningLoggerBase):
|
||||
super().__init__()
|
||||
self._name = name
|
||||
self._save_dir = save_dir
|
||||
self._anonymous = "allow" if anonymous else None
|
||||
self._anonymous = 'allow' if anonymous else None
|
||||
self._id = version or id
|
||||
self._tags = tags
|
||||
self._project = project
|
||||
@@ -79,14 +81,14 @@ class WandbLogger(LightningLoggerBase):
|
||||
"""
|
||||
if self._experiment is None:
|
||||
if self._offline:
|
||||
os.environ["WANDB_MODE"] = "dryrun"
|
||||
os.environ['WANDB_MODE'] = 'dryrun'
|
||||
self._experiment = wandb.init(
|
||||
name=self._name, dir=self._save_dir, project=self._project, anonymous=self._anonymous,
|
||||
id=self._id, resume="allow", tags=self._tags, entity=self._entity)
|
||||
id=self._id, resume='allow', tags=self._tags, entity=self._entity)
|
||||
return self._experiment
|
||||
|
||||
def watch(self, model, log="gradients", log_freq=100):
|
||||
wandb.watch(model, log, log_freq)
|
||||
def watch(self, model: nn.Module, log: str = 'gradients', log_freq: int = 100):
|
||||
wandb.watch(model, log=log, log_freq=log_freq)
|
||||
|
||||
@rank_zero_only
|
||||
def log_hyperparams(self, params: argparse.Namespace):
|
||||
@@ -94,12 +96,10 @@ class WandbLogger(LightningLoggerBase):
|
||||
|
||||
@rank_zero_only
|
||||
def log_metrics(self, metrics: Dict[str, float], step: Optional[int] = None):
|
||||
metrics["global_step"] = step
|
||||
if step is not None:
|
||||
metrics['global_step'] = step
|
||||
self.experiment.log(metrics)
|
||||
|
||||
def save(self):
|
||||
pass
|
||||
|
||||
@rank_zero_only
|
||||
def finalize(self, status: str = 'success'):
|
||||
try:
|
||||
|
||||
@@ -0,0 +1,152 @@
|
||||
import pickle
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import tests.models.utils as tutils
|
||||
from pytorch_lightning import Trainer
|
||||
from pytorch_lightning.loggers import LightningLoggerBase, rank_zero_only, LoggerCollection
|
||||
from tests.models import LightningTestModel
|
||||
|
||||
|
||||
def test_logger_collection():
|
||||
mock1 = MagicMock()
|
||||
mock2 = MagicMock()
|
||||
|
||||
logger = LoggerCollection([mock1, mock2])
|
||||
|
||||
assert logger[0] == mock1
|
||||
assert logger[1] == mock2
|
||||
|
||||
assert logger.experiment[0] == mock1.experiment
|
||||
assert logger.experiment[1] == mock2.experiment
|
||||
|
||||
logger.close()
|
||||
mock1.close.assert_called_once()
|
||||
mock2.close.assert_called_once()
|
||||
|
||||
|
||||
class CustomLogger(LightningLoggerBase):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.hparams_logged = None
|
||||
self.metrics_logged = None
|
||||
self.finalized = False
|
||||
|
||||
@property
|
||||
def experiment(self):
|
||||
return 'test'
|
||||
|
||||
@rank_zero_only
|
||||
def log_hyperparams(self, params):
|
||||
self.hparams_logged = params
|
||||
|
||||
@rank_zero_only
|
||||
def log_metrics(self, metrics, step):
|
||||
self.metrics_logged = metrics
|
||||
|
||||
@rank_zero_only
|
||||
def finalize(self, status):
|
||||
self.finalized_status = status
|
||||
|
||||
@property
|
||||
def name(self):
|
||||
return "name"
|
||||
|
||||
@property
|
||||
def version(self):
|
||||
return "1"
|
||||
|
||||
|
||||
def test_custom_logger(tmpdir):
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
logger = CustomLogger()
|
||||
|
||||
trainer_options = dict(
|
||||
max_epochs=1,
|
||||
train_percent_check=0.05,
|
||||
logger=logger,
|
||||
default_save_path=tmpdir
|
||||
)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
result = trainer.fit(model)
|
||||
assert result == 1, "Training failed"
|
||||
assert logger.hparams_logged == hparams
|
||||
assert logger.metrics_logged != {}
|
||||
assert logger.finalized_status == "success"
|
||||
|
||||
|
||||
def test_multiple_loggers(tmpdir):
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
logger1 = CustomLogger()
|
||||
logger2 = CustomLogger()
|
||||
|
||||
trainer_options = dict(
|
||||
max_epochs=1,
|
||||
train_percent_check=0.05,
|
||||
logger=[logger1, logger2],
|
||||
default_save_path=tmpdir
|
||||
)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
result = trainer.fit(model)
|
||||
assert result == 1, "Training failed"
|
||||
|
||||
assert logger1.hparams_logged == hparams
|
||||
assert logger1.metrics_logged != {}
|
||||
assert logger1.finalized_status == "success"
|
||||
|
||||
assert logger2.hparams_logged == hparams
|
||||
assert logger2.metrics_logged != {}
|
||||
assert logger2.finalized_status == "success"
|
||||
|
||||
|
||||
def test_multiple_loggers_pickle(tmpdir):
|
||||
"""Verify that pickling trainer with multiple loggers works."""
|
||||
|
||||
logger1 = CustomLogger()
|
||||
logger2 = CustomLogger()
|
||||
|
||||
trainer_options = dict(max_epochs=1, logger=[logger1, logger2])
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
pkl_bytes = pickle.dumps(trainer)
|
||||
trainer2 = pickle.loads(pkl_bytes)
|
||||
trainer2.logger.log_metrics({"acc": 1.0}, 0)
|
||||
|
||||
assert logger1.metrics_logged != {}
|
||||
assert logger2.metrics_logged != {}
|
||||
|
||||
|
||||
def test_adding_step_key(tmpdir):
|
||||
logged_step = 0
|
||||
|
||||
def _validation_end(outputs):
|
||||
nonlocal logged_step
|
||||
logged_step += 1
|
||||
return {"log": {"step": logged_step, "val_acc": logged_step / 10}}
|
||||
|
||||
def _log_metrics_decorator(log_metrics_fn):
|
||||
def decorated(metrics, step):
|
||||
if "val_acc" in metrics:
|
||||
assert step == logged_step
|
||||
return log_metrics_fn(metrics, step)
|
||||
|
||||
return decorated
|
||||
|
||||
model, hparams = tutils.get_model()
|
||||
model.validation_end = _validation_end
|
||||
trainer_options = dict(
|
||||
max_epochs=4,
|
||||
default_save_path=tmpdir,
|
||||
train_percent_check=0.001,
|
||||
val_percent_check=0.01,
|
||||
num_sanity_val_steps=0
|
||||
)
|
||||
trainer = Trainer(**trainer_options)
|
||||
trainer.logger.log_metrics = _log_metrics_decorator(trainer.logger.log_metrics)
|
||||
trainer.fit(model)
|
||||
@@ -0,0 +1,158 @@
|
||||
import os
|
||||
import pickle
|
||||
|
||||
import torch
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
import tests.models.utils as tutils
|
||||
from pytorch_lightning import Trainer
|
||||
from pytorch_lightning.utilities.debugging import MisconfigurationException
|
||||
from pytorch_lightning.loggers import CometLogger
|
||||
from tests.models import LightningTestModel
|
||||
|
||||
|
||||
def test_comet_logger(tmpdir, monkeypatch):
|
||||
"""Verify that basic functionality of Comet.ml logger works."""
|
||||
|
||||
# prevent comet logger from trying to print at exit, since
|
||||
# pytest's stdout/stderr redirection breaks it
|
||||
import atexit
|
||||
monkeypatch.setattr(atexit, 'register', lambda _: None)
|
||||
|
||||
tutils.reset_seed()
|
||||
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
comet_dir = os.path.join(tmpdir, 'cometruns')
|
||||
|
||||
# We test CometLogger in offline mode with local saves
|
||||
logger = CometLogger(
|
||||
save_dir=comet_dir,
|
||||
project_name='general',
|
||||
workspace='dummy-test',
|
||||
)
|
||||
|
||||
trainer_options = dict(
|
||||
default_save_path=tmpdir,
|
||||
max_epochs=1,
|
||||
train_percent_check=0.05,
|
||||
logger=logger
|
||||
)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
result = trainer.fit(model)
|
||||
trainer.logger.log_metrics({'acc': torch.ones(1)})
|
||||
|
||||
assert result == 1, 'Training failed'
|
||||
|
||||
|
||||
def test_comet_logger_online():
|
||||
"""Test comet online with mocks."""
|
||||
# Test api_key given
|
||||
with patch('pytorch_lightning.loggers.comet.CometExperiment') as comet:
|
||||
logger = CometLogger(
|
||||
api_key='key',
|
||||
workspace='dummy-test',
|
||||
project_name='general'
|
||||
)
|
||||
|
||||
_ = logger.experiment
|
||||
|
||||
comet.assert_called_once_with(
|
||||
api_key='key',
|
||||
workspace='dummy-test',
|
||||
project_name='general'
|
||||
)
|
||||
|
||||
# Test both given
|
||||
with patch('pytorch_lightning.loggers.comet.CometExperiment') as comet:
|
||||
logger = CometLogger(
|
||||
save_dir='test',
|
||||
api_key='key',
|
||||
workspace='dummy-test',
|
||||
project_name='general'
|
||||
)
|
||||
|
||||
_ = logger.experiment
|
||||
|
||||
comet.assert_called_once_with(
|
||||
api_key='key',
|
||||
workspace='dummy-test',
|
||||
project_name='general'
|
||||
)
|
||||
|
||||
# Test neither given
|
||||
with pytest.raises(MisconfigurationException):
|
||||
CometLogger(
|
||||
workspace='dummy-test',
|
||||
project_name='general'
|
||||
)
|
||||
|
||||
# Test already exists
|
||||
with patch('pytorch_lightning.loggers.comet.CometExistingExperiment') as comet_existing:
|
||||
logger = CometLogger(
|
||||
experiment_key='test',
|
||||
experiment_name='experiment',
|
||||
api_key='key',
|
||||
workspace='dummy-test',
|
||||
project_name='general'
|
||||
)
|
||||
|
||||
_ = logger.experiment
|
||||
|
||||
comet_existing.assert_called_once_with(
|
||||
api_key='key',
|
||||
workspace='dummy-test',
|
||||
project_name='general',
|
||||
previous_experiment='test'
|
||||
)
|
||||
|
||||
comet_existing().set_name.assert_called_once_with('experiment')
|
||||
|
||||
with patch('pytorch_lightning.loggers.comet.API') as api:
|
||||
CometLogger(
|
||||
api_key='key',
|
||||
workspace='dummy-test',
|
||||
project_name='general',
|
||||
rest_api_key='rest'
|
||||
)
|
||||
|
||||
api.assert_called_once_with('rest')
|
||||
|
||||
|
||||
def test_comet_pickle(tmpdir, monkeypatch):
|
||||
"""Verify that pickling trainer with comet logger works."""
|
||||
|
||||
# prevent comet logger from trying to print at exit, since
|
||||
# pytest's stdout/stderr redirection breaks it
|
||||
import atexit
|
||||
monkeypatch.setattr(atexit, 'register', lambda _: None)
|
||||
|
||||
tutils.reset_seed()
|
||||
|
||||
# hparams = tutils.get_hparams()
|
||||
# model = LightningTestModel(hparams)
|
||||
|
||||
comet_dir = os.path.join(tmpdir, 'cometruns')
|
||||
|
||||
# We test CometLogger in offline mode with local saves
|
||||
logger = CometLogger(
|
||||
save_dir=comet_dir,
|
||||
project_name='general',
|
||||
workspace='dummy-test',
|
||||
)
|
||||
|
||||
trainer_options = dict(
|
||||
default_save_path=tmpdir,
|
||||
max_epochs=1,
|
||||
logger=logger
|
||||
)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
pkl_bytes = pickle.dumps(trainer)
|
||||
trainer2 = pickle.loads(pkl_bytes)
|
||||
trainer2.logger.log_metrics({'acc': 1.0})
|
||||
@@ -0,0 +1,54 @@
|
||||
import os
|
||||
import pickle
|
||||
|
||||
import tests.models.utils as tutils
|
||||
from pytorch_lightning import Trainer
|
||||
from pytorch_lightning.loggers import MLFlowLogger
|
||||
from tests.models import LightningTestModel
|
||||
|
||||
|
||||
def test_mlflow_logger(tmpdir):
|
||||
"""Verify that basic functionality of mlflow logger works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
mlflow_dir = os.path.join(tmpdir, 'mlruns')
|
||||
logger = MLFlowLogger('test', tracking_uri=f'file:{os.sep * 2}{mlflow_dir}')
|
||||
|
||||
# Test already exists
|
||||
logger2 = MLFlowLogger('test', tracking_uri=f'file:{os.sep * 2}{mlflow_dir}')
|
||||
_ = logger2.run_id
|
||||
|
||||
# Try logging string
|
||||
logger.log_metrics({'acc': 'test'})
|
||||
|
||||
trainer_options = dict(
|
||||
default_save_path=tmpdir,
|
||||
max_epochs=1,
|
||||
train_percent_check=0.05,
|
||||
logger=logger
|
||||
)
|
||||
trainer = Trainer(**trainer_options)
|
||||
result = trainer.fit(model)
|
||||
|
||||
assert result == 1, 'Training failed'
|
||||
|
||||
|
||||
def test_mlflow_pickle(tmpdir):
|
||||
"""Verify that pickling trainer with mlflow logger works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
mlflow_dir = os.path.join(tmpdir, 'mlruns')
|
||||
logger = MLFlowLogger('test', tracking_uri=f'file:{os.sep * 2}{mlflow_dir}')
|
||||
trainer_options = dict(
|
||||
default_save_path=tmpdir,
|
||||
max_epochs=1,
|
||||
logger=logger
|
||||
)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
pkl_bytes = pickle.dumps(trainer)
|
||||
trainer2 = pickle.loads(pkl_bytes)
|
||||
trainer2.logger.log_metrics({'acc': 1.0})
|
||||
@@ -0,0 +1,99 @@
|
||||
import pickle
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
import tests.models.utils as tutils
|
||||
from pytorch_lightning import Trainer
|
||||
from pytorch_lightning.loggers import NeptuneLogger
|
||||
from tests.models import LightningTestModel
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def test_neptune_logger(tmpdir):
|
||||
"""Verify that basic functionality of neptune logger works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
logger = NeptuneLogger(offline_mode=True)
|
||||
|
||||
trainer_options = dict(
|
||||
default_save_path=tmpdir,
|
||||
max_epochs=1,
|
||||
train_percent_check=0.05,
|
||||
logger=logger
|
||||
)
|
||||
trainer = Trainer(**trainer_options)
|
||||
result = trainer.fit(model)
|
||||
|
||||
assert result == 1, 'Training failed'
|
||||
|
||||
|
||||
@patch('pytorch_lightning.loggers.neptune.neptune')
|
||||
def test_neptune_online(neptune):
|
||||
logger = NeptuneLogger(api_key='test', project_name='project')
|
||||
neptune.init.assert_called_once_with(api_token='test', project_qualified_name='project')
|
||||
|
||||
assert logger.name == neptune.create_experiment().name
|
||||
assert logger.version == neptune.create_experiment().id
|
||||
|
||||
|
||||
@patch('pytorch_lightning.loggers.neptune.neptune')
|
||||
def test_neptune_additional_methods(neptune):
|
||||
logger = NeptuneLogger(offline_mode=True)
|
||||
|
||||
logger.log_metric('test', torch.ones(1))
|
||||
neptune.create_experiment().log_metric.assert_called_once_with('test', torch.ones(1))
|
||||
neptune.create_experiment().log_metric.reset_mock()
|
||||
|
||||
logger.log_metric('test', 1.0)
|
||||
neptune.create_experiment().log_metric.assert_called_once_with('test', 1.0)
|
||||
neptune.create_experiment().log_metric.reset_mock()
|
||||
|
||||
logger.log_metric('test', 1.0, step=2)
|
||||
neptune.create_experiment().log_metric.assert_called_once_with('test', x=2, y=1.0)
|
||||
neptune.create_experiment().log_metric.reset_mock()
|
||||
|
||||
logger.log_text('test', 'text')
|
||||
neptune.create_experiment().log_metric.assert_called_once_with('test', 'text')
|
||||
neptune.create_experiment().log_metric.reset_mock()
|
||||
|
||||
logger.log_image('test', 'image file')
|
||||
neptune.create_experiment().log_image.assert_called_once_with('test', 'image file')
|
||||
neptune.create_experiment().log_image.reset_mock()
|
||||
|
||||
logger.log_image('test', 'image file', step=2)
|
||||
neptune.create_experiment().log_image.assert_called_once_with('test', x=2, y='image file')
|
||||
neptune.create_experiment().log_image.reset_mock()
|
||||
|
||||
logger.log_artifact('file')
|
||||
neptune.create_experiment().log_artifact.assert_called_once_with('file', None)
|
||||
|
||||
logger.set_property('property', 10)
|
||||
neptune.create_experiment().set_property.assert_called_once_with('property', 10)
|
||||
|
||||
logger.append_tags('one tag')
|
||||
neptune.create_experiment().append_tags.assert_called_once_with('one tag')
|
||||
neptune.create_experiment().append_tags.reset_mock()
|
||||
|
||||
logger.append_tags(['two', 'tags'])
|
||||
neptune.create_experiment().append_tags.assert_called_once_with('two', 'tags')
|
||||
|
||||
|
||||
def test_neptune_pickle(tmpdir):
|
||||
"""Verify that pickling trainer with neptune logger works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
logger = NeptuneLogger(offline_mode=True)
|
||||
|
||||
trainer_options = dict(
|
||||
default_save_path=tmpdir,
|
||||
max_epochs=1,
|
||||
logger=logger
|
||||
)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
pkl_bytes = pickle.dumps(trainer)
|
||||
trainer2 = pickle.loads(pkl_bytes)
|
||||
trainer2.logger.log_metrics({'acc': 1.0})
|
||||
@@ -0,0 +1,113 @@
|
||||
import pickle
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
import tests.models.utils as tutils
|
||||
from pytorch_lightning import Trainer
|
||||
from pytorch_lightning.loggers import (
|
||||
TensorBoardLogger
|
||||
)
|
||||
from tests.models import LightningTestModel
|
||||
|
||||
|
||||
def test_tensorboard_logger(tmpdir):
|
||||
"""Verify that basic functionality of Tensorboard logger works."""
|
||||
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
logger = TensorBoardLogger(save_dir=tmpdir, name="tensorboard_logger_test")
|
||||
|
||||
trainer_options = dict(max_epochs=1, train_percent_check=0.01, logger=logger)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
result = trainer.fit(model)
|
||||
|
||||
print("result finished")
|
||||
assert result == 1, "Training failed"
|
||||
|
||||
|
||||
def test_tensorboard_pickle(tmpdir):
|
||||
"""Verify that pickling trainer with Tensorboard logger works."""
|
||||
|
||||
logger = TensorBoardLogger(save_dir=tmpdir, name="tensorboard_pickle_test")
|
||||
|
||||
trainer_options = dict(max_epochs=1, logger=logger)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
pkl_bytes = pickle.dumps(trainer)
|
||||
trainer2 = pickle.loads(pkl_bytes)
|
||||
trainer2.logger.log_metrics({"acc": 1.0})
|
||||
|
||||
|
||||
def test_tensorboard_automatic_versioning(tmpdir):
|
||||
"""Verify that automatic versioning works"""
|
||||
|
||||
root_dir = tmpdir.mkdir("tb_versioning")
|
||||
root_dir.mkdir("version_0")
|
||||
root_dir.mkdir("version_1")
|
||||
|
||||
logger = TensorBoardLogger(save_dir=tmpdir, name="tb_versioning")
|
||||
|
||||
assert logger.version == 2
|
||||
|
||||
|
||||
def test_tensorboard_manual_versioning(tmpdir):
|
||||
"""Verify that manual versioning works"""
|
||||
|
||||
root_dir = tmpdir.mkdir("tb_versioning")
|
||||
root_dir.mkdir("version_0")
|
||||
root_dir.mkdir("version_1")
|
||||
root_dir.mkdir("version_2")
|
||||
|
||||
logger = TensorBoardLogger(save_dir=tmpdir, name="tb_versioning", version=1)
|
||||
|
||||
assert logger.version == 1
|
||||
|
||||
|
||||
def test_tensorboard_named_version(tmpdir):
|
||||
"""Verify that manual versioning works for string versions, e.g. '2020-02-05-162402' """
|
||||
|
||||
tmpdir.mkdir("tb_versioning")
|
||||
expected_version = "2020-02-05-162402"
|
||||
|
||||
logger = TensorBoardLogger(save_dir=tmpdir, name="tb_versioning", version=expected_version)
|
||||
logger.log_hyperparams({"a": 1, "b": 2}) # Force data to be written
|
||||
|
||||
assert logger.version == expected_version
|
||||
# Could also test existence of the directory but this fails
|
||||
# in the "minimum requirements" test setup
|
||||
|
||||
|
||||
def test_tensorboard_no_name(tmpdir):
|
||||
"""Verify that None or empty name works"""
|
||||
|
||||
logger = TensorBoardLogger(save_dir=tmpdir, name="")
|
||||
assert logger.root_dir == tmpdir
|
||||
|
||||
logger = TensorBoardLogger(save_dir=tmpdir, name=None)
|
||||
assert logger.root_dir == tmpdir
|
||||
|
||||
|
||||
@pytest.mark.parametrize("step_idx", [10, None])
|
||||
def test_tensorboard_log_metrics(tmpdir, step_idx):
|
||||
logger = TensorBoardLogger(tmpdir)
|
||||
metrics = {
|
||||
"float": 0.3,
|
||||
"int": 1,
|
||||
"FloatTensor": torch.tensor(0.1),
|
||||
"IntTensor": torch.tensor(1)
|
||||
}
|
||||
logger.log_metrics(metrics, step_idx)
|
||||
|
||||
|
||||
def test_tensorboard_log_hyperparams(tmpdir):
|
||||
logger = TensorBoardLogger(tmpdir)
|
||||
hparams = {
|
||||
"float": 0.3,
|
||||
"int": 1,
|
||||
"string": "abc",
|
||||
"bool": True
|
||||
}
|
||||
logger.log_hyperparams(hparams)
|
||||
@@ -0,0 +1,51 @@
|
||||
import pickle
|
||||
|
||||
import tests.models.utils as tutils
|
||||
from pytorch_lightning import Trainer
|
||||
from tests.models import LightningTestModel
|
||||
|
||||
|
||||
def test_testtube_logger(tmpdir):
|
||||
"""Verify that basic functionality of test tube logger works."""
|
||||
tutils.reset_seed()
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
logger = tutils.get_test_tube_logger(tmpdir, False)
|
||||
|
||||
assert logger.name == 'lightning_logs'
|
||||
|
||||
trainer_options = dict(
|
||||
default_save_path=tmpdir,
|
||||
max_epochs=1,
|
||||
train_percent_check=0.05,
|
||||
logger=logger
|
||||
)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
result = trainer.fit(model)
|
||||
|
||||
assert result == 1, 'Training failed'
|
||||
|
||||
|
||||
def test_testtube_pickle(tmpdir):
|
||||
"""Verify that pickling a trainer containing a test tube logger works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
hparams = tutils.get_hparams()
|
||||
|
||||
logger = tutils.get_test_tube_logger(tmpdir, False)
|
||||
logger.log_hyperparams(hparams)
|
||||
logger.save()
|
||||
|
||||
trainer_options = dict(
|
||||
default_save_path=tmpdir,
|
||||
max_epochs=1,
|
||||
train_percent_check=0.05,
|
||||
logger=logger
|
||||
)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
pkl_bytes = pickle.dumps(trainer)
|
||||
trainer2 = pickle.loads(pkl_bytes)
|
||||
trainer2.logger.log_metrics({'acc': 1.0})
|
||||
@@ -0,0 +1,78 @@
|
||||
import os
|
||||
import pickle
|
||||
|
||||
import pytest
|
||||
from unittest.mock import patch
|
||||
|
||||
import tests.models.utils as tutils
|
||||
from pytorch_lightning import Trainer
|
||||
from pytorch_lightning.loggers import WandbLogger
|
||||
|
||||
|
||||
@patch('pytorch_lightning.loggers.wandb.wandb')
|
||||
def test_wandb_logger(wandb):
|
||||
"""Verify that basic functionality of wandb logger works.
|
||||
Wandb doesn't work well with pytest so we have to mock it out here."""
|
||||
tutils.reset_seed()
|
||||
|
||||
logger = WandbLogger(anonymous=True, offline=True)
|
||||
|
||||
logger.log_metrics({'acc': 1.0})
|
||||
wandb.init().log.assert_called_once_with({'acc': 1.0})
|
||||
|
||||
wandb.init().log.reset_mock()
|
||||
logger.log_metrics({'acc': 1.0}, step=3)
|
||||
wandb.init().log.assert_called_once_with({'global_step': 3, 'acc': 1.0})
|
||||
|
||||
logger.log_hyperparams('test')
|
||||
wandb.init().config.update.assert_called_once_with('test')
|
||||
|
||||
logger.watch('model', 'log', 10)
|
||||
wandb.watch.assert_called_once_with('model', log='log', log_freq=10)
|
||||
|
||||
logger.finalize('fail')
|
||||
wandb.join.assert_called_once_with(1)
|
||||
|
||||
wandb.join.reset_mock()
|
||||
logger.finalize('success')
|
||||
wandb.join.assert_called_once_with(0)
|
||||
|
||||
wandb.join.reset_mock()
|
||||
wandb.join.side_effect = TypeError
|
||||
with pytest.raises(TypeError):
|
||||
logger.finalize('any')
|
||||
|
||||
wandb.join.assert_called()
|
||||
|
||||
assert logger.name == wandb.init().project_name()
|
||||
assert logger.version == wandb.init().id
|
||||
|
||||
|
||||
@patch('pytorch_lightning.loggers.wandb.wandb')
|
||||
def test_wandb_pickle(wandb):
|
||||
"""Verify that pickling trainer with wandb logger works.
|
||||
Wandb doesn't work well with pytest so we have to mock it out here."""
|
||||
tutils.reset_seed()
|
||||
|
||||
class Experiment:
|
||||
id = 'the_id'
|
||||
|
||||
wandb.init.return_value = Experiment()
|
||||
|
||||
logger = WandbLogger(id='the_id', offline=True)
|
||||
|
||||
trainer_options = dict(max_epochs=1, logger=logger)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
pkl_bytes = pickle.dumps(trainer)
|
||||
trainer2 = pickle.loads(pkl_bytes)
|
||||
|
||||
assert os.environ['WANDB_MODE'] == 'dryrun'
|
||||
assert trainer2.logger.__class__.__name__ == WandbLogger.__name__
|
||||
_ = trainer2.logger.experiment
|
||||
|
||||
wandb.init.assert_called()
|
||||
assert 'id' in wandb.init.call_args[1]
|
||||
assert wandb.init.call_args[1]['id'] == 'the_id'
|
||||
|
||||
del os.environ['WANDB_MODE']
|
||||
@@ -1,457 +0,0 @@
|
||||
import os
|
||||
import pickle
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
import tests.models.utils as tutils
|
||||
from pytorch_lightning import Trainer
|
||||
from pytorch_lightning.loggers import (
|
||||
LightningLoggerBase,
|
||||
rank_zero_only,
|
||||
TensorBoardLogger,
|
||||
MLFlowLogger,
|
||||
CometLogger,
|
||||
WandbLogger,
|
||||
NeptuneLogger
|
||||
)
|
||||
from tests.models import LightningTestModel
|
||||
|
||||
|
||||
def test_testtube_logger(tmpdir):
|
||||
"""Verify that basic functionality of test tube logger works."""
|
||||
tutils.reset_seed()
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
logger = tutils.get_test_tube_logger(tmpdir, False)
|
||||
|
||||
trainer_options = dict(
|
||||
default_save_path=tmpdir,
|
||||
max_epochs=1,
|
||||
train_percent_check=0.05,
|
||||
logger=logger
|
||||
)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
result = trainer.fit(model)
|
||||
|
||||
assert result == 1, "Training failed"
|
||||
|
||||
|
||||
def test_testtube_pickle(tmpdir):
|
||||
"""Verify that pickling a trainer containing a test tube logger works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
logger = tutils.get_test_tube_logger(tmpdir, False)
|
||||
logger.log_hyperparams(hparams)
|
||||
logger.save()
|
||||
|
||||
trainer_options = dict(
|
||||
default_save_path=tmpdir,
|
||||
max_epochs=1,
|
||||
train_percent_check=0.05,
|
||||
logger=logger
|
||||
)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
pkl_bytes = pickle.dumps(trainer)
|
||||
trainer2 = pickle.loads(pkl_bytes)
|
||||
trainer2.logger.log_metrics({"acc": 1.0})
|
||||
|
||||
|
||||
def test_mlflow_logger(tmpdir):
|
||||
"""Verify that basic functionality of mlflow logger works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
mlflow_dir = os.path.join(tmpdir, "mlruns")
|
||||
logger = MLFlowLogger("test", tracking_uri=f"file:{os.sep * 2}{mlflow_dir}")
|
||||
|
||||
trainer_options = dict(
|
||||
default_save_path=tmpdir,
|
||||
max_epochs=1,
|
||||
train_percent_check=0.05,
|
||||
logger=logger
|
||||
)
|
||||
trainer = Trainer(**trainer_options)
|
||||
result = trainer.fit(model)
|
||||
|
||||
print('result finished')
|
||||
assert result == 1, "Training failed"
|
||||
|
||||
|
||||
def test_mlflow_pickle(tmpdir):
|
||||
"""Verify that pickling trainer with mlflow logger works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
# hparams = tutils.get_hparams()
|
||||
# model = LightningTestModel(hparams)
|
||||
|
||||
mlflow_dir = os.path.join(tmpdir, "mlruns")
|
||||
logger = MLFlowLogger("test", tracking_uri=f"file:{os.sep * 2}{mlflow_dir}")
|
||||
trainer_options = dict(
|
||||
default_save_path=tmpdir,
|
||||
max_epochs=1,
|
||||
logger=logger
|
||||
)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
pkl_bytes = pickle.dumps(trainer)
|
||||
trainer2 = pickle.loads(pkl_bytes)
|
||||
trainer2.logger.log_metrics({"acc": 1.0})
|
||||
|
||||
|
||||
def test_comet_logger(tmpdir, monkeypatch):
|
||||
"""Verify that basic functionality of Comet.ml logger works."""
|
||||
|
||||
# prevent comet logger from trying to print at exit, since
|
||||
# pytest's stdout/stderr redirection breaks it
|
||||
import atexit
|
||||
monkeypatch.setattr(atexit, "register", lambda _: None)
|
||||
|
||||
tutils.reset_seed()
|
||||
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
comet_dir = os.path.join(tmpdir, "cometruns")
|
||||
|
||||
# We test CometLogger in offline mode with local saves
|
||||
logger = CometLogger(
|
||||
save_dir=comet_dir,
|
||||
project_name="general",
|
||||
workspace="dummy-test",
|
||||
)
|
||||
|
||||
trainer_options = dict(
|
||||
default_save_path=tmpdir,
|
||||
max_epochs=1,
|
||||
train_percent_check=0.05,
|
||||
logger=logger
|
||||
)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
result = trainer.fit(model)
|
||||
|
||||
print('result finished')
|
||||
assert result == 1, "Training failed"
|
||||
|
||||
|
||||
def test_comet_pickle(tmpdir, monkeypatch):
|
||||
"""Verify that pickling trainer with comet logger works."""
|
||||
|
||||
# prevent comet logger from trying to print at exit, since
|
||||
# pytest's stdout/stderr redirection breaks it
|
||||
import atexit
|
||||
monkeypatch.setattr(atexit, "register", lambda _: None)
|
||||
|
||||
tutils.reset_seed()
|
||||
|
||||
# hparams = tutils.get_hparams()
|
||||
# model = LightningTestModel(hparams)
|
||||
|
||||
comet_dir = os.path.join(tmpdir, "cometruns")
|
||||
|
||||
# We test CometLogger in offline mode with local saves
|
||||
logger = CometLogger(
|
||||
save_dir=comet_dir,
|
||||
project_name="general",
|
||||
workspace="dummy-test",
|
||||
)
|
||||
|
||||
trainer_options = dict(
|
||||
default_save_path=tmpdir,
|
||||
max_epochs=1,
|
||||
logger=logger
|
||||
)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
pkl_bytes = pickle.dumps(trainer)
|
||||
trainer2 = pickle.loads(pkl_bytes)
|
||||
trainer2.logger.log_metrics({"acc": 1.0})
|
||||
|
||||
|
||||
def test_wandb_logger(tmpdir):
|
||||
"""Verify that basic functionality of wandb logger works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
wandb_dir = os.path.join(tmpdir, "wandb")
|
||||
_ = WandbLogger(save_dir=wandb_dir, anonymous=True, offline=True)
|
||||
|
||||
|
||||
def test_wandb_pickle(tmpdir):
|
||||
"""Verify that pickling trainer with wandb logger works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
wandb_dir = str(tmpdir)
|
||||
logger = WandbLogger(save_dir=wandb_dir, anonymous=True, offline=True)
|
||||
assert logger is not None
|
||||
|
||||
|
||||
def test_neptune_logger(tmpdir):
|
||||
"""Verify that basic functionality of neptune logger works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
logger = NeptuneLogger(offline_mode=True)
|
||||
|
||||
trainer_options = dict(
|
||||
default_save_path=tmpdir,
|
||||
max_epochs=1,
|
||||
train_percent_check=0.05,
|
||||
logger=logger
|
||||
)
|
||||
trainer = Trainer(**trainer_options)
|
||||
result = trainer.fit(model)
|
||||
|
||||
print('result finished')
|
||||
assert result == 1, "Training failed"
|
||||
|
||||
|
||||
def test_neptune_pickle(tmpdir):
|
||||
"""Verify that pickling trainer with neptune logger works."""
|
||||
tutils.reset_seed()
|
||||
|
||||
# hparams = tutils.get_hparams()
|
||||
# model = LightningTestModel(hparams)
|
||||
|
||||
logger = NeptuneLogger(offline_mode=True)
|
||||
|
||||
trainer_options = dict(
|
||||
default_save_path=tmpdir,
|
||||
max_epochs=1,
|
||||
logger=logger
|
||||
)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
pkl_bytes = pickle.dumps(trainer)
|
||||
trainer2 = pickle.loads(pkl_bytes)
|
||||
trainer2.logger.log_metrics({"acc": 1.0})
|
||||
|
||||
|
||||
def test_tensorboard_logger(tmpdir):
|
||||
"""Verify that basic functionality of Tensorboard logger works."""
|
||||
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
logger = TensorBoardLogger(save_dir=tmpdir, name="tensorboard_logger_test")
|
||||
|
||||
trainer_options = dict(max_epochs=1, train_percent_check=0.01, logger=logger)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
result = trainer.fit(model)
|
||||
|
||||
print("result finished")
|
||||
assert result == 1, "Training failed"
|
||||
|
||||
|
||||
def test_tensorboard_pickle(tmpdir):
|
||||
"""Verify that pickling trainer with Tensorboard logger works."""
|
||||
|
||||
# hparams = tutils.get_hparams()
|
||||
# model = LightningTestModel(hparams)
|
||||
|
||||
logger = TensorBoardLogger(save_dir=tmpdir, name="tensorboard_pickle_test")
|
||||
|
||||
trainer_options = dict(max_epochs=1, logger=logger)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
pkl_bytes = pickle.dumps(trainer)
|
||||
trainer2 = pickle.loads(pkl_bytes)
|
||||
trainer2.logger.log_metrics({"acc": 1.0})
|
||||
|
||||
|
||||
def test_tensorboard_automatic_versioning(tmpdir):
|
||||
"""Verify that automatic versioning works"""
|
||||
|
||||
root_dir = tmpdir.mkdir("tb_versioning")
|
||||
root_dir.mkdir("version_0")
|
||||
root_dir.mkdir("version_1")
|
||||
|
||||
logger = TensorBoardLogger(save_dir=tmpdir, name="tb_versioning")
|
||||
|
||||
assert logger.version == 2
|
||||
|
||||
|
||||
def test_tensorboard_manual_versioning(tmpdir):
|
||||
"""Verify that manual versioning works"""
|
||||
|
||||
root_dir = tmpdir.mkdir("tb_versioning")
|
||||
root_dir.mkdir("version_0")
|
||||
root_dir.mkdir("version_1")
|
||||
root_dir.mkdir("version_2")
|
||||
|
||||
logger = TensorBoardLogger(save_dir=tmpdir, name="tb_versioning", version=1)
|
||||
|
||||
assert logger.version == 1
|
||||
|
||||
|
||||
def test_tensorboard_named_version(tmpdir):
|
||||
"""Verify that manual versioning works for string versions, e.g. '2020-02-05-162402' """
|
||||
|
||||
tmpdir.mkdir("tb_versioning")
|
||||
expected_version = "2020-02-05-162402"
|
||||
|
||||
logger = TensorBoardLogger(save_dir=tmpdir, name="tb_versioning", version=expected_version)
|
||||
logger.log_hyperparams({"a": 1, "b": 2}) # Force data to be written
|
||||
|
||||
assert logger.version == expected_version
|
||||
# Could also test existence of the directory but this fails in the "minimum requirements" test setup
|
||||
|
||||
|
||||
@pytest.mark.parametrize("step_idx", [10, None])
|
||||
def test_tensorboard_log_metrics(tmpdir, step_idx):
|
||||
logger = TensorBoardLogger(tmpdir)
|
||||
metrics = {
|
||||
"float": 0.3,
|
||||
"int": 1,
|
||||
"FloatTensor": torch.tensor(0.1),
|
||||
"IntTensor": torch.tensor(1)
|
||||
}
|
||||
logger.log_metrics(metrics, step_idx)
|
||||
|
||||
|
||||
def test_tensorboard_log_hyperparams(tmpdir):
|
||||
logger = TensorBoardLogger(tmpdir)
|
||||
hparams = {
|
||||
"float": 0.3,
|
||||
"int": 1,
|
||||
"string": "abc",
|
||||
"bool": True
|
||||
}
|
||||
logger.log_hyperparams(hparams)
|
||||
|
||||
|
||||
class CustomLogger(LightningLoggerBase):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.hparams_logged = None
|
||||
self.metrics_logged = None
|
||||
self.finalized = False
|
||||
|
||||
@property
|
||||
def experiment(self):
|
||||
return 'test'
|
||||
|
||||
@rank_zero_only
|
||||
def log_hyperparams(self, params):
|
||||
self.hparams_logged = params
|
||||
|
||||
@rank_zero_only
|
||||
def log_metrics(self, metrics, step):
|
||||
self.metrics_logged = metrics
|
||||
|
||||
@rank_zero_only
|
||||
def finalize(self, status):
|
||||
self.finalized_status = status
|
||||
|
||||
@property
|
||||
def name(self):
|
||||
return "name"
|
||||
|
||||
@property
|
||||
def version(self):
|
||||
return "1"
|
||||
|
||||
|
||||
def test_custom_logger(tmpdir):
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
logger = CustomLogger()
|
||||
|
||||
trainer_options = dict(
|
||||
max_epochs=1,
|
||||
train_percent_check=0.05,
|
||||
logger=logger,
|
||||
default_save_path=tmpdir
|
||||
)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
result = trainer.fit(model)
|
||||
assert result == 1, "Training failed"
|
||||
assert logger.hparams_logged == hparams
|
||||
assert logger.metrics_logged != {}
|
||||
assert logger.finalized_status == "success"
|
||||
|
||||
|
||||
def test_multiple_loggers(tmpdir):
|
||||
hparams = tutils.get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
logger1 = CustomLogger()
|
||||
logger2 = CustomLogger()
|
||||
|
||||
trainer_options = dict(
|
||||
max_epochs=1,
|
||||
train_percent_check=0.05,
|
||||
logger=[logger1, logger2],
|
||||
default_save_path=tmpdir
|
||||
)
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
result = trainer.fit(model)
|
||||
assert result == 1, "Training failed"
|
||||
|
||||
assert logger1.hparams_logged == hparams
|
||||
assert logger1.metrics_logged != {}
|
||||
assert logger1.finalized_status == "success"
|
||||
|
||||
assert logger2.hparams_logged == hparams
|
||||
assert logger2.metrics_logged != {}
|
||||
assert logger2.finalized_status == "success"
|
||||
|
||||
|
||||
def test_multiple_loggers_pickle(tmpdir):
|
||||
"""Verify that pickling trainer with multiple loggers works."""
|
||||
|
||||
logger1 = CustomLogger()
|
||||
logger2 = CustomLogger()
|
||||
|
||||
trainer_options = dict(max_epochs=1, logger=[logger1, logger2])
|
||||
|
||||
trainer = Trainer(**trainer_options)
|
||||
pkl_bytes = pickle.dumps(trainer)
|
||||
trainer2 = pickle.loads(pkl_bytes)
|
||||
trainer2.logger.log_metrics({"acc": 1.0}, 0)
|
||||
|
||||
assert logger1.metrics_logged != {}
|
||||
assert logger2.metrics_logged != {}
|
||||
|
||||
|
||||
def test_adding_step_key(tmpdir):
|
||||
logged_step = 0
|
||||
|
||||
def _validation_end(outputs):
|
||||
nonlocal logged_step
|
||||
logged_step += 1
|
||||
return {"log": {"step": logged_step, "val_acc": logged_step / 10}}
|
||||
|
||||
def _log_metrics_decorator(log_metrics_fn):
|
||||
def decorated(metrics, step):
|
||||
if "val_acc" in metrics:
|
||||
assert step == logged_step
|
||||
return log_metrics_fn(metrics, step)
|
||||
|
||||
return decorated
|
||||
|
||||
model, hparams = tutils.get_model()
|
||||
model.validation_end = _validation_end
|
||||
trainer_options = dict(
|
||||
max_epochs=4,
|
||||
default_save_path=tmpdir,
|
||||
train_percent_check=0.001,
|
||||
val_percent_check=0.01,
|
||||
num_sanity_val_steps=0
|
||||
)
|
||||
trainer = Trainer(**trainer_options)
|
||||
trainer.logger.log_metrics = _log_metrics_decorator(trainer.logger.log_metrics)
|
||||
trainer.fit(model)
|
||||
Reference in New Issue
Block a user