mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
Refactor: name modules (#548)
* refactor: rename some modules * add deprecation warnings * fix paths
This commit is contained in:
committed by
William Falcon
parent
fea7cc87f6
commit
9785a3e78e
@@ -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):
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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"]
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
+2
-2
@@ -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):
|
||||
+1
-1
@@ -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):
|
||||
@@ -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,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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user