diff --git a/pl_examples/domain_templates/generative_adversarial_net.py b/pl_examples/domain_templates/generative_adversarial_net.py index cbe21fe2..99a57f1a 100644 --- a/pl_examples/domain_templates/generative_adversarial_net.py +++ b/pl_examples/domain_templates/generative_adversarial_net.py @@ -173,7 +173,7 @@ class GAN(LightningModule): # log sampled images sample_imgs = self(z) grid = torchvision.utils.make_grid(sample_imgs) - self.logger.experiment.add_image(f'generated_images', grid, self.current_epoch) + self.logger.experiment.add_image('generated_images', grid, self.current_epoch) def main(hparams): diff --git a/pytorch_lightning/core/model_saving.py b/pytorch_lightning/core/model_saving.py index 8c363023..13dbfd25 100644 --- a/pytorch_lightning/core/model_saving.py +++ b/pytorch_lightning/core/model_saving.py @@ -4,8 +4,7 @@ """ from pytorch_lightning.utilities import rank_zero_warn +from pytorch_lightning.core.saving import * # noqa: F403 rank_zero_warn("`model_saving` module has been renamed to `saving` since v0.6.0." " The deprecated module name will be removed in v0.8.0.", DeprecationWarning) - -from pytorch_lightning.core.saving import * # noqa: F403 diff --git a/pytorch_lightning/core/root_module.py b/pytorch_lightning/core/root_module.py index b8e602da..afbd8919 100644 --- a/pytorch_lightning/core/root_module.py +++ b/pytorch_lightning/core/root_module.py @@ -4,8 +4,7 @@ """ from pytorch_lightning.utilities import rank_zero_warn +from pytorch_lightning.core.lightning import * # noqa: F403 rank_zero_warn("`root_module` module has been renamed to `lightning` since v0.6.0." " The deprecated module name will be removed in v0.8.0.", DeprecationWarning) - -from pytorch_lightning.core.lightning import * # noqa: F403 diff --git a/pytorch_lightning/logging/__init__.py b/pytorch_lightning/logging/__init__.py index 9d027b34..7d0323a6 100644 --- a/pytorch_lightning/logging/__init__.py +++ b/pytorch_lightning/logging/__init__.py @@ -4,8 +4,7 @@ """ from pytorch_lightning.utilities import rank_zero_warn +from pytorch_lightning.loggers import * # noqa: F403 rank_zero_warn("`logging` package has been renamed to `loggers` since v0.7.0" " The deprecated package name will be removed in v0.9.0.", DeprecationWarning) - -from pytorch_lightning.loggers import * # noqa: F403 diff --git a/pytorch_lightning/logging/comet.py b/pytorch_lightning/logging/comet.py index ce854292..7fe59468 100644 --- a/pytorch_lightning/logging/comet.py +++ b/pytorch_lightning/logging/comet.py @@ -3,8 +3,7 @@ """ from pytorch_lightning.utilities import rank_zero_warn +from pytorch_lightning.loggers.comet import CometLogger # noqa: F403 rank_zero_warn("`logging.comet` module has been renamed to `loggers.comet` since v0.7.0." " The deprecated module name will be removed in v0.9.0.", DeprecationWarning) - -from pytorch_lightning.loggers.comet import CometLogger # noqa: F403 diff --git a/pytorch_lightning/logging/mlflow.py b/pytorch_lightning/logging/mlflow.py index 15b7fd81..75d1c0db 100644 --- a/pytorch_lightning/logging/mlflow.py +++ b/pytorch_lightning/logging/mlflow.py @@ -3,8 +3,7 @@ """ from pytorch_lightning.utilities import rank_zero_warn +from pytorch_lightning.loggers.mlflow import MLFlowLogger # noqa: F403 rank_zero_warn("`logging.mlflow` module has been renamed to `loggers.mlflow` since v0.7.0." " The deprecated module name will be removed in v0.9.0.", DeprecationWarning) - -from pytorch_lightning.loggers.mlflow import MLFlowLogger # noqa: F403 diff --git a/pytorch_lightning/logging/neptune.py b/pytorch_lightning/logging/neptune.py index af6e18c1..8d63cb48 100644 --- a/pytorch_lightning/logging/neptune.py +++ b/pytorch_lightning/logging/neptune.py @@ -3,8 +3,7 @@ """ from pytorch_lightning.utilities import rank_zero_warn +from pytorch_lightning.loggers.neptune import NeptuneLogger # noqa: F403 rank_zero_warn("`logging.neptune` module has been renamed to `loggers.neptune` since v0.7.0." " The deprecated module name will be removed in v0.9.0.", DeprecationWarning) - -from pytorch_lightning.loggers.neptune import NeptuneLogger # noqa: F403 diff --git a/pytorch_lightning/logging/test_tube.py b/pytorch_lightning/logging/test_tube.py index 3648db61..5c629e84 100644 --- a/pytorch_lightning/logging/test_tube.py +++ b/pytorch_lightning/logging/test_tube.py @@ -3,8 +3,7 @@ """ from pytorch_lightning.utilities import rank_zero_warn +from pytorch_lightning.loggers.test_tube import TestTubeLogger # noqa: F403 rank_zero_warn("`logging.test_tube` module has been renamed to `loggers.test_tube` since v0.7.0." " The deprecated module name will be removed in v0.9.0.", DeprecationWarning) - -from pytorch_lightning.loggers.test_tube import TestTubeLogger # noqa: F403 diff --git a/pytorch_lightning/logging/wandb.py b/pytorch_lightning/logging/wandb.py index 98a753c0..ee5ad1b6 100644 --- a/pytorch_lightning/logging/wandb.py +++ b/pytorch_lightning/logging/wandb.py @@ -3,8 +3,7 @@ """ from pytorch_lightning.utilities import rank_zero_warn +from pytorch_lightning.loggers.wandb import WandbLogger # noqa: F403 rank_zero_warn("`logging.wandb` module has been renamed to `loggers.wandb` since v0.7.0." " The deprecated module name will be removed in v0.9.0.", DeprecationWarning) - -from pytorch_lightning.loggers.wandb import WandbLogger # noqa: F403 diff --git a/pytorch_lightning/pt_overrides/override_data_parallel.py b/pytorch_lightning/pt_overrides/override_data_parallel.py index 34a65e3c..0515c206 100644 --- a/pytorch_lightning/pt_overrides/override_data_parallel.py +++ b/pytorch_lightning/pt_overrides/override_data_parallel.py @@ -4,9 +4,8 @@ """ from pytorch_lightning.utilities import rank_zero_warn +from pytorch_lightning.overrides.data_parallel import ( # noqa: F402 + get_a_var, parallel_apply, LightningDataParallel, LightningDistributedDataParallel) rank_zero_warn("`override_data_parallel` module has been renamed to `data_parallel` since v0.6.0." " The deprecated module name will be removed in v0.8.0.", DeprecationWarning) - -from pytorch_lightning.overrides.data_parallel import ( # noqa: F402 - get_a_var, parallel_apply, LightningDataParallel, LightningDistributedDataParallel) diff --git a/pytorch_lightning/root_module/decorators.py b/pytorch_lightning/root_module/decorators.py index 7031273b..2d4b9ef2 100644 --- a/pytorch_lightning/root_module/decorators.py +++ b/pytorch_lightning/root_module/decorators.py @@ -4,8 +4,7 @@ """ from pytorch_lightning.utilities import rank_zero_warn +from pytorch_lightning.core.decorators import * # noqa: F403 rank_zero_warn("`root_module.decorators` module has been renamed to `core.decorators` since v0.6.0." " The deprecated module name will be removed in v0.8.0.", DeprecationWarning) - -from pytorch_lightning.core.decorators import * # noqa: F403 diff --git a/pytorch_lightning/root_module/grads.py b/pytorch_lightning/root_module/grads.py index 81811411..b6ea9875 100644 --- a/pytorch_lightning/root_module/grads.py +++ b/pytorch_lightning/root_module/grads.py @@ -4,8 +4,7 @@ """ from pytorch_lightning.utilities import rank_zero_warn +from pytorch_lightning.core.grads import * # noqa: F403 rank_zero_warn("`root_module.grads` module has been renamed to `core.grads` since v0.6.0." " The deprecated module name will be removed in v0.8.0.", DeprecationWarning) - -from pytorch_lightning.core.grads import * # noqa: F403 diff --git a/pytorch_lightning/root_module/hooks.py b/pytorch_lightning/root_module/hooks.py index 0214a391..c6b492f6 100644 --- a/pytorch_lightning/root_module/hooks.py +++ b/pytorch_lightning/root_module/hooks.py @@ -4,8 +4,7 @@ """ from pytorch_lightning.utilities import rank_zero_warn +from pytorch_lightning.core.hooks import * # noqa: F403 rank_zero_warn("`root_module.hooks` module has been renamed to `core.hooks` since v0.6.0." " The deprecated module name will be removed in v0.8.0.", DeprecationWarning) - -from pytorch_lightning.core.hooks import * # noqa: F403 diff --git a/pytorch_lightning/root_module/memory.py b/pytorch_lightning/root_module/memory.py index 89d3d281..e63ed528 100644 --- a/pytorch_lightning/root_module/memory.py +++ b/pytorch_lightning/root_module/memory.py @@ -4,8 +4,7 @@ """ from pytorch_lightning.utilities import rank_zero_warn +from pytorch_lightning.core.memory import * # noqa: F403 rank_zero_warn("`root_module.memory` module has been renamed to `core.memory` since v0.6.0." " The deprecated module name will be removed in v0.8.0.", DeprecationWarning) - -from pytorch_lightning.core.memory import * # noqa: F403 diff --git a/pytorch_lightning/root_module/model_saving.py b/pytorch_lightning/root_module/model_saving.py index 67cf6a6a..3227905d 100644 --- a/pytorch_lightning/root_module/model_saving.py +++ b/pytorch_lightning/root_module/model_saving.py @@ -4,8 +4,7 @@ """ from pytorch_lightning.utilities import rank_zero_warn +from pytorch_lightning.core.saving import * # noqa: F403 rank_zero_warn("`root_module.model_saving` module has been renamed to `core.saving` since v0.6.0." " The deprecated module name will be removed in v0.8.0.", DeprecationWarning) - -from pytorch_lightning.core.saving import * # noqa: F403 diff --git a/pytorch_lightning/root_module/root_module.py b/pytorch_lightning/root_module/root_module.py index 3f3e9fad..b0e49860 100644 --- a/pytorch_lightning/root_module/root_module.py +++ b/pytorch_lightning/root_module/root_module.py @@ -4,8 +4,7 @@ """ from pytorch_lightning.utilities import rank_zero_warn +from pytorch_lightning.core.lightning import * # noqa: F403 rank_zero_warn("`root_module.root_module` module has been renamed to `core.lightning` since v0.6.0." " The deprecated module name will be removed in v0.8.0.", DeprecationWarning) - -from pytorch_lightning.core.lightning import * # noqa: F403 diff --git a/pytorch_lightning/trainer/optimizers.py b/pytorch_lightning/trainer/optimizers.py index ccabcddb..ea33e3e0 100644 --- a/pytorch_lightning/trainer/optimizers.py +++ b/pytorch_lightning/trainer/optimizers.py @@ -90,7 +90,7 @@ class TrainerOptimizersMixin(ABC): for scheduler in schedulers: if isinstance(scheduler, dict): if 'scheduler' not in scheduler: - raise ValueError(f'Lr scheduler should have key `scheduler`', + raise ValueError('Lr scheduler should have key `scheduler`', ' with item being a lr scheduler') scheduler['reduce_on_plateau'] = isinstance( scheduler['scheduler'], optim.lr_scheduler.ReduceLROnPlateau) diff --git a/pytorch_lightning/trainer/training_tricks.py b/pytorch_lightning/trainer/training_tricks.py index 61471a84..b8f29a4c 100644 --- a/pytorch_lightning/trainer/training_tricks.py +++ b/pytorch_lightning/trainer/training_tricks.py @@ -136,7 +136,7 @@ class TrainerTrainingTricksMixin(ABC): raise MisconfigurationException(f'Field {batch_arg_name} not found in `model.hparams`') if hasattr(model.train_dataloader, 'patch_loader_code'): - raise MisconfigurationException(f'The batch scaling feature cannot be used with dataloaders' + raise MisconfigurationException('The batch scaling feature cannot be used with dataloaders' ' passed directly to `.fit()`. Please disable the feature or' ' incorporate the dataloader into the model.') diff --git a/tests/models/test_module_hooks.py b/tests/models/test_module_hooks.py index 0a90a388..8b855ba4 100644 --- a/tests/models/test_module_hooks.py +++ b/tests/models/test_module_hooks.py @@ -9,6 +9,7 @@ import tests.base.utils as tutils def test_training_epoch_end_metrics_collection(tmpdir): """ Test that progress bar metrics also get collected at the end of an epoch. """ num_epochs = 3 + class CurrentModel(EvalModelTemplate): def training_step(self, *args, **kwargs): diff --git a/tests/trainer/test_dataloaders.py b/tests/trainer/test_dataloaders.py index 15940848..d157768f 100644 --- a/tests/trainer/test_dataloaders.py +++ b/tests/trainer/test_dataloaders.py @@ -11,6 +11,42 @@ from pytorch_lightning.utilities.exceptions import MisconfigurationException from tests.base import EvalModelTemplate +def test_fit_train_loader_only(tmpdir): + + model = EvalModelTemplate() + train_dataloader = model.train_dataloader() + + model.train_dataloader = None + model.val_dataloader = None + model.test_dataloader = None + + model.validation_step = None + model.validation_epoch_end = None + + model.test_step = None + model.test_epoch_end = None + + trainer = Trainer(fast_dev_run=True, default_root_dir=tmpdir) + trainer.fit(model, train_dataloader=train_dataloader) + + +def test_fit_val_loader_only(tmpdir): + + model = EvalModelTemplate() + train_dataloader = model.train_dataloader() + val_dataloader = model.val_dataloader() + + model.train_dataloader = None + model.val_dataloader = None + model.test_dataloader = None + + model.test_step = None + model.test_epoch_end = None + + trainer = Trainer(fast_dev_run=True, default_root_dir=tmpdir) + trainer.fit(model, train_dataloader=train_dataloader, val_dataloaders=val_dataloader) + + @pytest.mark.parametrize("dataloader_options", [ dict(train_percent_check=-0.1), dict(train_percent_check=1.1),