Refactor: name modules (#548)

* refactor: rename some modules

* add deprecation warnings

* fix paths
This commit is contained in:
Jirka Borovec
2019-11-26 22:39:18 -05:00
committed by William Falcon
parent fea7cc87f6
commit 9785a3e78e
33 changed files with 379 additions and 325 deletions
@@ -16,7 +16,7 @@ from torch.utils.data.distributed import DistributedSampler
from torchvision.datasets import MNIST
import pytorch_lightning as pl
from pytorch_lightning.root_module.root_module import LightningModule
from pytorch_lightning.core.lightning import LightningModule
class LightningTemplateModel(LightningModule):
+2 -2
View File
@@ -25,8 +25,8 @@ if __LIGHTNING_SETUP__:
# process, as it may not be compiled yet
else:
from .trainer.trainer import Trainer
from .root_module.root_module import LightningModule
from .root_module.decorators import data_loader
from .core.lightning import LightningModule
from .core.decorators import data_loader
__all__ = [
'Trainer',
+1 -1
View File
@@ -4,7 +4,7 @@ import logging
import warnings
import numpy as np
from pytorch_lightning.pt_overrides.override_data_parallel import LightningDistributedDataParallel
from pytorch_lightning.overrides.data_parallel import LightningDistributedDataParallel
class Callback(object):
View File
@@ -6,14 +6,14 @@ from argparse import Namespace
import torch
import torch.distributed as dist
from pytorch_lightning.root_module.decorators import data_loader
from pytorch_lightning.root_module.grads import GradInformation
from pytorch_lightning.root_module.hooks import ModelHooks
from pytorch_lightning.root_module.memory import ModelSummary
from pytorch_lightning.root_module.model_saving import ModelIO
from pytorch_lightning.core.decorators import data_loader
from pytorch_lightning.core.grads import GradInformation
from pytorch_lightning.core.hooks import ModelHooks
from pytorch_lightning.core.memory import ModelSummary
from pytorch_lightning.core.saving import ModelIO
from pytorch_lightning.trainer.trainer_io import load_hparams_from_tags_csv
import logging
from pytorch_lightning.pt_overrides.override_data_parallel import LightningDistributedDataParallel
from pytorch_lightning.overrides.data_parallel import LightningDistributedDataParallel
class LightningModule(GradInformation, ModelIO, ModelHooks):
+6
View File
@@ -0,0 +1,6 @@
import warnings
warnings.warn("`model_saving` module has been renamed to `saving` since v0.5.3"
" and will be removed in v0.8.0", DeprecationWarning)
from pytorch_lightning.core.saving import ModelIO # noqa: E402
+6
View File
@@ -0,0 +1,6 @@
import warnings
warnings.warn("`root_module` module has been renamed to `lightning` since v0.5.3"
" and will be removed in v0.8.0", DeprecationWarning)
from pytorch_lightning.core.lightning import LightningModule # noqa: E402
+3 -3
View File
@@ -2,17 +2,17 @@ from os import environ
from .base import LightningLoggerBase, rank_zero_only
try:
from .test_tube_logger import TestTubeLogger
from .test_tube import TestTubeLogger
except ImportError:
pass
try:
from .mlflow_logger import MLFlowLogger
from .mlflow import MLFlowLogger
except ImportError:
pass
try:
# needed to prevent ImportError and duplicated logs.
environ["COMET_DISABLE_AUTO_LOGGING"] = "1"
from .comet_logger import CometLogger
from .comet import CometLogger
except ImportError:
del environ["COMET_DISABLE_AUTO_LOGGING"]
+126
View File
@@ -0,0 +1,126 @@
from logging import getLogger
try:
from comet_ml import Experiment as CometExperiment
from comet_ml import OfflineExperiment as CometOfflineExperiment
from comet_ml.papi import API
except ImportError:
raise ImportError('Missing comet_ml package.')
from torch import is_tensor
from .base import LightningLoggerBase, rank_zero_only
from ..utilities.debugging import MisconfigurationException
logger = getLogger(__name__)
class CometLogger(LightningLoggerBase):
def __init__(self, api_key=None, save_dir=None, workspace=None,
rest_api_key=None, project_name=None, experiment_name=None, **kwargs):
"""
Initialize a Comet.ml logger. Requires either an API Key (online mode) or a local directory path (offline mode)
:param str api_key: Required in online mode. API key, found on Comet.ml
:param str save_dir: Required in offline mode. The path for the directory to save local comet logs
:param str workspace: Optional. Name of workspace for this user
:param str project_name: Optional. Send your experiment to a specific project.
Otherwise will be sent to Uncategorized Experiments.
If project name does not already exists Comet.ml will create a new project.
:param str rest_api_key: Optional. Rest API key found in Comet.ml settings.
This is used to determine version number
:param str experiment_name: Optional. String representing the name for this particular experiment on Comet.ml
"""
super().__init__()
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:
self.mode = "online"
self.api_key = api_key
elif save_dir is not None:
self.mode = "offline"
self.save_dir = save_dir
else:
# If neither api_key nor save_dir are passed as arguments, raise an exception
raise MisconfigurationException("CometLogger requires either api_key or save_dir during initialization.")
logger.info(f"CometLogger will be initialized in {self.mode} mode")
self.workspace = workspace
self.project_name = project_name
self._kwargs = kwargs
if rest_api_key is not None:
# Comet.ml rest API, used to determine version number
self.rest_api_key = rest_api_key
self.comet_api = API(self.rest_api_key)
else:
self.rest_api_key = None
self.comet_api = None
if experiment_name:
try:
self.name = experiment_name
except TypeError as e:
logger.exception("Failed to set experiment name for comet.ml logger")
@property
def experiment(self):
if self._experiment is not None:
return self._experiment
if self.mode == "online":
self._experiment = CometExperiment(
api_key=self.api_key,
workspace=self.workspace,
project_name=self.project_name,
**self._kwargs
)
else:
self._experiment = CometOfflineExperiment(
offline_directory=self.save_dir,
workspace=self.workspace,
project_name=self.project_name,
**self._kwargs
)
return self._experiment
@rank_zero_only
def log_hyperparams(self, params):
self.experiment.log_parameters(vars(params))
@rank_zero_only
def log_metrics(self, metrics, step_num=None):
# Comet.ml expects metrics to be a dictionary of detached tensors on CPU
for key, val in metrics.items():
if is_tensor(val):
metrics[key] = val.cpu().detach()
self.experiment.log_metrics(metrics, step=step_num)
@rank_zero_only
def finalize(self, status):
self.experiment.end()
@property
def name(self):
return self.experiment.project_name
@name.setter
def name(self, value):
self.experiment.set_name(value)
@property
def version(self):
if self.project_name and self.rest_api_key:
# Determines the number of experiments in this project, and returns the next integer as the version number
nb_exps = len(self.comet_api.get_experiments(self.workspace, self.project_name))
return nb_exps + 1
else:
return None
+4 -124
View File
@@ -1,126 +1,6 @@
from logging import getLogger
import warnings
try:
from comet_ml import Experiment as CometExperiment
from comet_ml import OfflineExperiment as CometOfflineExperiment
from comet_ml.papi import API
except ImportError:
raise ImportError('Missing comet_ml package.')
warnings.warn("`comet_logger` module has been renamed to `comet` since v0.5.3"
" and will be removed in v0.8.0", DeprecationWarning)
from torch import is_tensor
from .base import LightningLoggerBase, rank_zero_only
from ..utilities.debugging import MisconfigurationException
logger = getLogger(__name__)
class CometLogger(LightningLoggerBase):
def __init__(self, api_key=None, save_dir=None, workspace=None,
rest_api_key=None, project_name=None, experiment_name=None, **kwargs):
"""
Initialize a Comet.ml logger. Requires either an API Key (online mode) or a local directory path (offline mode)
:param str api_key: Required in online mode. API key, found on Comet.ml
:param str save_dir: Required in offline mode. The path for the directory to save local comet logs
:param str workspace: Optional. Name of workspace for this user
:param str project_name: Optional. Send your experiment to a specific project.
Otherwise will be sent to Uncategorized Experiments.
If project name does not already exists Comet.ml will create a new project.
:param str rest_api_key: Optional. Rest API key found in Comet.ml settings.
This is used to determine version number
:param str experiment_name: Optional. String representing the name for this particular experiment on Comet.ml
"""
super().__init__()
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:
self.mode = "online"
self.api_key = api_key
elif save_dir is not None:
self.mode = "offline"
self.save_dir = save_dir
else:
# If neither api_key nor save_dir are passed as arguments, raise an exception
raise MisconfigurationException("CometLogger requires either api_key or save_dir during initialization.")
logger.info(f"CometLogger will be initialized in {self.mode} mode")
self.workspace = workspace
self.project_name = project_name
self._kwargs = kwargs
if rest_api_key is not None:
# Comet.ml rest API, used to determine version number
self.rest_api_key = rest_api_key
self.comet_api = API(self.rest_api_key)
else:
self.rest_api_key = None
self.comet_api = None
if experiment_name:
try:
self.name = experiment_name
except TypeError as e:
logger.exception("Failed to set experiment name for comet.ml logger")
@property
def experiment(self):
if self._experiment is not None:
return self._experiment
if self.mode == "online":
self._experiment = CometExperiment(
api_key=self.api_key,
workspace=self.workspace,
project_name=self.project_name,
**self._kwargs
)
else:
self._experiment = CometOfflineExperiment(
offline_directory=self.save_dir,
workspace=self.workspace,
project_name=self.project_name,
**self._kwargs
)
return self._experiment
@rank_zero_only
def log_hyperparams(self, params):
self.experiment.log_parameters(vars(params))
@rank_zero_only
def log_metrics(self, metrics, step_num=None):
# Comet.ml expects metrics to be a dictionary of detached tensors on CPU
for key, val in metrics.items():
if is_tensor(val):
metrics[key] = val.cpu().detach()
self.experiment.log_metrics(metrics, step=step_num)
@rank_zero_only
def finalize(self, status):
self.experiment.end()
@property
def name(self):
return self.experiment.project_name
@name.setter
def name(self, value):
self.experiment.set_name(value)
@property
def version(self):
if self.project_name and self.rest_api_key:
# Determines the number of experiments in this project, and returns the next integer as the version number
nb_exps = len(self.comet_api.get_experiments(self.workspace, self.project_name))
return nb_exps + 1
else:
return None
from pytorch_lightning.logging.comet import CometLogger # noqa: E402
+70
View File
@@ -0,0 +1,70 @@
from logging import getLogger
from time import time
try:
import mlflow
except ImportError:
raise ImportError('Missing mlflow package.')
from .base import LightningLoggerBase, rank_zero_only
logger = getLogger(__name__)
class MLFlowLogger(LightningLoggerBase):
def __init__(self, experiment_name, tracking_uri=None, tags=None):
super().__init__()
self.experiment = mlflow.tracking.MlflowClient(tracking_uri)
self.experiment_name = experiment_name
self._run_id = None
self.tags = tags
@property
def run_id(self):
if self._run_id is not None:
return self._run_id
experiment = self.experiment.get_experiment_by_name(self.experiment_name)
if experiment is None:
logger.warning(
f"Experiment with name f{self.experiment_name} not found. Creating it."
)
self.experiment.create_experiment(self.experiment_name)
experiment = self.experiment.get_experiment_by_name(self.experiment_name)
run = self.experiment.create_run(experiment.experiment_id, tags=self.tags)
self._run_id = run.info.run_id
return self._run_id
@rank_zero_only
def log_hyperparams(self, params):
for k, v in vars(params).items():
self.experiment.log_param(self.run_id, k, v)
@rank_zero_only
def log_metrics(self, metrics, step_num=None):
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}"
)
continue
self.experiment.log_metric(self.run_id, k, v, timestamp_ms, step_num)
def save(self):
pass
@rank_zero_only
def finalize(self, status="FINISHED"):
if status == 'success':
status = 'FINISHED'
self.experiment.set_terminated(self.run_id, status)
@property
def name(self):
return self.experiment_name
@property
def version(self):
return self._run_id
+4 -68
View File
@@ -1,70 +1,6 @@
from logging import getLogger
from time import time
import warnings
try:
import mlflow
except ImportError:
raise ImportError('Missing mlflow package.')
warnings.warn("`mlflow_logger` module has been renamed to `mlflow` since v0.5.3"
" and will be removed in v0.8.0", DeprecationWarning)
from .base import LightningLoggerBase, rank_zero_only
logger = getLogger(__name__)
class MLFlowLogger(LightningLoggerBase):
def __init__(self, experiment_name, tracking_uri=None, tags=None):
super().__init__()
self.experiment = mlflow.tracking.MlflowClient(tracking_uri)
self.experiment_name = experiment_name
self._run_id = None
self.tags = tags
@property
def run_id(self):
if self._run_id is not None:
return self._run_id
experiment = self.experiment.get_experiment_by_name(self.experiment_name)
if experiment is None:
logger.warning(
f"Experiment with name f{self.experiment_name} not found. Creating it."
)
self.experiment.create_experiment(self.experiment_name)
experiment = self.experiment.get_experiment_by_name(self.experiment_name)
run = self.experiment.create_run(experiment.experiment_id, tags=self.tags)
self._run_id = run.info.run_id
return self._run_id
@rank_zero_only
def log_hyperparams(self, params):
for k, v in vars(params).items():
self.experiment.log_param(self.run_id, k, v)
@rank_zero_only
def log_metrics(self, metrics, step_num=None):
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}"
)
continue
self.experiment.log_metric(self.run_id, k, v, timestamp_ms, step_num)
def save(self):
pass
@rank_zero_only
def finalize(self, status="FINISHED"):
if status == 'success':
status = 'FINISHED'
self.experiment.set_terminated(self.run_id, status)
@property
def name(self):
return self.experiment_name
@property
def version(self):
return self._run_id
from pytorch_lightning.logging.mlflow import MLFlowLogger # noqa: E402
+109
View File
@@ -0,0 +1,109 @@
try:
from test_tube import Experiment
except ImportError:
raise ImportError('Missing test-tube package.')
from .base import LightningLoggerBase, rank_zero_only
class TestTubeLogger(LightningLoggerBase):
__test__ = False
def __init__(
self, save_dir, name="default", description=None, debug=False,
version=None, create_git_tag=False
):
super().__init__()
self.save_dir = save_dir
self._name = name
self.description = description
self.debug = debug
self._version = version
self.create_git_tag = create_git_tag
self._experiment = None
@property
def experiment(self):
if self._experiment is not None:
return self._experiment
self._experiment = Experiment(
save_dir=self.save_dir,
name=self._name,
debug=self.debug,
version=self.version,
description=self.description,
create_git_tag=self.create_git_tag,
rank=self.rank,
)
return self._experiment
@rank_zero_only
def log_hyperparams(self, params):
# TODO: HACK figure out where this is being set to true
self.experiment.debug = self.debug
self.experiment.argparse(params)
@rank_zero_only
def log_metrics(self, metrics, step_num=None):
# TODO: HACK figure out where this is being set to true
self.experiment.debug = self.debug
self.experiment.log(metrics, global_step=step_num)
@rank_zero_only
def save(self):
# TODO: HACK figure out where this is being set to true
self.experiment.debug = self.debug
self.experiment.save()
@rank_zero_only
def finalize(self, status):
# TODO: HACK figure out where this is being set to true
self.experiment.debug = self.debug
self.save()
self.close()
@rank_zero_only
def close(self):
# TODO: HACK figure out where this is being set to true
self.experiment.debug = self.debug
exp = self.experiment
exp.close()
@property
def rank(self):
return self._rank
@rank.setter
def rank(self, value):
self._rank = value
if self._experiment is not None:
self.experiment.rank = value
@property
def name(self):
if self._experiment is None:
return self._name
else:
return self.experiment.name
@property
def version(self):
if self._experiment is None:
return self._version
else:
return self.experiment.version
# Test tube experiments are not pickleable, so we need to override a few
# methods to get DDP working. See
# https://docs.python.org/3/library/pickle.html#handling-stateful-objects
# for more info.
def __getstate__(self):
state = self.__dict__.copy()
state["_experiment"] = self.experiment.get_meta_copy()
return state
def __setstate__(self, state):
self._experiment = state["_experiment"].get_non_ddp_exp()
del state["_experiment"]
self.__dict__.update(state)
+4 -107
View File
@@ -1,109 +1,6 @@
try:
from test_tube import Experiment
except ImportError:
raise ImportError('Missing test-tube package.')
import warnings
from .base import LightningLoggerBase, rank_zero_only
warnings.warn("`test_tube_logger` module has been renamed to `test_tube` since v0.5.3"
" and will be removed in v0.8.0", DeprecationWarning)
class TestTubeLogger(LightningLoggerBase):
__test__ = False
def __init__(
self, save_dir, name="default", description=None, debug=False,
version=None, create_git_tag=False
):
super().__init__()
self.save_dir = save_dir
self._name = name
self.description = description
self.debug = debug
self._version = version
self.create_git_tag = create_git_tag
self._experiment = None
@property
def experiment(self):
if self._experiment is not None:
return self._experiment
self._experiment = Experiment(
save_dir=self.save_dir,
name=self._name,
debug=self.debug,
version=self.version,
description=self.description,
create_git_tag=self.create_git_tag,
rank=self.rank,
)
return self._experiment
@rank_zero_only
def log_hyperparams(self, params):
# TODO: HACK figure out where this is being set to true
self.experiment.debug = self.debug
self.experiment.argparse(params)
@rank_zero_only
def log_metrics(self, metrics, step_num=None):
# TODO: HACK figure out where this is being set to true
self.experiment.debug = self.debug
self.experiment.log(metrics, global_step=step_num)
@rank_zero_only
def save(self):
# TODO: HACK figure out where this is being set to true
self.experiment.debug = self.debug
self.experiment.save()
@rank_zero_only
def finalize(self, status):
# TODO: HACK figure out where this is being set to true
self.experiment.debug = self.debug
self.save()
self.close()
@rank_zero_only
def close(self):
# TODO: HACK figure out where this is being set to true
self.experiment.debug = self.debug
exp = self.experiment
exp.close()
@property
def rank(self):
return self._rank
@rank.setter
def rank(self, value):
self._rank = value
if self._experiment is not None:
self.experiment.rank = value
@property
def name(self):
if self._experiment is None:
return self._name
else:
return self.experiment.name
@property
def version(self):
if self._experiment is None:
return self._version
else:
return self.experiment.version
# Test tube experiments are not pickleable, so we need to override a few
# methods to get DDP working. See
# https://docs.python.org/3/library/pickle.html#handling-stateful-objects
# for more info.
def __getstate__(self):
state = self.__dict__.copy()
state["_experiment"] = self.experiment.get_meta_copy()
return state
def __setstate__(self, state):
self._experiment = state["_experiment"].get_non_ddp_exp()
del state["_experiment"]
self.__dict__.update(state)
from pytorch_lightning.logging.test_tube import TestTubeLogger # noqa: E402
@@ -0,0 +1,7 @@
import warnings
warnings.warn("`override_data_parallel` module has been renamed to `data_parallel` since v0.5.3"
" and will be removed in v0.8.0", DeprecationWarning)
from pytorch_lightning.overrides.data_parallel import ( # noqa: E402
get_a_var, parallel_apply, LightningDataParallel, LightningDistributedDataParallel)
@@ -0,0 +1,6 @@
import warnings
warnings.warn("`pt_overrides` package has been renamed to `overrides` since v0.5.3"
" and will be removed in v0.8.0", DeprecationWarning)
from pytorch_lightning.overrides import override_data_parallel # noqa: E402
@@ -0,0 +1,7 @@
import warnings
warnings.warn("`root_module` package has been renamed to `core` since v0.5.3"
" and will be removed in v0.8.0", DeprecationWarning)
from pytorch_lightning.core import ( # noqa: E402
decorators, grads, hooks, root_module, memory, model_saving)
+3 -3
View File
@@ -1,6 +1,6 @@
from .lm_test_module import LightningTestModel
from .lm_test_module_base import LightningTestModelBase
from .lm_test_module_mixins import (
from .test_module import LightningTestModel
from .test_module_base import LightningTestModelBase
from .test_module_mixins import (
LightningValidationStepMixin,
LightningValidationMixin,
LightningValidationStepMultipleDataloadersMixin,
@@ -1,7 +1,7 @@
import torch
from .lm_test_module_base import LightningTestModelBase
from .lm_test_module_mixins import LightningValidationMixin, LightningTestMixin
from .test_module_base import LightningTestModelBase
from .test_module_mixins import LightningValidationMixin, LightningTestMixin
class LightningTestModel(LightningValidationMixin, LightningTestMixin, LightningTestModelBase):
@@ -16,7 +16,7 @@ except ImportError:
raise ImportError('Missing test-tube package.')
from pytorch_lightning import data_loader
from pytorch_lightning.root_module.root_module import LightningModule
from pytorch_lightning.core.lightning import LightningModule
class LightningTestModelBase(LightningModule):
+4 -2
View File
@@ -1,7 +1,9 @@
import torch
from pytorch_lightning.pt_overrides.override_data_parallel import (
LightningDistributedDataParallel, LightningDataParallel)
from pytorch_lightning.overrides.data_parallel import (
LightningDistributedDataParallel,
LightningDataParallel,
)
from pytorch_lightning.utilities.debugging import MisconfigurationException
try:
+1 -1
View File
@@ -1,6 +1,6 @@
import torch
from pytorch_lightning.root_module import memory
from pytorch_lightning.core import memory
class TrainerLoggingMixin(object):
@@ -1,4 +1,4 @@
from pytorch_lightning.root_module.root_module import LightningModule
from pytorch_lightning.core.lightning import LightningModule
class TrainerModelHooksMixin(object):
+4 -2
View File
@@ -8,8 +8,10 @@ import logging
import torch
import torch.distributed as dist
from pytorch_lightning.pt_overrides.override_data_parallel import (
LightningDistributedDataParallel, LightningDataParallel)
from pytorch_lightning.overrides.data_parallel import (
LightningDistributedDataParallel,
LightningDataParallel,
)
class TrainerIOMixin(object):
+1 -1
View File
@@ -6,7 +6,7 @@ from pytorch_lightning import Trainer
from pytorch_lightning.callbacks import (
ModelCheckpoint,
)
from pytorch_lightning.root_module import memory
from pytorch_lightning.core import memory
from pytorch_lightning.testing import (
LightningTestModel,
)