updated test models with lazy decorators

This commit is contained in:
William Falcon
2019-07-25 10:56:03 -04:00
parent 39b15855ed
commit 24a3246bc1
5 changed files with 16 additions and 31 deletions
+2 -1
View File
@@ -1,2 +1,3 @@
from .models import Trainer
from .root_module.root_module import LightningModule
from .root_module.root_module import LightningModule
from .root_module.decorators import data_loader
@@ -10,6 +10,7 @@ from torch import optim
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
import pytorch_lightning as ptl
from pytorch_lightning.root_module.root_module import LightningModule
@@ -200,35 +201,17 @@ class LightningTemplateModel(LightningModule):
return loader
@property
@ptl.data_loader
def tng_dataloader(self):
if self._tng_dataloader is None:
try:
self._tng_dataloader = self.__dataloader(train=True)
except Exception as e:
print(e)
raise e
return self._tng_dataloader
return self.__dataloader(train=True)
@property
@ptl.data_loader
def val_dataloader(self):
if self._val_dataloader is None:
try:
self._val_dataloader = self.__dataloader(train=False)
except Exception as e:
print(e)
raise e
return self._val_dataloader
return self.__dataloader(train=False)
@property
@ptl.data_loader
def test_dataloader(self):
if self._test_dataloader is None:
try:
self._test_dataloader = self.__dataloader(train=False)
except Exception as e:
print(e)
raise e
return self._test_dataloader
return self.__dataloader(train=False)
@staticmethod
def add_model_specific_args(parent_parser, root_dir): # pragma: no cover
@@ -1,6 +1,6 @@
import torch.nn as nn
import numpy as np
from pytorch_lightning.root_module.root_module import LightningModule
from pytorch_lightning import LightningModule
from test_tube import HyperOptArgumentParser
from torchvision.datasets import MNIST
import torchvision.transforms as transforms
@@ -149,7 +149,7 @@ class ExampleModel1(LightningModule):
return loader
@property
@data_loader
def tng_dataloader(self):
if self._tng_dataloader is None:
try:
@@ -0,0 +1 @@
from .decorators import data_loader
+4 -4
View File
@@ -3,7 +3,7 @@ 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
import pytorch_lightning as ptl
class LightningModule(GradInformation, ModelIO, ModelHooks):
@@ -84,7 +84,7 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
for param in self.parameters():
param.requires_grad = True
@data_loader
@ptl.data_loader
def tng_dataloader(self):
"""
Implement a function to load an h5py of this data
@@ -92,7 +92,7 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
"""
raise NotImplementedError
@data_loader
@ptl.data_loader
def test_dataloader(self):
"""
Implement a function to load an h5py of this data
@@ -100,7 +100,7 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
"""
raise NotImplementedError
@data_loader
@ptl.data_loader
def val_dataloader(self):
"""
Implement a function to load an h5py of this data