diff --git a/pytorch_lightning/callbacks/pt_callbacks.py b/pytorch_lightning/callbacks/pt_callbacks.py index f07c6a28..24d46e75 100644 --- a/pytorch_lightning/callbacks/pt_callbacks.py +++ b/pytorch_lightning/callbacks/pt_callbacks.py @@ -3,7 +3,7 @@ import shutil import numpy as np -from ..pt_overrides.override_data_parallel import LightningDistributedDataParallel +from pytorch_lightning.pt_overrides.override_data_parallel import LightningDistributedDataParallel class Callback(object): diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index 1309e4c5..2dbf46eb 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -13,11 +13,11 @@ from torch.utils.data.distributed import DistributedSampler import torch.multiprocessing as mp import torch.distributed as dist -from ..root_module.memory import get_gpu_memory_map -from ..root_module.model_saving import TrainerIO -from ..pt_overrides.override_data_parallel import ( +from pytorch_lightning.root_module.memory import get_gpu_memory_map +from pytorch_lightning.root_module.model_saving import TrainerIO +from pytorch_lightning.pt_overrides.override_data_parallel import ( LightningDistributedDataParallel, LightningDataParallel) -from ..utilities.debugging import MisconfigurationException +from pytorch_lightning.utilities.debugging import MisconfigurationException try: from apex import amp @@ -261,7 +261,7 @@ class Trainer(TrainerIO): if '.ckpt' in name: epoch = name.split('epoch_')[1] - epoch = int(re.sub('[^0-9]', '' ,epoch)) + epoch = int(re.sub('[^0-9]', '', epoch)) if epoch > last_epoch: last_epoch = epoch diff --git a/pytorch_lightning/root_module/model_saving.py b/pytorch_lightning/root_module/model_saving.py index ffd76387..7d6d2c8c 100644 --- a/pytorch_lightning/root_module/model_saving.py +++ b/pytorch_lightning/root_module/model_saving.py @@ -3,7 +3,7 @@ import re import torch -from ..pt_overrides.override_data_parallel import ( +from pytorch_lightning.pt_overrides.override_data_parallel import ( LightningDistributedDataParallel, LightningDataParallel) diff --git a/pytorch_lightning/root_module/root_module.py b/pytorch_lightning/root_module/root_module.py index 96dbfdcd..700e6db1 100644 --- a/pytorch_lightning/root_module/root_module.py +++ b/pytorch_lightning/root_module/root_module.py @@ -1,10 +1,10 @@ import torch -from .memory import ModelSummary -from .grads import GradInformation -from .model_saving import ModelIO, load_hparams_from_tags_csv -from .hooks import ModelHooks -from .decorators import data_loader +from pytorch_lightning.root_module.memory import ModelSummary +from pytorch_lightning.root_module.grads import GradInformation +from pytorch_lightning.root_module.model_saving import ModelIO, load_hparams_from_tags_csv +from pytorch_lightning.root_module.hooks import ModelHooks +from pytorch_lightning.root_module.decorators import data_loader class LightningModule(GradInformation, ModelIO, ModelHooks): diff --git a/pytorch_lightning/testing/lm_test_module.py b/pytorch_lightning/testing/lm_test_module.py index 61ecf874..8fe4cfc0 100644 --- a/pytorch_lightning/testing/lm_test_module.py +++ b/pytorch_lightning/testing/lm_test_module.py @@ -11,7 +11,7 @@ from torchvision.datasets import MNIST from torchvision import transforms from test_tube import HyperOptArgumentParser -from ..root_module.root_module import LightningModule +from pytorch_lightning.root_module.root_module import LightningModule from pytorch_lightning import data_loader diff --git a/tests/test_models.py b/tests/test_models.py index 746f97ef..f6e792fe 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -123,8 +123,6 @@ def test_amp_gpu_ddp(): run_gpu_model_test(trainer_options, model, hparams) - - def test_cpu_slurm_save_load(): """ Verify model save/load/checkpoint on CPU