mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-20 13:10:42 +08:00
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:
co-authored by
William Falcon
parent
b9b5a93f0f
commit
ffd6e693de
@@ -2,7 +2,7 @@
|
||||
|
||||
import torch
|
||||
|
||||
from .base import LightningTestModelBase
|
||||
from .base import LightningTestModelBase, LightningTestModelBaseWithoutDataloader
|
||||
from .mixins import (
|
||||
LightningValidationStepMixin,
|
||||
LightningValidationMixin,
|
||||
|
||||
+14
-6
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user