fix deprecated tng and abstract ligntning (#644)

This commit is contained in:
Jirka Borovec
2020-01-13 22:20:38 -05:00
committed by William Falcon
parent 8dc8a8bfd3
commit f7db44e750
4 changed files with 16 additions and 18 deletions
+13 -15
View File
@@ -1,6 +1,8 @@
import os
import warnings
import collections
import logging
from abc import ABC, abstractmethod
from argparse import Namespace
import torch
@@ -11,12 +13,11 @@ 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.trainer.training_io import load_hparams_from_tags_csv
from pytorch_lightning.overrides.data_parallel import LightningDistributedDataParallel
class LightningModule(GradInformation, ModelIO, ModelHooks):
class LightningModule(ABC, GradInformation, ModelIO, ModelHooks):
"""
A LightningModule has the following properties which you can access at any time
@@ -96,6 +97,7 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
self.use_ddp2 = False
self.use_amp = False
@abstractmethod
def forward(self, *args, **kwargs):
"""
Expand model in into whatever you need.
@@ -103,8 +105,8 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
:param x:
:return:
"""
raise NotImplementedError
@abstractmethod
def training_step(self, *args, **kwargs):
"""return loss, dict with metrics for tqdm
@@ -169,7 +171,6 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
You can also return a -1 instead of a dict to stop the current loop. This is useful
if you want to break out of the current training epoch early.
"""
raise NotImplementedError
def training_end(self, *args, **kwargs):
"""return loss, dict with metrics for tqdm
@@ -602,7 +603,7 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
# all ports should be in the 10k+ range
default_port = int(default_port) + 15000
except Exception as e:
except Exception:
default_port = 12910
# if user gave a port number, use that one instead
@@ -648,6 +649,7 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
return model, optimizers
@abstractmethod
def configure_optimizers(self):
"""Return a list of optimizers and a list of schedulers (could be empty)
@@ -694,7 +696,6 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
override the `optimizer_step` hook.
"""
raise NotImplementedError
def optimizer_step(self, epoch, batch_idx, optimizer, optimizer_idx, second_order_closure=None):
"""Do something instead of the standard optimizer behavior
@@ -817,6 +818,7 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
return splits
@data_loader
@abstractmethod
def train_dataloader(self):
"""Implement a PyTorch DataLoader
@@ -844,7 +846,6 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
"""
raise NotImplementedError
@data_loader
def tng_dataloader(self):
@@ -852,13 +853,10 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
.. warning:: Deprecated in v0.5.0. use train_dataloader instead.
"""
try:
output = self.tng_dataloader()
warnings.warn("`tng_dataloader` has been renamed to `train_dataloader` since v0.5.0"
" and will be removed in v0.8.0", DeprecationWarning)
return output
except NotImplementedError:
raise NotImplementedError
output = self.train_dataloader()
warnings.warn("`tng_dataloader` has been renamed to `train_dataloader` since v0.5.0"
" and will be removed in v0.8.0", DeprecationWarning)
return output
@data_loader
def test_dataloader(self):
+1 -1
View File
@@ -26,7 +26,7 @@ from pytorch_lightning.trainer.evaluation_loop import TrainerEvaluationLoopMixin
from pytorch_lightning.trainer.logging import TrainerLoggingMixin
from pytorch_lightning.trainer.model_hooks import TrainerModelHooksMixin
from pytorch_lightning.trainer.training_loop import TrainerTrainLoopMixin
from pytorch_lightning.trainer.trainer_io import TrainerIOMixin
from pytorch_lightning.trainer.training_io import TrainerIOMixin
from pytorch_lightning.trainer.training_tricks import TrainerTrainingTricksMixin
from pytorch_lightning.utilities.debugging import MisconfigurationException
+2 -2
View File
@@ -15,7 +15,7 @@ from pytorch_lightning.testing import (
LightningValidationMultipleDataloadersMixin,
LightningTestMultipleDataloadersMixin,
)
from pytorch_lightning.trainer import trainer_io
from pytorch_lightning.trainer import training_io
from pytorch_lightning.trainer.logging import TrainerLoggingMixin
@@ -187,7 +187,7 @@ def test_loading_meta_tags(tmpdir):
tags_path = logger.experiment.get_data_path(
logger.experiment.name, logger.experiment.version
) + '/meta_tags.csv'
tags = trainer_io.load_hparams_from_tags_csv(tags_path)
tags = training_io.load_hparams_from_tags_csv(tags_path)
assert tags.batch_size == 32 and tags.hidden_dim == 1000