new way of passing dataloaders (#759)

* new way of passing dataloaders

* fixed docs

* fixed codestyle to follow flake8

* allow val/test be list of dataloaders and smarter checking

* added test

* fix flake error

* fix linking to new test model

* split into multiple test

* fix naming and typo

* minor documentation changes

* remove random file

* Update trainer.py

* Update trainer.py

* Update trainer.py

* Update trainer.py

* Update trainer.py

* Update trainer.py

* better error/warning message

* final adjustments

* update CHANGELOG.md

Co-authored-by: William Falcon <waf2107@columbia.edu>
This commit is contained in:
Nicki Skafte
2020-02-19 06:00:08 -05:00
committed by GitHub
co-authored by William Falcon
parent b9b5a93f0f
commit ffd6e693de
7 changed files with 267 additions and 13 deletions
+1 -1
View File
@@ -2,7 +2,7 @@
import torch
from .base import LightningTestModelBase
from .base import LightningTestModelBase, LightningTestModelBaseWithoutDataloader
from .mixins import (
LightningValidationStepMixin,
LightningValidationMixin,
+14 -6
View File
@@ -36,7 +36,7 @@ class TestingMNIST(MNIST):
self.targets = self.targets[:num_samples]
class LightningTestModelBase(LightningModule):
class TestModelBase(LightningModule):
"""
Base LightningModule for testing. Implements only the required
interface
@@ -48,7 +48,7 @@ class LightningTestModelBase(LightningModule):
:param hparams:
"""
# init superclass
super(LightningTestModelBase, self).__init__()
super(TestModelBase, self).__init__()
self.hparams = hparams
self.batch_size = hparams.batch_size
@@ -178,10 +178,6 @@ class LightningTestModelBase(LightningModule):
return loader
@data_loader
def train_dataloader(self):
return self._dataloader(train=True)
@staticmethod
def add_model_specific_args(parent_parser, root_dir): # pragma: no cover
"""
@@ -218,3 +214,15 @@ class LightningTestModelBase(LightningModule):
options=[32, 64, 128, 256], tunable=False,
help='batch size will be divided over all gpus being used across all nodes')
return parser
class LightningTestModelBase(TestModelBase):
""" with pre-defined train dataloader """
@data_loader
def train_dataloader(self):
return self._dataloader(train=True)
class LightningTestModelBaseWithoutDataloader(TestModelBase):
""" without pre-defined train dataloader """
pass