diff --git a/pl_examples/basic_examples/lightning_module_template.py b/pl_examples/basic_examples/lightning_module_template.py index 8a892d10..eb1fe9f1 100644 --- a/pl_examples/basic_examples/lightning_module_template.py +++ b/pl_examples/basic_examples/lightning_module_template.py @@ -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): diff --git a/pytorch_lightning/__init__.py b/pytorch_lightning/__init__.py index 2e32e627..cbff186d 100644 --- a/pytorch_lightning/__init__.py +++ b/pytorch_lightning/__init__.py @@ -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', diff --git a/pytorch_lightning/callbacks/pt_callbacks.py b/pytorch_lightning/callbacks/pt_callbacks.py index 81b5be70..4940f00e 100644 --- a/pytorch_lightning/callbacks/pt_callbacks.py +++ b/pytorch_lightning/callbacks/pt_callbacks.py @@ -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): diff --git a/pytorch_lightning/core/__init__.py b/pytorch_lightning/core/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/pytorch_lightning/root_module/decorators.py b/pytorch_lightning/core/decorators.py similarity index 100% rename from pytorch_lightning/root_module/decorators.py rename to pytorch_lightning/core/decorators.py diff --git a/pytorch_lightning/root_module/grads.py b/pytorch_lightning/core/grads.py similarity index 100% rename from pytorch_lightning/root_module/grads.py rename to pytorch_lightning/core/grads.py diff --git a/pytorch_lightning/root_module/hooks.py b/pytorch_lightning/core/hooks.py similarity index 100% rename from pytorch_lightning/root_module/hooks.py rename to pytorch_lightning/core/hooks.py diff --git a/pytorch_lightning/root_module/root_module.py b/pytorch_lightning/core/lightning.py similarity index 96% rename from pytorch_lightning/root_module/root_module.py rename to pytorch_lightning/core/lightning.py index c5703e8b..536e10c3 100644 --- a/pytorch_lightning/root_module/root_module.py +++ b/pytorch_lightning/core/lightning.py @@ -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): diff --git a/pytorch_lightning/root_module/memory.py b/pytorch_lightning/core/memory.py similarity index 100% rename from pytorch_lightning/root_module/memory.py rename to pytorch_lightning/core/memory.py diff --git a/pytorch_lightning/core/model_saving.py b/pytorch_lightning/core/model_saving.py new file mode 100644 index 00000000..2a8fe521 --- /dev/null +++ b/pytorch_lightning/core/model_saving.py @@ -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 diff --git a/pytorch_lightning/core/root_module.py b/pytorch_lightning/core/root_module.py new file mode 100644 index 00000000..ecccdb62 --- /dev/null +++ b/pytorch_lightning/core/root_module.py @@ -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 diff --git a/pytorch_lightning/root_module/model_saving.py b/pytorch_lightning/core/saving.py similarity index 100% rename from pytorch_lightning/root_module/model_saving.py rename to pytorch_lightning/core/saving.py diff --git a/pytorch_lightning/logging/__init__.py b/pytorch_lightning/logging/__init__.py index dc47188a..d0644b65 100644 --- a/pytorch_lightning/logging/__init__.py +++ b/pytorch_lightning/logging/__init__.py @@ -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"] diff --git a/pytorch_lightning/logging/comet.py b/pytorch_lightning/logging/comet.py new file mode 100644 index 00000000..21dad421 --- /dev/null +++ b/pytorch_lightning/logging/comet.py @@ -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 diff --git a/pytorch_lightning/logging/comet_logger.py b/pytorch_lightning/logging/comet_logger.py index 21dad421..fff4f8d9 100644 --- a/pytorch_lightning/logging/comet_logger.py +++ b/pytorch_lightning/logging/comet_logger.py @@ -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 diff --git a/pytorch_lightning/logging/mlflow.py b/pytorch_lightning/logging/mlflow.py new file mode 100644 index 00000000..e9b2abd5 --- /dev/null +++ b/pytorch_lightning/logging/mlflow.py @@ -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 diff --git a/pytorch_lightning/logging/mlflow_logger.py b/pytorch_lightning/logging/mlflow_logger.py index e9b2abd5..ab7cd263 100644 --- a/pytorch_lightning/logging/mlflow_logger.py +++ b/pytorch_lightning/logging/mlflow_logger.py @@ -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 diff --git a/pytorch_lightning/logging/test_tube.py b/pytorch_lightning/logging/test_tube.py new file mode 100644 index 00000000..23e8806f --- /dev/null +++ b/pytorch_lightning/logging/test_tube.py @@ -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) diff --git a/pytorch_lightning/logging/test_tube_logger.py b/pytorch_lightning/logging/test_tube_logger.py index 23e8806f..f0aac342 100644 --- a/pytorch_lightning/logging/test_tube_logger.py +++ b/pytorch_lightning/logging/test_tube_logger.py @@ -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 diff --git a/pytorch_lightning/overrides/__init__.py b/pytorch_lightning/overrides/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/pytorch_lightning/pt_overrides/override_data_parallel.py b/pytorch_lightning/overrides/data_parallel.py similarity index 100% rename from pytorch_lightning/pt_overrides/override_data_parallel.py rename to pytorch_lightning/overrides/data_parallel.py diff --git a/pytorch_lightning/overrides/override_data_parallel.py b/pytorch_lightning/overrides/override_data_parallel.py new file mode 100644 index 00000000..6cb5b6c9 --- /dev/null +++ b/pytorch_lightning/overrides/override_data_parallel.py @@ -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) diff --git a/pytorch_lightning/pt_overrides/__init__.py b/pytorch_lightning/pt_overrides/__init__.py index e69de29b..fe0c3588 100644 --- a/pytorch_lightning/pt_overrides/__init__.py +++ b/pytorch_lightning/pt_overrides/__init__.py @@ -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 diff --git a/pytorch_lightning/root_module/__init__.py b/pytorch_lightning/root_module/__init__.py index e69de29b..9f13369b 100644 --- a/pytorch_lightning/root_module/__init__.py +++ b/pytorch_lightning/root_module/__init__.py @@ -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) diff --git a/pytorch_lightning/testing/__init__.py b/pytorch_lightning/testing/__init__.py index 8097c0b4..1e33f0e6 100644 --- a/pytorch_lightning/testing/__init__.py +++ b/pytorch_lightning/testing/__init__.py @@ -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, diff --git a/pytorch_lightning/testing/lm_test_module.py b/pytorch_lightning/testing/test_module.py similarity index 67% rename from pytorch_lightning/testing/lm_test_module.py rename to pytorch_lightning/testing/test_module.py index dd0992e0..703787e4 100644 --- a/pytorch_lightning/testing/lm_test_module.py +++ b/pytorch_lightning/testing/test_module.py @@ -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): diff --git a/pytorch_lightning/testing/lm_test_module_base.py b/pytorch_lightning/testing/test_module_base.py similarity index 98% rename from pytorch_lightning/testing/lm_test_module_base.py rename to pytorch_lightning/testing/test_module_base.py index 89bc3f29..c417e0e1 100644 --- a/pytorch_lightning/testing/lm_test_module_base.py +++ b/pytorch_lightning/testing/test_module_base.py @@ -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): diff --git a/pytorch_lightning/testing/lm_test_module_mixins.py b/pytorch_lightning/testing/test_module_mixins.py similarity index 100% rename from pytorch_lightning/testing/lm_test_module_mixins.py rename to pytorch_lightning/testing/test_module_mixins.py diff --git a/pytorch_lightning/trainer/dp_mixin.py b/pytorch_lightning/trainer/dp_mixin.py index efa10b49..1176b805 100644 --- a/pytorch_lightning/trainer/dp_mixin.py +++ b/pytorch_lightning/trainer/dp_mixin.py @@ -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: diff --git a/pytorch_lightning/trainer/logging_mixin.py b/pytorch_lightning/trainer/logging_mixin.py index 13e5afb6..0dae4da1 100644 --- a/pytorch_lightning/trainer/logging_mixin.py +++ b/pytorch_lightning/trainer/logging_mixin.py @@ -1,6 +1,6 @@ import torch -from pytorch_lightning.root_module import memory +from pytorch_lightning.core import memory class TrainerLoggingMixin(object): diff --git a/pytorch_lightning/trainer/model_hooks_mixin.py b/pytorch_lightning/trainer/model_hooks_mixin.py index 537d342d..23c82485 100644 --- a/pytorch_lightning/trainer/model_hooks_mixin.py +++ b/pytorch_lightning/trainer/model_hooks_mixin.py @@ -1,4 +1,4 @@ -from pytorch_lightning.root_module.root_module import LightningModule +from pytorch_lightning.core.lightning import LightningModule class TrainerModelHooksMixin(object): diff --git a/pytorch_lightning/trainer/trainer_io.py b/pytorch_lightning/trainer/trainer_io.py index 11bbb1d5..9dff1bf7 100644 --- a/pytorch_lightning/trainer/trainer_io.py +++ b/pytorch_lightning/trainer/trainer_io.py @@ -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): diff --git a/tests/test_gpu_models.py b/tests/test_gpu_models.py index 17f8a8fd..8d294f23 100644 --- a/tests/test_gpu_models.py +++ b/tests/test_gpu_models.py @@ -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, )