mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
Fixing tests (#936)
* abs import * rename test model * update trainer * revert test_step check * move tags * fix test_step * clean tests * fix template * update dataset path * fix parent order
This commit is contained in:
@@ -188,6 +188,8 @@ class LightningTemplateModel(pl.LightningModule):
|
||||
return [optimizer], [scheduler]
|
||||
|
||||
def __dataloader(self, train):
|
||||
# this is neede when you want some info about dataset before binding to trainer
|
||||
self.prepare_data()
|
||||
# init data generators
|
||||
transform = transforms.Compose([transforms.ToTensor(),
|
||||
transforms.Normalize((0.5,), (1.0,))])
|
||||
@@ -208,10 +210,8 @@ class LightningTemplateModel(pl.LightningModule):
|
||||
def prepare_data(self):
|
||||
transform = transforms.Compose([transforms.ToTensor(),
|
||||
transforms.Normalize((0.5,), (1.0,))])
|
||||
dataset = MNIST(root=self.hparams.data_root, train=True,
|
||||
transform=transform, download=True)
|
||||
dataset = MNIST(root=self.hparams.data_root, train=False,
|
||||
transform=transform, download=True)
|
||||
_ = MNIST(root=self.hparams.data_root, train=True,
|
||||
transform=transform, download=True)
|
||||
|
||||
def train_dataloader(self):
|
||||
log.info('Training data loader called.')
|
||||
|
||||
Reference in New Issue
Block a user